Files
sanguo_vnpy_v2/sanguo_factor/batch_eval.py
T
claude_dev 67fae2fef2 perf(factor): 财务特征join_asof分块化——NAS全量1480万行grid OOM根治 [nas]
根因: 全市场grid(5555股×2670日×33列≈4G)单次join_asof + batch_eval侧
特征帧+alpha_df双全量副本,峰值10G+,7.9G NAS必爆(合成3股测不出)。

- adapter重构: iter_fundamental_feature_chunks生成器——grid按股分批构建
  (BATCH_CODES=500,批间无全量grid副本),事件右表(报告期+forecast)全批共用一份;
  build_fundamental_features改为chunks concat(单一代码路径)
- batch_eval: 特征join移到del bars之后(省1G bars常驻),逐块filter→join
  alpha_df分片→concat,不再持有特征帧全量副本;断点续跑已完成的财务因子不再触发join
- 消两处join_asof UserWarning: 显式按键sort后抑制polars 1.42 by分组无法
  校验sortedness的无信息提示(sort即正确性保险;set_sorted实测压不住)
- 等值测试: 6股合成域 batch=1/2/6 逐值等值(分块不改变结果)
- NAS真数据探针(600真股×2018-2026×batch500): roe_ttm覆盖0.904,
  峰值RSS 1017MB(含全量statement加载),零OOM;生产规模外推~2G内

[nas]

Co-Authored-By: Claude Code <noreply@anthropic.com>
2026-09-08 09:31:52 +08:00

209 lines
9.4 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.
# sanguo_factor/batch_eval.py
"""批量评估引擎:全A bars → 进程内逐因子 calculate_by_expression → 指标 → 落盘.
不走 AlphaDataset.prepare_data(其 spawn 池对每个表达式 pickle 整个 DataFrame,
9M 行 × 258 因子的传输开销不可接受);calculate_by_expression 纯进程内 polars,
内存随单因子天然有界。现有单因子分析链路(analyzer/alphalens tears)零改动.
"""
import sys
import os
import time
_VNPY_SRC = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "vnpy_v4.4.0"))
if _VNPY_SRC not in sys.path:
sys.path.insert(0, _VNPY_SRC)
import pandas as pd
import polars as pl
from datetime import datetime as _dt
from .universe import load_universe_bars, WARMUP_BARS
from .registry import get_factor
from . import eval_store
from .metrics import summarize_factor
from .fast_ops import register_fast_ops
from . import fundamental_library # noqa: F401 财务因子 import 即注册(alpha_datasets 同模式)
from .fundamental_adapter import DEFAULT_STATIC_DIR
def _forward_return_matrices(close_wide: pd.DataFrame, periods=(1, 5, 10)) -> dict[int, pd.DataFrame]:
out = {}
for p in periods:
out[p] = close_wide.shift(-p) / close_wide - 1.0
return out
def run_batch_eval(
factor_names: list[str],
start: str,
end: str,
eval_db: str,
label: str,
universe: str = "all_a",
symbols: list[str] | None = None,
limit: int | None = None,
cfg=None,
progress_cb=None,
vnpy_db_override: str | None = None,
fund_data_dir: str | None = None,
run_id: str | None = None,
) -> dict:
"""跑一轮批量评估,结果增量写入 eval_db,返回摘要.
fund_data_dir: 财务静态域根目录(None → cfg.data_paths["static_dir"] →
NAS 默认 /volume1/stock/sanguo_vnpy_v2/data/static);仅当因子列表含
category="fundamental" 时才读取并 join(量价批零开销)。
"""
from vnpy.alpha.dataset.utility import calculate_by_expression
# Register fast polars operators (idempotent)
register_fast_ops()
if cfg is None:
from sanguo_data.config import load_config, find_config_path
cfg = load_config(find_config_path())
vnpy_db = vnpy_db_override or cfg.data_paths["vnpy_db"]
t0 = time.time()
bars = load_universe_bars(vnpy_db, start, end, symbols=symbols, limit=limit)
if bars.height == 0:
raise ValueError(f"股票池为空: vnpy_db={vnpy_db} window={start}~{end}")
alpha_df = bars.select(["vt_symbol", "datetime", "open", "high", "low", "close",
"volume", "turnover", "vwap"])
# 财务因子批: bars 释放前仅抽取小体量 codes/dates(特征 join 移到 del bars 后,
# 分块进行——全量特征帧+alpha_df 双全量副本在 7.9G NAS 必 OOM)
fund_names = [n for n in factor_names
if (get_factor(n) or {}).get("category") == "fundamental"]
if fund_names:
fund_static_dir = fund_data_dir or cfg.data_paths.get("static_dir") or DEFAULT_STATIC_DIR
fund_codes = bars["vt_symbol"].unique().to_list()
fund_days = bars["datetime"].unique().sort()
# Pre-compute per-symbol warmup cutoff dates (bar_idx >= WARMUP_BARS 的首日)
# 用于替代 per-factor hash join,改为 pivot 后 pandas 广播掩码置 NaN
cutoffs = (
bars.filter(pl.col("bar_idx") >= WARMUP_BARS)
.group_by("vt_symbol").agg(pl.col("datetime").min().alias("cutoff"))
)
cutoff_map = dict(zip(cutoffs["vt_symbol"].to_list(), cutoffs["cutoff"].to_list()))
start_dt = _dt.strptime(start, "%Y-%m-%d")
end_dt = _dt.strptime(end, "%Y-%m-%d")
# close_wide 同款 pandas pivot 快路(同 _eval_one 修法:polars pivot 10M 行分钟级)
_cw = bars.select(["datetime", "vt_symbol", "close"]).to_pandas()
close_wide = _cw.pivot(index="datetime", columns="vt_symbol", values="close").sort_index()
close_wide.index = pd.to_datetime(close_wide.index)
close_wide.columns.name = None
rets = _forward_return_matrices(close_wide)
# 保存symbols_count(删除bars前)
symbols_count = bars["vt_symbol"].n_unique()
# 内存减负:bars(原始+bar_idx ~1G)派生完毕即释放——NAS 8G 盒防 OOM 周期性被杀
import gc
del bars
gc.collect()
universe_label = universe if symbols is None else "custom"
eval_store.init_db(eval_db)
if run_id is None:
run_id = eval_store.create_run(
eval_db, label=label, universe=universe_label,
symbols_count=symbols_count, factors_total=len(factor_names),
start=start, end=end, params={"limit": limit, "symbols": symbols[:20] if symbols else None}, # 截断防 4400 只全量塞 JSON
)
done_before = set()
else:
# 断点续跑:复用既有 run 行,跳过已落库因子(容器被 CI 重启后无损接续)
from sanguo_factor.eval_store import get_rows as _get_rows
done_before = {r["factor"] for r in _get_rows(eval_db, run_id)}
factor_names = [f for f in factor_names if f not in done_before]
if not factor_names:
# 空续跑:全部已完成时直接收尾不炸
eval_store.finish_run(eval_db, run_id, "done", factors_done=len(done_before))
return {
"run_id": run_id,
"factors_total": len(done_before),
"factors_done": len(done_before),
"errors": [],
"elapsed_sec": round(time.time() - t0, 1),
"symbols_count": int(symbols_count),
}
errors: list[str] = []
done = 0
buffer: list[dict] = []
# 财务特征分块 join: 逐块 filter→join alpha_df 分片→concat(bars 已释放;
# 单块特征帧用完即弃,峰值 ≈ alpha_df + 已 join 分片累积,无全量特征副本)
has_fund = any((get_factor(n) or {}).get("category") == "fundamental"
for n in factor_names)
if has_fund:
from .fundamental_adapter import iter_fundamental_feature_chunks
parts = []
for feat_chunk in iter_fundamental_feature_chunks(
fund_codes, start, end, data_dir=fund_static_dir,
trading_dates=fund_days):
syms = feat_chunk["vt_symbol"].unique().to_list()
parts.append(
alpha_df.filter(pl.col("vt_symbol").is_in(syms))
.join(feat_chunk, on=["vt_symbol", "datetime"], how="left"))
alpha_df = pl.concat(parts, how="vertical")
parts.clear()
for i, name in enumerate(factor_names):
row = _eval_one(name, alpha_df, cutoff_map, start_dt, end_dt, rets, calculate_by_expression)
if "error" in row["metrics"]:
errors.append(name)
buffer.append(row)
done += 1
# 逐因子即存:容器重启风暴下单次最多丢在途 1 因子(buffer=20 曾丢 6 个)
if len(buffer) >= 1 or done == len(factor_names):
eval_store.save_results(eval_db, run_id, buffer)
buffer = []
if progress_cb:
progress_cb(done, len(factor_names), name)
eval_store.finish_run(eval_db, run_id, "done", factors_done=len(done_before) + done)
return {
"run_id": run_id,
"factors_total": len(done_before) + len(factor_names),
"factors_done": len(done_before) + done,
"errors": errors,
"elapsed_sec": round(time.time() - t0, 1),
"symbols_count": int(symbols_count),
}
def _eval_one(name: str, alpha_df: pl.DataFrame, cutoff_map: dict[str, _dt],
start_dt: _dt, end_dt: _dt, rets: dict[int, pd.DataFrame], calculate_by_expression) -> dict:
factor = get_factor(name)
if factor is None:
return {"factor": name, "category": "unknown", "expression": "",
"metrics": {"error": f"因子未注册: {name}"}}
try:
res = calculate_by_expression(alpha_df, factor["expression"])
# 优化:先 pivot 全表,再 pandas 切窗口+预热期广播掩码置 NaN(替代 per-factor hash join
# 语义等价性:join 版把「窗口外或 bar_idx<WARMUP」的行剔除;新版这些位置变 NaN
# 下游 rank_corr_rows/metrics 所有指标以 NaN=缺失自动剔除,两者等价
# polars pivot 在 10M 行宽表化分钟级(py-spy实锤);唯一键下 pandas pivot=factorize+reshape 秒级
_pd = res.to_pandas()
F = _pd.pivot(index="datetime", columns="vt_symbol", values="data").sort_index()
F.index = pd.to_datetime(F.index)
F.columns.name = None
F = F.loc[start_dt:end_dt] # 评估窗口切片
# 预热期前置 NaN:每列用对应 cutoff_map[vt_symbol] 进行广播比较
cut_arr = pd.Series([cutoff_map.get(s) for s in F.columns], index=F.columns).to_numpy()
F = F.where(F.index.to_numpy()[:, None] >= cut_arr[None, :])
if F.empty or F.shape[0] == 0 or F.shape[1] == 0:
return {"factor": name, "category": factor["category"],
"expression": factor["expression"],
"metrics": {"error": "因子矩阵为空(处理后无评估窗行,检查 datetime dtype/窗口)"}}
metrics = summarize_factor(F, rets[1], rets[5], rets[10])
return {"factor": name, "category": factor["category"],
"expression": factor["expression"], "metrics": metrics}
except Exception as e: # 单因子失败不拖垮整批
return {"factor": name, "category": factor["category"],
"expression": factor["expression"],
"metrics": {"error": f"{type(e).__name__}: {e}"}}