字段库(本次新增的表与接口): - `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。
450 lines
16 KiB
Python
450 lines
16 KiB
Python
"""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
|