Files
qlib/backend/app/quant/selection.py
T
Simon 2e90f3eeac feat(backend): 字段库(condition_field)+ 因子参数化(模板/受控参数)+ 单位换算底座
字段库(本次新增的表与接口):
- `condition_field` 表 + `/api/condition-fields`:中文名/说明可编辑、可停用;
  `kind`/单位阶梯/`base_unit` 由代码注册表收敛(改类型 422,伪字段 422,
  越界单位 422),停用的字段不再进条件下拉,但既有策略仍按名字解析。
- 说明书里的数值条件按字段注册表补**基准单位**后缀(字段间比较不加,不猜单位)。

因子参数化(键即身份,冻结口径):
- 模板 + 参数注册表(`quant/factors.py`):`ParamSpec`(类型/范围/枚举/默认值/说明)+
  `FactorTemplate`(公式/依赖列/参数);规范键把**全部**参数写进名字,如
  `momentum(window=90,direction=lower_is_better)`,所以改参数 = 新建一个身份,
  旧因子/既有策略/已归档实验都不变义;`momentum(window=90)`(缺参数)明确拒绝 ——
  缺项要靠模板默认值补齐,而默认值是可改的代码细节,一旦改动会追溯性改义。
- 参数只在受控范围内取值(窗口 2~500、方向二选一),越界/未知模板/多给参数一律 422
  并列出允许范围,不静默截断、不悄悄取默认值;内置实例的启用开关由代码决定(422)。
- `/api/factors` 暴露 `template`/`params`/`param_specs`/`label`/`source`/`enabled`/
  `resolvable`;新增 `/api/factors/templates`、`POST /api/factors`、`PATCH /api/factors`;
  `get_factor = resolve_factor` 兼容全部旧调用点,参数化键也是一等条件字段。
- 迁移链:c5d6(存量策略陈旧说明重算)→ d6e7(condition_field)→ a7c1
  (factor_definition.enabled + name varchar(128))。

测试:新增 test_condition_fields.py / test_factor_params.py;全量 pytest 500 passed。
2026-10-01 16:33:32 +08:00

450 lines
16 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 DAILY_BASIC_NUMERIC_FIELDS, FinancialIndicator
from app.domain.entities.selection import (
SelectionCandidate,
SelectionQuery,
SelectionResult,
SelectionStatistics,
)
from app.quant.composite import build_score_panel
from app.quant.factors import FactorError, compute_factor, get_factor
_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「历史回测与当前选股使用同一套引擎」的一致性。
"""
return build_score_panel(daily, factor_specs) # 未知因子在此抛 FactorError
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) -> set[str]:
"""条件引用的行情/指标列(fundamental/static 走元数据与财务表,不需要面板列)。
query 可以是 SelectionQuery,也可以是任何带 `conditions` 的对象
(ResearchSpec 亦然)—— 回测与选股共用本函数,保证列裁剪一致。
"""
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
if f in DAILY_BASIC_NUMERIC_FIELDS:
needed.add(f) # 每日指标列(dv_ratio / pe / pb / total_mv …)
continue
try: # 其余按已注册因子处理
defn, _fn = get_factor(f)
except FactorError:
raise ValueError(
f"条件字段未知:{f}(可用: 行情列/ma20/ma60/已注册因子/"
"参数化因子键(如 momentum(window=90,direction=higher_is_better))/"
"每日指标列(dv_ratio 等)/static.*/fundamental.*)"
) from None
needed.update(defn.requires)
return needed
def build_condition_fields(
daily: pd.DataFrame,
conditions,
obs: pd.Timestamp,
) -> dict[str, pd.Series]:
"""条件各字段在 obs(<= as_of 的最近交易日)的截面值。
返回 {字段名: Series(index=symbol)},覆盖:
- 行情原列:close / open / high / low / volume / amount
- 技术派生:ma20 / ma60
- 每日指标列:dv_ratio / dv_ttm / pe / pb / total_mv …(由 Service 并入 daily)
- 已注册因子:momentum_60 等(在 <=obs 的截断数据上计算,无未来函数)
回测与选股共用本函数(v2 §25 一致性)。
"""
view = daily[pd.to_datetime(daily["trade_date"]) <= obs]
if view.empty:
return {}
close = view.pivot(index="trade_date", columns="symbol", values="close").sort_index()
close.index = pd.to_datetime(close.index)
fields: dict[str, pd.Series] = {}
for col in ("open", "high", "low", "volume", "amount"):
if col in view.columns:
panel = view.pivot(index="trade_date", columns="symbol", values=col).sort_index()
panel.index = pd.to_datetime(panel.index)
if obs in panel.index:
fields[col] = panel.loc[obs]
if obs in close.index:
fields["close"] = close.loc[obs]
ma20 = close.rolling(20).mean()
ma60 = close.rolling(60).mean()
if obs in ma20.index:
fields["ma20"] = ma20.loc[obs]
if obs in ma60.index:
fields["ma60"] = ma60.loc[obs]
wanted: set[str] = set()
for cond in conditions:
for f in (cond.field, cond.ref):
if not f or f.startswith((_STATIC_PREFIX, _FUNDAMENTAL_PREFIX)):
continue
if f in _TECH_DERIVED or f in fields:
continue
wanted.add(f)
for f in sorted(wanted):
if f in DAILY_BASIC_NUMERIC_FIELDS:
if f in view.columns:
panel = view.pivot(index="trade_date", columns="symbol", values=f).sort_index()
panel.index = pd.to_datetime(panel.index)
if obs in panel.index:
fields[f] = panel.loc[obs]
continue
try:
_defn, panel = compute_factor(f, view)
except FactorError:
continue # 已在 condition_needed_columns 报错;此处防御
if obs in panel.index:
fields[f] = panel.loc[obs]
return fields
def eligible_symbols(
candidates,
conditions,
statics: dict[str, dict],
fields: dict[str, pd.Series],
financial: dict[str, FinancialIndicator],
) -> dict[str, list[str]]:
"""逐股求值全部条件(AND),返回 {symbol: [各条件通过情况文案]}(仅通过者)。
`candidates` 限定参与求值的股票(通常 = universe 过滤后的 symbol 列表)。
回测(ResearchSpec.conditions)与选股(SelectionQuery.conditions)共用,
确保「历史某日 Selection == 回测当日 Selection」(v2 §25 / v3 §28)。
"""
passed: dict[str, list[str]] = {}
for sym in candidates:
statuses: list[str] = []
all_ok = True
for cond in conditions:
ok = _eval_condition(cond, sym, statics, fields, financial)
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[sym] = statuses
return passed
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"),
)
# 共享字段面板 + 求值器(回测与选股同一实现,v2 §25 一致性)
tech = build_condition_fields(daily, query.conditions, obs)
statics = {s.symbol: s.model_dump() for s in stocks}
passed = eligible_symbols(sorted(statics), query.conditions, statics, tech, financial or {})
candidates: list[SelectionCandidate] = []
for rank, sym in enumerate(sorted(passed), start=1):
candidates.append(
SelectionCandidate(
symbol=sym,
rank=rank,
score=1.0,
filter_status=passed[sym],
selection_reason=[f"通过全部 {len(query.conditions)} 条条件"],
)
)
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 外不通过;数值/字符串分别处理。
任何类型不匹配(拿日期字段去比大小、in 的右侧不是列表…)一律返回 False,
**绝不抛异常**:条件是用户可以随手改的输入,一个手滑的字段名不该把选股/回测
打成 500。字段库(quant.condition_fields)会在源头上拒绝不可比较的字段,
这里只是最后一道防线。
"""
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 ("in", "not_in"):
try:
return left in right if op == "in" else left not in right
except TypeError: # 右侧不是容器 → 该条件无法求值
return False
if op in ("gt", "gte", "lt", "lte"):
try:
return _num_cmp(left, right, op) # 尝试数值,字符串会 ValueError → False
except (TypeError, ValueError):
return 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