67fae2fef2
根因: 全市场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>
209 lines
9.4 KiB
Python
209 lines
9.4 KiB
Python
# 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}"}}
|