Files
qlib/backend/app/quant/selection.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

377 lines
14 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.
"""Selection Engine(ARCHITECTURE_v2 §14)—— 纯 pandas 执行层。
当前实现 method=score:因子加权复合分 → TopN/Top% 截断,输出 SelectionResult。
M6.2 在同一模块加入 method=condition(结构化条件选股)。
未来函数纪律:面板只在 <= observation_date 的数据上计算;observation_date 是
<= as_of 的最近可用交易日(as_of 显式传入即历史选股,None 则到数据最新)。
data 长表由 Service 装配(已按 universe 过滤 symbol、含预热窗口)。
"""
from __future__ import annotations
from datetime import date
import pandas as pd
from app.domain.entities.market import FinancialIndicator
from app.domain.entities.selection import (
SelectionCandidate,
SelectionQuery,
SelectionResult,
SelectionStatistics,
)
from app.quant.factors import FactorError, compute_factor, get_factor
from app.quant.local_engine import build_factor_panels, composite_score
_UNIMPLEMENTED_DEFAULT = [
"exclude_suspended 依赖停牌数据,当前未建模(结果可能包含停牌股)",
]
def score_panel_for_factors(daily: pd.DataFrame, factor_specs) -> pd.DataFrame:
"""因子加权复合分面板(index=trade_date, columns=symbol)。
回测(LocalEngine)与选股(run_score_selection)共用同一构建 ——
保证 v2 §25/§27「历史回测与当前选股使用同一套引擎」的一致性。
"""
panels = build_factor_panels(daily, factor_specs) # 未知因子在此抛 FactorError
return composite_score(panels)
def resolve_observation_date(daily: pd.DataFrame, as_of: date | None) -> pd.Timestamp | None:
"""<= as_of 的最近可用交易日;as_of=None 取数据最新一日。"""
if daily.empty:
return None
dates = pd.to_datetime(daily["trade_date"])
if as_of is None:
return dates.max()
avail = dates[dates <= pd.Timestamp(as_of)]
return avail.max() if len(avail) else None
def factor_columns(query: SelectionQuery) -> set[str]:
"""score 模式所需行情数值列(数据装配裁剪用)。"""
needed = {"close"}
for fs in query.factors:
try:
defn, _fn = get_factor(fs.name)
except FactorError:
continue # 未知因子由执行期统一报错(score_selection 中 build_factor_panels)
needed.update(defn.requires)
return needed
def run_score_selection(
daily: pd.DataFrame,
query: SelectionQuery,
as_of: date | None,
) -> SelectionResult:
"""因子评分选股(v2 §14.1B):复合分 → 排序 → TopN/Top%。"""
if query.method != "score":
raise ValueError(f"run_score_selection 需要 method=score,当前 {query.method}")
obs = resolve_observation_date(daily, as_of)
if obs is None:
resolved = as_of or date.today()
return SelectionResult(
as_of_date=resolved,
method=query.method,
statistics=SelectionStatistics(),
candidates=[],
unimplemented=list(_UNIMPLEMENTED_DEFAULT),
config_snapshot=query.model_dump(mode="json"),
)
resolved = obs.date()
# 只允许使用 <= obs 的数据(面板计算在截断后数据上进行)
view = daily[pd.to_datetime(daily["trade_date"]) <= obs]
if view.empty:
return SelectionResult(
as_of_date=resolved,
method=query.method,
statistics=SelectionStatistics(),
candidates=[],
unimplemented=list(_UNIMPLEMENTED_DEFAULT),
config_snapshot=query.model_dump(mode="json"),
)
score = score_panel_for_factors(view, query.factors).loc[obs].dropna().sort_values(
ascending=False
)
# 每因子在 obs 行的原始值(factor_values 供展示与解释;与 build_factor_panels 同数据)
raw: dict[str, pd.Series] = {}
for fs in query.factors:
_defn, panel = compute_factor(fs.name, view)
if obs in panel.index:
raw[fs.name] = panel.loc[obs]
candidates_df = _truncate(score, query)
evaluated = int(len(score)) # score 已 dropna,长度即有分股票数
candidates: list[SelectionCandidate] = []
for rank, (sym, sc) in enumerate(candidates_df.items(), start=1):
factor_values = {
name: _to_float(series.get(sym))
for name, series in raw.items()
if isinstance(series, pd.Series)
}
factor_values = {k: v for k, v in factor_values.items() if v is not None}
candidates.append(
SelectionCandidate(
symbol=sym,
rank=rank,
score=round(float(sc), 6),
factor_values=factor_values,
selection_reason=_score_reason(query, sym, raw),
)
)
return SelectionResult(
as_of_date=resolved,
method=query.method,
statistics=SelectionStatistics(
universe_size=_symbol_count(view),
evaluated=evaluated,
selected=len(candidates),
),
candidates=candidates,
unimplemented=list(_UNIMPLEMENTED_DEFAULT),
config_snapshot=query.model_dump(mode="json"),
)
def _truncate(score: pd.Series, query: SelectionQuery) -> pd.Series:
"""按 top_n / top_pct / min_score 截断(入参已按分数降序)。"""
s = score
if query.min_score is not None:
s = s[s >= query.min_score]
if query.top_pct is not None:
n = max(int(round(len(s) * query.top_pct)), 1)
s = s.head(n)
elif query.top_n is not None:
s = s.head(query.top_n)
return s
def _score_reason(query: SelectionQuery, symbol: str, raw: dict[str, pd.Series]) -> list[str]:
"""生成可读的入选理由:列每个因子的观测值与权重。"""
reasons: list[str] = []
for fs in query.factors:
try:
defn, _fn = get_factor(fs.name)
except FactorError:
continue
series = raw.get(fs.name)
val = _to_float(series.get(symbol)) if isinstance(series, pd.Series) else None
if val is None:
continue
good = defn.direction == "higher_is_better"
reasons.append(
f"{fs.name}={val:.4f}(权重 {fs.weight},{'越高越好' if good else '越低越好'})"
)
return reasons
def _symbol_count(daily: pd.DataFrame) -> int:
return int(daily["symbol"].nunique()) if not daily.empty and "symbol" in daily else 0
# ---------- method=condition:结构化条件选股(M6.2) ----------
# 技术字段:预计算派生量 + 行情原列(原列需在装配列中才可用)
_TECH_DERIVED = ("ma20", "ma60")
_STATIC_PREFIX = "static."
_FUNDAMENTAL_PREFIX = "fundamental."
def condition_needed_columns(query: SelectionQuery) -> set[str]:
"""条件引用的行情列(fundamental/static 走元数据与财务表,不需要行情列)。"""
needed = {"close"}
names = [c.field for c in query.conditions] + [
c.ref for c in query.conditions if c.ref and not c.ref.startswith(_FUNDAMENTAL_PREFIX)
]
for f in names:
if not f or f.startswith((_STATIC_PREFIX, _FUNDAMENTAL_PREFIX)):
continue
if f in {"open", "high", "low", "close", "volume", "amount", *_TECH_DERIVED}:
if f not in _TECH_DERIVED:
needed.add(f)
continue
try: # 其余按已注册因子处理
defn, _fn = get_factor(f)
except FactorError:
raise ValueError(
f"条件字段未知:{f}(可用: 行情列/ma20/ma60/已注册因子/static.*/fundamental.*)"
) from None
needed.update(defn.requires)
return needed
def run_condition_selection(
daily: pd.DataFrame,
stocks: list,
query: SelectionQuery,
as_of: date | None,
financial: dict[str, FinancialIndicator] | None = None,
) -> SelectionResult:
"""条件选股(v2 §14.1A):全部条件 AND 通过者入选(无排序;truncation 不适用)。
fields 域:static.*(股票基础)、close/volume/amount/ma20/ma60/已注册因子(行情)、
fundamental.*(announce_date <= as_of 的最新已公告财务值 —— 防未来函数由 Service 取数保证)。
"""
if query.method != "condition":
raise ValueError(f"run_condition_selection 需要 method=condition,当前 {query.method}")
obs = resolve_observation_date(daily, as_of)
resolved = (obs.date() if obs is not None else as_of) or date.today()
if obs is None or daily.empty:
return SelectionResult(
as_of_date=resolved, method=query.method,
statistics=SelectionStatistics(), candidates=[],
unimplemented=list(_UNIMPLEMENTED_DEFAULT),
config_snapshot=query.model_dump(mode="json"),
)
view = daily[pd.to_datetime(daily["trade_date"]) <= obs]
close = view.pivot(index="trade_date", columns="symbol", values="close").sort_index()
close.index = pd.to_datetime(close.index)
# 技术字段面板(obs 行)
tech: dict[str, pd.Series] = {}
for col in ("close", "open", "high", "low", "volume", "amount"):
if col in view.columns and col != "close":
panel = view.pivot(index="trade_date", columns="symbol", values=col).sort_index()
panel.index = pd.to_datetime(panel.index)
tech[col] = panel.loc[obs]
tech["close"] = close.loc[obs]
tech["ma20"] = close.rolling(20).mean().loc[obs]
tech["ma60"] = close.rolling(60).mean().loc[obs]
# 因子字段按需计算
for cond in query.conditions:
for f in (cond.field, cond.ref):
if f is None or f.startswith((_STATIC_PREFIX, _FUNDAMENTAL_PREFIX)) or f in tech:
continue
if f in _TECH_DERIVED or f in ("close", "open", "high", "low", "volume", "amount"):
continue
try:
_defn, panel = compute_factor(f, view)
except FactorError:
continue # 已在 condition_needed_columns 报错;此处防御
if obs in panel.index:
tech[f] = panel.loc[obs]
statics = {s.symbol: s.model_dump() for s in stocks}
candidates: list[SelectionCandidate] = []
passed_symbols: list[str] = []
for sym in sorted(statics):
statuses: list[str] = []
all_ok = True
for cond in query.conditions:
ok = _eval_condition(cond, sym, statics, tech, financial or {})
statuses.append(f"{cond.field} {cond.op} {cond.ref or cond.value}: {'通过' if ok else '未通过'}")
all_ok = all_ok and ok
if all_ok:
passed_symbols.append(sym)
candidates.append(
SelectionCandidate(
symbol=sym,
rank=0, # 占位,末尾统一编号
score=1.0,
filter_status=statuses,
selection_reason=[f"通过全部 {len(query.conditions)} 条条件"],
)
)
for rank, c in enumerate(candidates, start=1):
c.rank = rank
return SelectionResult(
as_of_date=resolved,
method=query.method,
statistics=SelectionStatistics(
universe_size=len(statics),
evaluated=len(statics),
selected=len(candidates),
),
candidates=candidates,
unimplemented=list(_UNIMPLEMENTED_DEFAULT) + [
"条件选股为纯过滤(AND),未排序/未截断;如需排序请在 factors 中提供评分",
],
config_snapshot=query.model_dump(mode="json"),
)
def _eval_condition(
cond,
sym: str,
statics: dict,
tech: dict[str, pd.Series],
financial: dict,
) -> bool:
"""求值单条条件:value 与 ref 二选一;left 与 right 同为 field 或 field vs 字面量。"""
left = _field_value(cond.field, sym, statics, tech, financial)
if cond.ref is not None:
right = _field_value(cond.ref, sym, statics, tech, financial)
else:
right = cond.value
return _compare(left, right, cond.op)
def _field_value(field, sym, statics, tech, financial):
if field.startswith(_STATIC_PREFIX):
return statics.get(sym, {}).get(field[len(_STATIC_PREFIX):])
if field.startswith(_FUNDAMENTAL_PREFIX):
fin = financial.get(sym)
return getattr(fin, field[len(_FUNDAMENTAL_PREFIX):], None) if fin else None
series = tech.get(field)
if series is None:
return None
v = series.get(sym)
return None if v is None or (isinstance(v, float) and v != v) else v # NaN → None
def _compare(left, right, op: str) -> bool:
"""混合比较:None 视为不可用 → 除 ne 外不通过;数值/字符串分别处理。"""
if op == "ne":
return left != right
if left is None or right is None:
return False
try:
if isinstance(left, (int, float)) or isinstance(right, (int, float)):
return _num_cmp(float(left), float(right), op)
except (TypeError, ValueError):
pass
# 字符串/其它:支持 eq/ne/in/not_in
if op == "eq":
return left == right
if op == "in":
return left in right
if op == "not_in":
return left not in right
if op in ("gt", "gte", "lt", "lte"):
return _num_cmp(left, right, op) # 尝试数值,字符串会 ValueError → False
return False
def _num_cmp(a: float, b: float, op: str) -> bool:
if op == "gt":
return a > b
if op == "gte":
return a >= b
if op == "lt":
return a < b
if op == "lte":
return a <= b
if op == "eq":
return a == b
return a != b
def _to_float(v) -> float | None:
if v is None:
return None
try:
f = float(v)
except (TypeError, ValueError):
return None
if f != f: # NaN
return None
return f