feat(selection): M6.0 选股契约与评分引擎(SelectionQuery/Result + select(as_of))
- domain/entities/selection.py:SelectionQuery(universe+method+factors+top_n/top_pct/ min_score+as_of+预热)与 SelectionResult/Candidate/Statistics(v2 §14.2/§21.1 DTO); ConditionSpec 字段就位供 M6.2 条件选股 - quant/selection.py:Selection Engine method=score —— 复合分(zscore×权重×方向) → TopN/Top% 截断;observation_date=<=as_of 最近交易日(防未来函数,v2 §9); 候选带 factor_values 与 selection_reason(可解释) - application/services/selection_service.py:选股用例(universe 过滤 → 装配 → 引擎) - quant/service.py:抽取公共 load_daily_df 供研究/选股共用(行为不变) - tests/test_selection.py:11 例 —— TopN/排序/理由、as_of 防未来函数、ST/上市天数/ 退市过滤、top_pct/min_score、空数据与查询校验;全量 pytest 通过
This commit is contained in:
@@ -88,6 +88,36 @@ def _frame_from_stream(rows: Iterable[tuple], columns: list[str]) -> pd.DataFram
|
||||
return df
|
||||
|
||||
|
||||
def load_daily_df(
|
||||
daily_repo,
|
||||
symbols: list[str],
|
||||
start: date,
|
||||
end: date,
|
||||
columns: list[str],
|
||||
) -> 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)), 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))
|
||||
else: # 兜底:逐只查询
|
||||
bars = []
|
||||
for sym in symbols:
|
||||
bars.extend(daily_repo.get_range(sym, start, end))
|
||||
return bars_to_daily_df(bars)
|
||||
|
||||
|
||||
class ResearchService:
|
||||
"""研究用例入口(因子测试 / 回测)。依赖注入 Repository 与引擎。"""
|
||||
|
||||
@@ -120,26 +150,8 @@ class ResearchService:
|
||||
# 回测前预留因子 warmup(lookback≤120 交易日,取 300 自然日余量)
|
||||
data_start = start - timedelta(days=300)
|
||||
stocks = filter_stocks(self._stock_repo.list(), spec.universe, as_of=start)
|
||||
if not stocks:
|
||||
return pd.DataFrame()
|
||||
symbols = [s.symbol for s in stocks]
|
||||
|
||||
# 引擎所需列裁剪(LocalEngine 只取 close + 因子字段;Qlib 回测取全 OHLCV)
|
||||
required = self._engine.required_columns(spec)
|
||||
streamer = getattr(self._daily_repo, "stream_range_many_columns", None)
|
||||
if streamer is not None:
|
||||
try:
|
||||
return _frame_from_stream(
|
||||
streamer(symbols, data_start, end, sorted(required)), sorted(required)
|
||||
)
|
||||
except Exception: # noqa: BLE001 —— 流式路径失败回退旧路径(兼容非 SQL 实现)
|
||||
pass
|
||||
# 旧路径:逐实体(供内存 / Fake 仓储等实现使用)
|
||||
get_many = getattr(self._daily_repo, "get_range_many", None)
|
||||
if get_many is not None:
|
||||
bars = list(get_many(symbols, data_start, end))
|
||||
else: # 兜底:逐只查询
|
||||
bars = []
|
||||
for s in stocks:
|
||||
bars.extend(self._daily_repo.get_range(s.symbol, data_start, end))
|
||||
return bars_to_daily_df(bars)
|
||||
return load_daily_df(
|
||||
self._daily_repo, [s.symbol for s in stocks], data_start, end, sorted(required)
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user