Files
qlib/backend/app/quant/service.py
T
Simon 0d05bfd187 feat(factor): C1 因子相关性分析(横截面 Spearman 矩阵 + API)
- evaluation.factor_correlation_report:多因子共同日期 ∩ 后逐日横截面 Spearman 相关
  取均值 → FactorCorrelationReport(冗余剔除前置,v3 §12 Correlation→Redundancy)
- ResearchService.run_factor_correlation + POST /api/factor-correlations
  (universe/factors/period;与其它研究同装配口径)
- tests/test_factor_correlation.py:矩阵对角=1/近线性±相关符号/对称/无共同日期补零、
  API 冒烟;全量 pytest 通过
2026-09-09 07:31:49 +08:00

161 lines
6.0 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,
FactorCorrelationReport,
FactorTestReport,
ResearchSpec,
)
from app.domain.repositories.market import (
DailyBarRepository,
StockRepository,
)
from app.quant.composite import build_factor_panels
from app.quant.engine import QuantEngine
from app.quant.evaluation import factor_correlation_report
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 run_factor_correlation(self, spec: ResearchSpec) -> FactorCorrelationReport:
"""多因子两两相关(v3 §12):同 universe/period 装配 → 横截面相关矩阵。"""
daily = self._load_daily(spec)
panels = {fs.name: build_factor_panels(daily, [fs])[0][1] for fs in spec.factors}
return factor_correlation_report(panels)
# ---- 数据装配 ----
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,
)