Files
qlib/backend/app/quant/service.py
T
Simon ef09d5b419 feat(quant): M7.3 研究行情口径显式化(默认不复权 none,可切 qfq)
- DailyBarRepository.get_range_many / stream_range_many_columns 增加 adjust 参数
  (默认 'none')→ SQL 层过滤口径,消除 stock_daily 混 source/adjust 污染因子的风险
- ResearchSpec / SelectionQuery 增加 price_adjustment(none|qfq),随 config_snapshot
  落库可溯源;ResearchService._load_daily 与 SelectionService 装配按口径取数
- tests/test_price_adjustment.py:repo 读取按 adjust 过滤(none/qfq 各自命中)、
  spec 默认与字段记录;全量 pytest 通过
2026-09-09 00:32:55 +08:00

147 lines
5.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""研究服务:把 Research Specification 编排为数据获取 + 引擎执行。
本层是业务入口:API / Agent 只能调用这里的用例(AGENT.md §16/§17),
禁止直接拼接引擎配置。数据一律经 Repository 获取(防未来函数由查询层保证)。
内存优化:大数据面板优先走 Repository 的流式列裁剪查询
(stream_range_many_columns,SQL 侧转 REAL、分批拉取),避免 ORM 对象 /
Decimal 全量物化;老实现回退到 get_range_many 逐实体路径。
"""
from __future__ import annotations
from collections.abc import Iterable
from datetime import date, timedelta
import pandas as pd
from app.domain.entities.research import (
BacktestResult,
FactorTestReport,
ResearchSpec,
)
from app.domain.repositories.market import (
DailyBarRepository,
StockRepository,
)
from app.quant.engine import QuantEngine
from app.quant.universe import filter_stocks # noqa: F401 —— 选股/回测共用范围过滤
# 流式路径每攒多少行落一个 DataFrame 分片(控制 concat 峰值)
_FRAME_CHUNK_ROWS = 50_000
def bars_to_daily_df(bars) -> pd.DataFrame:
"""DailyBar 列表 → 引擎长表 DataFrame(symbol/trade_date/ohlc/volume/amount)。
领域实体中的 Decimal 在此转 float,供 pandas 数值运算(保持 DataFrame 全数值列)。
"""
df = pd.DataFrame([b.model_dump() for b in bars])
if not df.empty:
for col in ("open", "high", "low", "close", "volume", "amount"):
if col in df.columns:
df[col] = df[col].astype(float)
return df
def _frame_from_stream(rows: Iterable[tuple], columns: list[str]) -> pd.DataFrame:
"""把流式 (symbol, trade_date_iso, *float_cols) 分批拼成 float 长表。
全程只保留分批 DataFrame + 最终一份结果,避免整批 tuple/Decimal 同时驻留。
返回列:symbol, trade_date(datetime64), *columns(float64)。
"""
cols = ["symbol", "trade_date", *columns]
buf: list[tuple] = []
pieces: list[pd.DataFrame] = []
for row in rows:
buf.append(row)
if len(buf) >= _FRAME_CHUNK_ROWS:
pieces.append(pd.DataFrame(buf, columns=cols))
buf = []
if buf:
pieces.append(pd.DataFrame(buf, columns=cols))
if not pieces:
return pd.DataFrame()
df = pd.concat(pieces, ignore_index=True)
df["trade_date"] = pd.to_datetime(df["trade_date"])
for col in columns: # NULL → NaN,统一 float64
df[col] = pd.to_numeric(df[col], errors="coerce")
return df
def load_daily_df(
daily_repo,
symbols: list[str],
start: date,
end: date,
columns: list[str],
adjust: str = "none",
) -> pd.DataFrame:
"""从 Repository 装配行情长表(供研究/选股共用)。
优先走流式列裁剪(stream_range_many_columns,SQL 侧转 REAL、分批),
失败或实现缺失时回退 get_range_many / 逐只 get_range。
"""
if not symbols:
return pd.DataFrame()
streamer = getattr(daily_repo, "stream_range_many_columns", None)
if streamer is not None:
try:
return _frame_from_stream(
streamer(symbols, start, end, sorted(columns), adjust=adjust), sorted(columns)
)
except Exception: # noqa: BLE001 —— 流式路径失败回退旧路径(兼容非 SQL 实现)
pass
get_many = getattr(daily_repo, "get_range_many", None)
if get_many is not None:
bars = list(get_many(symbols, start, end, adjust=adjust))
else: # 兜底:逐只查询
bars = []
for sym in symbols:
bars.extend(daily_repo.get_range(sym, start, end))
return bars_to_daily_df(bars)
class ResearchService:
"""研究用例入口(因子测试 / 回测)。依赖注入 Repository 与引擎。"""
def __init__(
self,
stock_repo: StockRepository,
daily_repo: DailyBarRepository,
engine: QuantEngine,
) -> None:
self._stock_repo = stock_repo
self._daily_repo = daily_repo
self._engine = engine
def run_factor_test(self, spec: ResearchSpec, horizon_days: int = 21) -> FactorTestReport:
if spec.type != "factor_test":
raise ValueError("factor_test 用例需要 spec.type=factor_test")
daily = self._load_daily(spec)
return self._engine.run_factor_test(daily, spec, horizon_days=horizon_days)
def run_backtest(self, spec: ResearchSpec) -> BacktestResult:
if spec.type != "backtest":
raise ValueError("backtest 用例需要 spec.type=backtest")
daily = self._load_daily(spec)
return self._engine.run_backtest(daily, spec)
# ---- 数据装配 ----
def _load_daily(self, spec: ResearchSpec) -> pd.DataFrame:
start, end = spec.period
# 回测前预留因子 warmup(lookback≤120 交易日,取 300 自然日余量)
data_start = start - timedelta(days=300)
stocks = filter_stocks(self._stock_repo.list(), spec.universe, as_of=start)
# 引擎所需列裁剪(LocalEngine 只取 close + 因子字段;Qlib 回测取全 OHLCV)
required = self._engine.required_columns(spec)
return load_daily_df(
self._daily_repo,
[s.symbol for s in stocks],
data_start,
end,
sorted(required),
adjust=spec.price_adjustment,
)