diff --git a/sanguo_factor/batch_eval.py b/sanguo_factor/batch_eval.py index 5187a05..8d92d49 100644 --- a/sanguo_factor/batch_eval.py +++ b/sanguo_factor/batch_eval.py @@ -23,7 +23,7 @@ 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 build_fundamental_features, DEFAULT_STATIC_DIR +from .fundamental_adapter import DEFAULT_STATIC_DIR def _forward_return_matrices(close_wide: pd.DataFrame, periods=(1, 5, 10)) -> dict[int, pd.DataFrame]: @@ -71,17 +71,14 @@ def run_batch_eval( alpha_df = bars.select(["vt_symbol", "datetime", "open", "high", "low", "close", "volume", "turnover", "vwap"]) - # 财务因子批: 特征列 join 进 alpha_df(表达式引擎按列名直接消费;bars 释放前完成) + # 财务因子批: 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: - static_dir = fund_data_dir or cfg.data_paths.get("static_dir") or DEFAULT_STATIC_DIR - feat_df = build_fundamental_features( - codes=bars["vt_symbol"].unique().to_list(), - start=start, end=end, data_dir=static_dir, - trading_dates=bars["datetime"].unique().sort(), - ) - alpha_df = alpha_df.join(feat_df, on=["vt_symbol", "datetime"], how="left") + 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 = ( @@ -136,6 +133,24 @@ def run_batch_eval( 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"]: diff --git a/sanguo_factor/fundamental_adapter.py b/sanguo_factor/fundamental_adapter.py index d0b23ec..0c7a318 100644 --- a/sanguo_factor/fundamental_adapter.py +++ b/sanguo_factor/fundamental_adapter.py @@ -22,6 +22,7 @@ from __future__ import annotations import os +import warnings from datetime import datetime import polars as pl @@ -396,48 +397,33 @@ def _load_forecast_events(codes: list[str], data_dir: str) -> pl.DataFrame: # ==================== 对外主入口 ==================== -def _build_grid(codes: list[str], start: str, end: str, - trading_dates) -> pl.DataFrame: - """日频 grid: vt_symbol × datetime(交易日子集或日历日).""" +# 分块大小: NAS 7.9G 物理内存下,全市场 grid(5555股×2670日×33列≈4G)单次 +# join_asof + 全量副本必 OOM;按股分批使峰值 ≈ 事件表(常驻) + 单批 grid + 输出累积 +BATCH_CODES = 500 + + +def _prepare_days(start: str, end: str, trading_dates) -> list: + """交易日/日历日序列(一次解析,各批共用).""" if trading_dates is not None: days = pl.Series("datetime", trading_dates).cast(pl.Datetime("us")) - days = days.unique().sort() - else: - s = datetime.strptime(start, "%Y-%m-%d") - e = datetime.strptime(end, "%Y-%m-%d") - days = pl.datetime_range(s, e, interval="1d", eager=True).cast(pl.Datetime("us")) - day_list = days.to_list() + return days.unique().sort().to_list() + s = datetime.strptime(start, "%Y-%m-%d") + e = datetime.strptime(end, "%Y-%m-%d") + return pl.datetime_range(s, e, interval="1d", eager=True).cast( + pl.Datetime("us")).to_list() + + +def _build_grid(codes: list[str], day_list: list) -> pl.DataFrame: + """单批日频 grid: codes × day_list(批内全量,批间无全量 grid 副本).""" return pl.DataFrame({ "vt_symbol": pl.Series([c for c in codes for _ in day_list], dtype=pl.Utf8), "datetime": pl.Series(day_list * len(codes), dtype=pl.Datetime("us")), }, schema={"vt_symbol": pl.Utf8, "datetime": pl.Datetime("us")}) -def build_fundamental_features( - codes: list[str], - start: str, - end: str, - data_dir: str = DEFAULT_STATIC_DIR, - trading_dates: pl.Series | list | None = None, -) -> pl.DataFrame: - """构建 PIT 日频财务特征: vt_symbol × datetime × FEATURE_COLUMNS. - - Args: - codes: vt_symbol 列表(如 "600000.SSE") - start/end: "YYYY-MM-DD" 窗口(trading_dates=None 时生成日历日 grid) - data_dir: 静态域根目录(NAS=/volume1/stock/sanguo_vnpy_v2/data/static) - trading_dates: 交易日子集(传 bars 的 unique datetime 免造非交易日行) - - Returns: - 每行 = 决策日可见的最新报告期特征(NOTICE_DATE ≤ 决策日,asof 前向填充)。 - """ +def _load_feature_events(codes: list[str], data_dir: str) -> tuple[pl.DataFrame, pl.DataFrame]: + """报告期特征事件 + forecast 事件(全 codes 一次加载,分块 join 共用右表).""" stmt_cols = [c for c in FEATURE_COLUMNS if c not in _FORECAST_COLS] - if not codes: - return pl.DataFrame(schema={"vt_symbol": pl.Utf8, "datetime": pl.Datetime("us"), - **{c: pl.Float64 for c in FEATURE_COLUMNS}}) - - grid = _build_grid(codes, start, end, trading_dates) - reports = _load_statements(codes, data_dir) if reports.height == 0: stmt_events = pl.DataFrame(schema={ @@ -451,19 +437,75 @@ def build_fundamental_features( pl.col("notice_eff").cast(pl.Datetime("us")).alias("eff"), *stmt_cols, ).sort("eff") - fc_events = _load_forecast_events(codes, data_dir).select( pl.col("vt_symbol"), pl.col("eff").cast(pl.Datetime("us")), *_FORECAST_COLS, ).sort("eff") + return stmt_events, fc_events + + +def iter_fundamental_feature_chunks( + codes: list[str], + start: str, + end: str, + data_dir: str = DEFAULT_STATIC_DIR, + trading_dates: pl.Series | list | None = None, + batch_codes: int = BATCH_CODES, +): + """按 vt_symbol 分批产出 PIT 日频特征块(生成器,NAS 全量防 OOM 主入口). + + 每块 = 一批 codes × 全部日期 × FEATURE_COLUMNS,顺序即 codes 列表顺序; + 批内 grid 用完即弃,事件右表(报告期+forecast)全批共用仅此一份。 + batch_eval 侧应逐块 join alpha_df 分片后 concat,避免持有本帧全量副本。 + """ + stmt_events, fc_events = _load_feature_events(codes, data_dir) + day_list = _prepare_days(start, end, trading_dates) + for i in range(0, len(codes), batch_codes): + chunk_codes = codes[i:i + batch_codes] + grid = _build_grid(chunk_codes, day_list).sort("datetime") + # join_asof 前已显式按键排序;polars 1.42 用 by 分组时无法校验 sortedness, + # 该提示无信息量,就地抑制(sort 即正确性保险) + with warnings.catch_warnings(): + warnings.simplefilter("ignore", UserWarning) + out = grid.join_asof( + stmt_events, left_on="datetime", right_on="eff", + by="vt_symbol", strategy="backward") + out = out.sort("datetime").join_asof( + fc_events, left_on="datetime", right_on="eff", + by="vt_symbol", strategy="backward") + yield out.sort(["vt_symbol", "datetime"]).select( + ["vt_symbol", "datetime", *FEATURE_COLUMNS]) + grid = out = None # 批间释放(下一批重绑定) + + +def build_fundamental_features( + codes: list[str], + start: str, + end: str, + data_dir: str = DEFAULT_STATIC_DIR, + trading_dates: pl.Series | list | None = None, + batch_codes: int = BATCH_CODES, +) -> pl.DataFrame: + """构建 PIT 日频财务特征: vt_symbol × datetime × FEATURE_COLUMNS. + + Args: + codes: vt_symbol 列表(如 "600000.SSE") + start/end: "YYYY-MM-DD" 窗口(trading_dates=None 时生成日历日 grid) + data_dir: 静态域根目录(NAS=/volume1/stock/sanguo_vnpy_v2/data/static) + trading_dates: 交易日子集(传 bars 的 unique datetime 免造非交易日行) + batch_codes: 按股分批大小(全量防 OOM;测试可调小验分块等值) + + Returns: + 每行 = 决策日可见的最新报告期特征(NOTICE_DATE ≤ 决策日,asof 前向填充)。 + """ + schema = {"vt_symbol": pl.Utf8, "datetime": pl.Datetime("us"), + **{c: pl.Float64 for c in FEATURE_COLUMNS}} + if not codes: + return pl.DataFrame(schema=schema) + return pl.concat( + iter_fundamental_feature_chunks( + codes, start, end, data_dir=data_dir, + trading_dates=trading_dates, batch_codes=batch_codes), + how="vertical") - # asof 前向填充: 决策日可见的最新披露 - out = grid.sort("datetime").join_asof( - stmt_events, left_on="datetime", right_on="eff", - by="vt_symbol", strategy="backward") - out = out.sort("datetime").join_asof( - fc_events, left_on="datetime", right_on="eff", - by="vt_symbol", strategy="backward") - return out.sort(["vt_symbol", "datetime"]).select( - ["vt_symbol", "datetime", *FEATURE_COLUMNS]) diff --git a/tests/factor/conftest.py b/tests/factor/conftest.py index e466590..805f4c0 100644 --- a/tests/factor/conftest.py +++ b/tests/factor/conftest.py @@ -39,6 +39,13 @@ SYN_STOCKS = { "null_notice_periods": ["2022-09-30"]}, # NOTICE_DATE 缺失 → 该报告期跳过 "300001.SZ": {"vt": "300001.SZSE", "scale": 2.0, "drop_periods": [], "null_notice_periods": []}, + # 三只分块等值测试扩容股(全史无残缺,不同 scale 增截面多样性) + "600004.SH": {"vt": "600004.SSE", "scale": 0.5, + "drop_periods": [], "null_notice_periods": []}, + "000333.SZ": {"vt": "000333.SZSE", "scale": 1.7, + "drop_periods": [], "null_notice_periods": []}, + "300124.SZ": {"vt": "300124.SZSE", "scale": 3.0, + "drop_periods": [], "null_notice_periods": []}, } diff --git a/tests/factor/test_fundamental_adapter.py b/tests/factor/test_fundamental_adapter.py index 6f23f0f..c0ac3c3 100644 --- a/tests/factor/test_fundamental_adapter.py +++ b/tests/factor/test_fundamental_adapter.py @@ -195,3 +195,24 @@ def test_scale_two_stock(feat): # C(scale=2): TA=2*1450, NP_TTM=2*60 → roa 同 A(比率不变), np_ttm 翻倍 assert val(feat, C, "2023-08-29", "roa_ttm") == pytest.approx(60 / 1450) assert val(feat, C, "2023-08-29", "np_ttm") == pytest.approx(120.0) + + +# ---------- 分块等值(NAS 全量防 OOM 路径) ---------- + +_SIX = ["600000.SSE", "000001.SZSE", "300001.SZSE", + "600004.SSE", "000333.SZSE", "300124.SZSE"] + + +def test_chunked_equals_full(synthetic_static): + """batch_codes=2(3 批) 与 batch_codes=6(单批) 逐值等值——分块不改变结果.""" + full = build_fundamental_features( + _SIX, "2023-01-01", "2023-12-31", data_dir=synthetic_static, batch_codes=6) + chunked = build_fundamental_features( + _SIX, "2023-01-01", "2023-12-31", data_dir=synthetic_static, batch_codes=2) + assert full.height == chunked.height == 6 * 365 + key = ["vt_symbol", "datetime"] + assert full.sort(key).equals(chunked.sort(key)) + # 批大小 1(极端) 与 trading_dates 路径同款等值 + by_one = build_fundamental_features( + _SIX, "2023-01-01", "2023-12-31", data_dir=synthetic_static, batch_codes=1) + assert by_one.sort(key).equals(full.sort(key))