From 0b3f6bb87c56a046b4b53ed21da1a6d1c3311a72 Mon Sep 17 00:00:00 2001 From: claude_dev Date: Tue, 25 Aug 2026 01:43:52 +0800 Subject: [PATCH] =?UTF-8?q?perf(factor):=20=E9=80=90=E5=9B=A0=E5=AD=90join?= =?UTF-8?q?=E7=A0=8D=E9=99=A4=E2=80=94=E2=80=94pivot=E5=85=A8=E8=A1=A8?= =?UTF-8?q?=E5=90=8Epandas=E5=88=87=E7=AA=97=E5=8F=A3+=E9=A2=84=E7=83=ADcu?= =?UTF-8?q?toff=E5=B9=BF=E6=92=AD=E6=8E=A9=E7=A0=81=E7=BD=AENaN(NAS?= =?UTF-8?q?=E5=AE=9E=E6=B5=8B94.5=E4=B8=87=E8=A1=8Cjoin+pivot=2011.4s/?= =?UTF-8?q?=E5=9B=A0=E5=AD=90=E6=98=AF=E5=85=A8=E9=87=8F8h=E7=93=B6?= =?UTF-8?q?=E9=A2=88;NaN=E8=AF=AD=E4=B9=89=E4=B8=8Ejoin=E5=89=94=E9=99=A4?= =?UTF-8?q?=E7=AD=89=E4=BB=B7,=E6=8C=87=E6=A0=87=E5=B1=82=E8=87=AA?= =?UTF-8?q?=E5=8A=A8=E5=89=94=E7=BC=BA=E5=A4=B1)=20[vps]?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- sanguo_factor/batch_eval.py | 37 +++++++++++++++++++++++++------------ 1 file changed, 25 insertions(+), 12 deletions(-) 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}