perf(factor): 逐因子join砍除——pivot全表后pandas切窗口+预热cutoff广播掩码置NaN(NAS实测94.5万行join+pivot 11.4s/因子是全量8h瓶颈;NaN语义与join剔除等价,指标层自动剔缺失) [vps]
This commit is contained in:
+25
-12
@@ -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<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)
|
||||
if F.empty or len(F.index) == 0:
|
||||
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": "因子矩阵为空(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}
|
||||
|
||||
Reference in New Issue
Block a user