diff --git a/sanguo_factor/batch_eval.py b/sanguo_factor/batch_eval.py index 8a7a110..e004bb9 100644 --- a/sanguo_factor/batch_eval.py +++ b/sanguo_factor/batch_eval.py @@ -15,8 +15,9 @@ if _VNPY_SRC not in sys.path: import pandas as pd import polars as pl +from datetime import datetime as _dt -from .universe import load_universe_bars, evaluation_filter +from .universe import load_universe_bars, WARMUP_BARS from .registry import get_factor from . import eval_store from .metrics import summarize_factor @@ -57,7 +58,15 @@ def run_batch_eval( alpha_df = bars.select(["vt_symbol", "datetime", "open", "high", "low", "close", "volume", "turnover", "open_interest", "vwap"]) - eval_rows = evaluation_filter(bars, start, end).select(["vt_symbol", "datetime"]) + # 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"]) @@ -79,7 +88,7 @@ def run_batch_eval( done = 0 buffer: list[dict] = [] for i, name in enumerate(factor_names): - row = _eval_one(name, alpha_df, eval_rows, rets, calculate_by_expression) + 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) @@ -101,24 +110,28 @@ def run_batch_eval( } -def _eval_one(name: str, alpha_df: pl.DataFrame, eval_rows: pl.DataFrame, - rets: dict[int, pd.DataFrame], calculate_by_expression) -> dict: +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"]) - f_long = res.join(eval_rows, on=["vt_symbol", "datetime"], how="inner") - F = ( - f_long.pivot(index="datetime", on="vt_symbol", values="data") - .sort("datetime").to_pandas().set_index("datetime") - ) + # 优化:先 pivot 全表,再 pandas 切窗口+预热期广播掩码置 NaN(替代 per-factor hash join) + # 语义等价性:join 版把「窗口外或 bar_idx= 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": "因子矩阵为空(join 后无评估窗行,检查 datetime dtype/窗口)"}} + "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}