perf(factor): 财务特征join_asof分块化——NAS全量1480万行grid OOM根治 [nas]
根因: 全市场grid(5555股×2670日×33列≈4G)单次join_asof + batch_eval侧 特征帧+alpha_df双全量副本,峰值10G+,7.9G NAS必爆(合成3股测不出)。 - adapter重构: iter_fundamental_feature_chunks生成器——grid按股分批构建 (BATCH_CODES=500,批间无全量grid副本),事件右表(报告期+forecast)全批共用一份; build_fundamental_features改为chunks concat(单一代码路径) - batch_eval: 特征join移到del bars之后(省1G bars常驻),逐块filter→join alpha_df分片→concat,不再持有特征帧全量副本;断点续跑已完成的财务因子不再触发join - 消两处join_asof UserWarning: 显式按键sort后抑制polars 1.42 by分组无法 校验sortedness的无信息提示(sort即正确性保险;set_sorted实测压不住) - 等值测试: 6股合成域 batch=1/2/6 逐值等值(分块不改变结果) - NAS真数据探针(600真股×2018-2026×batch500): roe_ttm覆盖0.904, 峰值RSS 1017MB(含全量statement加载),零OOM;生产规模外推~2G内 [nas] Co-Authored-By: Claude Code <noreply@anthropic.com>
This commit is contained in:
@@ -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"]:
|
||||
|
||||
@@ -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])
|
||||
|
||||
@@ -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": []},
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user