Files
qlib/backend/app/quant/service.py
T
Simon 9cc4bfccac feat(universe): B1-1 指数历史成分(index_weight)+ Universe 按 as_of 成分过滤
- index_weight 表(migration f5e0d1c2b3a4,MySQL 已应用;index_code+date+symbol 唯一)
  + IndexWeight 实体 + IndexConstituentRepository(members_at:取 <=as_of 最近一期快照,
  Survivorship-free / 无未来成分;latest_date)
- UniverseSpec.index_code + universe.filter_stocks members 交集 + resolve_members;
  Research/Selection/Signal/Replay 服务注入 index repo(历史成分过滤,选股/回测共用)
- tests/test_index_universe.py:快照历史成分(成分变更不入早期结果)、幂等、
  空快照期空集、index_code 过滤下 as_of 一致性;全量 pytest 通过
2026-09-09 07:27:13 +08:00

152 lines
5.5 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, resolve_members # 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,
index_repo=None,
) -> None:
self._stock_repo = stock_repo
self._daily_repo = daily_repo
self._engine = engine
self._index_repo = index_repo
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,
members=resolve_members(self._index_repo, spec.universe, 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,
)