Files
sanguo_vnpy_v2/sanguo_factor/batch_eval.py
T

165 lines
7.1 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
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,
run_id: str | None = None,
) -> dict:
"""跑一轮批量评估,结果增量写入 eval_db,返回摘要."""
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", "open_interest", "vwap"])
# 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 = (
bars.select(["datetime", "vt_symbol", "close"])
.pivot(index="datetime", on="vt_symbol", values="close")
.sort("datetime").to_pandas().set_index("datetime")
)
close_wide.index = pd.to_datetime(close_wide.index)
rets = _forward_return_matrices(close_wide)
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=bars["vt_symbol"].n_unique(), 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(bars["vt_symbol"].n_unique()),
}
errors: list[str] = []
done = 0
buffer: list[dict] = []
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
if len(buffer) >= 20 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(bars["vt_symbol"].n_unique()),
}
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=缺失自动剔除,两者等价
Fw = res.pivot(index="datetime", on="vt_symbol", values="data").sort("datetime")
F = Fw.to_pandas().set_index("datetime")
F.index = pd.to_datetime(F.index)
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}"}}