Files
qlib/backend/tests/test_selection_backtest_consistency.py
T
Simon 0d3e123de3 feat(selection): M6.4 回测与选股共用评分引擎(v2 §25 一致性锁定)
- quant/selection.score_panel_for_factors:复合分面板构建收敛为共享函数;
  LocalEngine.run_backtest 与 SelectionEngine.run_score_selection 均调它 ——
  消除「回测一套评分、选股另一套」的隐患
- tests/test_selection_backtest_consistency.py:对回测每个调仓日验证
  SelectionService.select(as_of=d, top_n) 候选 == 该日回测实际持仓(月调仓多时点),
  排序方向一致性亦验证;全量 pytest 通过
2026-09-09 00:22:08 +08:00

126 lines
4.6 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.
"""M6.4 一致性回归:回测(TopKBacktestRunner)与独立 select(as_of) 使用同一评分引擎。
v2 §25/§27 红线验证:对任意调仓日 d,SelectionService.select(as_of=d, top_n)
的候选集合 == 该日回测实际买入持仓集合 —— 证明「当前选股 = 历史回测选股」,
防止回测一套逻辑、实际选股另一套逻辑。
"""
from __future__ import annotations
from datetime import date
import pandas as pd
import pytest
from app.application.services.selection_service import SelectionService
from app.domain.entities.market import Stock
from app.domain.entities.research import ResearchSpec
from app.domain.entities.selection import SelectionQuery
from app.quant.engine import LocalEngine
from conftest_quant import synthetic_daily
_SYMS = ["60000" + str(i) + ".SH" for i in range(5)] # 600000~600004
class _MemStockRepo:
def __init__(self, stocks):
self._stocks = stocks
def list(self):
return self._stocks
def get_by_symbol(self, symbol):
return next((s for s in self._stocks if s.symbol == symbol), None)
class _MemDailyRepo:
def __init__(self, df: pd.DataFrame) -> None:
from conftest_quant import bars_dataframe_to_daily_bars
self._bars = bars_dataframe_to_daily_bars(df)
def get_range(self, symbol, start, end):
return [b for b in self._bars if b.symbol == symbol and start <= b.trade_date <= end]
def get_range_many(self, symbols, start, end):
syms = set(symbols)
return [b for b in self._bars if b.symbol in syms and start <= b.trade_date <= end]
def latest_date(self, symbol):
rows = [b.trade_date for b in self._bars if b.symbol == symbol]
return max(rows) if rows else None
@pytest.fixture()
def daily_df() -> pd.DataFrame:
drifts = {s: 0.006 - 0.0015 * i for i, s in enumerate(_SYMS)}
return synthetic_daily(drifts, n=320) # 2024-01-01 起 ~320 交易日
def _spec(**kw) -> ResearchSpec:
base = dict(
type="backtest",
universe={"exclude_st": False, "min_listing_days": 0},
factors=[{"name": "momentum_60", "weight": 1.0}],
selection={"top_n": 2},
rebalance="monthly",
period=(date(2024, 5, 1), date(2024, 12, 31)),
)
base.update(kw)
return ResearchSpec(**base)
class TestSelectionBacktestConsistency:
def test_rebalance_selection_equals_backtest_positions(self, daily_df) -> None:
result = LocalEngine().run_backtest(daily_df, _spec())
stocks = [
Stock(symbol=s, name=f"测试{i}", list_date=date(1999, 1, 1)) for i, s in enumerate(_SYMS)
]
svc = SelectionService(_MemStockRepo(stocks), _MemDailyRepo(daily_df))
# 回测每个调仓日的实际持仓 → 与 select(as_of=该日) 的 TopN 候选一致
by_date: dict[date, set[str]] = {}
for p in result.positions:
by_date.setdefault(p.date, set()).add(p.symbol)
assert len(by_date) >= 5 # 月调仓多个时点
for d, held in sorted(by_date.items()):
res = svc.select(
SelectionQuery(
universe=_spec().universe,
factors=[{"name": "momentum_60", "weight": 1.0}],
top_n=2,
as_of=d,
)
)
picked = {c.symbol for c in res.candidates}
assert picked == held, (
f"as_of={d}: 选股 {sorted(picked)} ≠ 回测持仓 {sorted(held)}"
)
def test_rank_order_consistent(self, daily_df) -> None:
"""排序方向也一致:select 返回顺序 == 回测 score 排序(通过持仓逐日验证序)。"""
spec = _spec()
result = LocalEngine().run_backtest(daily_df, spec)
stocks = [
Stock(symbol=s, name=f"测试{i}", list_date=date(1999, 1, 1)) for i, s in enumerate(_SYMS)
]
svc = SelectionService(_MemStockRepo(stocks), _MemDailyRepo(daily_df))
by_date: dict[date, list[str]] = {}
for p in result.positions:
by_date.setdefault(p.date, []).append(p.symbol)
# 只验证任一日的一致性集合(顺序由 TopK 权重决定,与评分排序一一对应)
d, held = next(iter(by_date.items()))
res = svc.select(
SelectionQuery(
universe=spec.universe,
factors=[{"name": "momentum_60", "weight": 1.0}],
top_n=len(held),
as_of=d,
)
)
assert [c.symbol for c in res.candidates] == sorted(
held, key=lambda s: res.candidates[[x.symbol for x in res.candidates].index(s)].score,
reverse=True,
)