Files
sanguo_vnpy_v2/sanguo_factor/fundamental_pit.py
T
2026-09-10 07:18:28 +08:00

224 lines
10 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# sanguo_factor/fundamental_pit.py
"""PIT 对齐 + 分块 join 主入口: 报告期/forecast/兑现差/四新域事件流 →
日频特征块(自 fundamental_adapter.py 拆出,纯结构重构).
R3 PIT 红线: NOTICE_DATE ≤ 决策日才可见(asof backward 前向填充);
按股分批 join_asof 防 NAS 7.9G 内存 OOM.
"""
from __future__ import annotations
import warnings
from datetime import datetime
import polars as pl
from .fundamental_domains import _load_extra_domains
from .fundamental_forecast import _load_forecast_events
from .fundamental_report_features import _compute_report_features, _safe_ratio
from .fundamental_schema import (
_BEAT_COLS,
_DOMAIN_COLS,
_FORECAST_COLS,
DEFAULT_STATIC_DIR,
FEATURE_COLUMNS,
)
from .fundamental_statements import _load_statements
# ==================== 对外主入口 ====================
# 分块大小: 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"))
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 _load_feature_events(codes: list[str], data_dir: str) -> tuple[pl.DataFrame, pl.DataFrame, pl.DataFrame]:
"""报告期特征事件 + forecast 事件 + 兑现差事件(全 codes 一次加载,
分块 join 共用右表)."""
stmt_cols = [c for c in FEATURE_COLUMNS
if c not in _FORECAST_COLS and c not in _BEAT_COLS
and c not in _DOMAIN_COLS]
reports = _load_statements(codes, data_dir)
if reports.height == 0:
stmt_events = pl.DataFrame(schema={
"vt_symbol": pl.Utf8, "eff": pl.Datetime("us"), **{c: pl.Float64 for c in stmt_cols}})
beat_events = pl.DataFrame(schema={
"vt_symbol": pl.Utf8, "eff": pl.Datetime("us"), "forecast_beat": pl.Float64})
else:
feat = _compute_report_features(reports)
# NOTICE_DATE 缺失报告期整期跳过(红线: 宁缺毋假)
feat = feat.filter(pl.col("notice_eff").is_not_null())
stmt_events = feat.select(
pl.col("vt_symbol"),
pl.col("notice_eff").cast(pl.Datetime("us")).alias("eff"),
*stmt_cols,
).sort("eff")
beat_events = _build_beat_events(feat, codes, data_dir)
fc_events, _ = _load_forecast_events(codes, data_dir)
fc_events = fc_events.select(
pl.col("vt_symbol"),
pl.col("eff").cast(pl.Datetime("us")),
*_FORECAST_COLS,
).sort("eff")
return stmt_events, fc_events, beat_events
def _build_beat_events(feat: pl.DataFrame, codes: list[str], data_dir: str) -> pl.DataFrame:
"""F07 预告兑现差事件流: (实际NP − 预告中值)/abs(预告中值),同 REPORT_DATE 配对.
前视红线(P1 任务书): 兑现差含实际 NP,只有实际报告披露后才可知——
PIT 锚 = max(该报告期 income 有效披露日 notice_eff, 预告公告日),
不早于两者较晚者(预告公告晚于年报的罕见情形不被提前泄露)。
实际 NP 用报告期累计归母净利(预告口径即期间累计);预告中值=0/缺 → NaN。
"""
schema = {"vt_symbol": pl.Utf8, "eff": pl.Datetime("us"), "forecast_beat": pl.Float64}
_, fc_pair = _load_forecast_events(codes, data_dir)
if fc_pair.height == 0:
return pl.DataFrame(schema=schema)
beat = (
feat.select("vt_symbol", "REPORT_DATE", "notice_eff", pl.col("np"))
.join(fc_pair, on=["vt_symbol", "REPORT_DATE"], how="inner")
.with_columns(
_safe_ratio(pl.col("np") - pl.col("_fc_mid"),
pl.col("_fc_mid").abs()).alias("forecast_beat"),
pl.max_horizontal("notice_eff", "eff").alias("_anchor"),
)
.filter(pl.col("forecast_beat").is_not_null() & pl.col("_anchor").is_not_null())
)
return beat.select(
pl.col("vt_symbol"),
pl.col("_anchor").cast(pl.Datetime("us")).alias("eff"),
pl.col("forecast_beat"),
).sort("eff")
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,
columns: list[str] | None = None,
):
"""按 vt_symbol 分批产出 PIT 日频特征块(生成器,NAS 全量防 OOM 主入口).
每块 = 一批 codes × 全部日期 × columns(默认 FEATURE_COLUMNS 全量),
顺序即 codes 列表顺序;批内 grid 用完即弃,事件右表(报告期+forecast+
兑现差)全批共用仅此一份。batch_eval 侧应逐块 join alpha_df 分片后
concat,避免持有本帧全量副本;columns 子集可只 join 本批表达式引用列,
全量 68 列 × 1480 万行 ≈ 8G——按引用瘦身是 NAS 7.9G 内存的关键杠杆。
"""
out_cols = list(columns) if columns is not None else list(FEATURE_COLUMNS)
stmt_events, fc_events, beat_events = _load_feature_events(codes, data_dir)
extra = _load_extra_domains(codes, data_dir, start, end, out_cols)
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)
# 三条事件流各自 asof 后丢弃右表键 eff(留置会以 eff_right 后缀
# 累积,第三次 join 撞名)
out = grid.join_asof(
stmt_events, left_on="datetime", right_on="eff",
by="vt_symbol", strategy="backward").drop("eff")
out = out.sort("datetime").join_asof(
fc_events, left_on="datetime", right_on="eff",
by="vt_symbol", strategy="backward").drop("eff")
out = out.sort("datetime").join_asof(
beat_events, left_on="datetime", right_on="eff",
by="vt_symbol", strategy="backward").drop("eff")
# ---- P1-B 四新域 ----
div = extra["dividend"]
if div.height:
# E11: cum(d) cum(d365) = 窗 (d365, d] 事件合计(开区间下界);
# 域内无事件(≤d) → NaN 不填 0(与缺文件/域不覆盖区分)
div365 = div.rename({"_send_cum": "_send_cum_365"})
out = out.sort("datetime").join_asof(
div, left_on="datetime", right_on="eff",
by="vt_symbol", strategy="backward").drop("eff")
out = (
out.with_columns(
(pl.col("datetime") - pl.duration(days=365)).alias("_dt365"))
.sort("_dt365")
.join_asof(div365, left_on="_dt365", right_on="eff",
by="vt_symbol", strategy="backward")
.drop("eff", "_dt365")
.with_columns(
pl.when(pl.col("_send_cum").is_null()).then(None)
.otherwise(pl.col("_send_cum")
- pl.col("_send_cum_365").fill_null(0.0))
.alias("send_total_12m")))
for key in ("gdhs", "topholder"):
if extra[key].height:
out = out.sort("datetime").join_asof(
extra[key], left_on="datetime", right_on="eff",
by="vt_symbol", strategy="backward").drop("eff")
if extra["vb"].height:
out = out.sort("datetime").join_asof(
extra["vb"], left_on="datetime", right_on="eff",
by="vt_symbol", strategy="backward").drop("eff")
# 缺域/无事件兜底: 输出列缺失 → null 列(优雅降级,select 不炸)
absent = [c for c in out_cols if c not in out.columns]
if absent:
out = out.with_columns(
[pl.lit(None, dtype=pl.Float64).alias(c) for c in absent])
yield out.sort(["vt_symbol", "datetime"]).select(
["vt_symbol", "datetime", *out_cols])
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,
columns: list[str] | None = None,
) -> pl.DataFrame:
"""构建 PIT 日频财务特征: vt_symbol × datetime × 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;测试可调小验分块等值)
columns: 输出特征列子集(默认 FEATURE_COLUMNS 全量;引用瘦身用)
Returns:
每行 = 决策日可见的最新报告期特征(NOTICE_DATE ≤ 决策日,asof 前向填充)。
"""
out_cols = list(columns) if columns is not None else list(FEATURE_COLUMNS)
schema = {"vt_symbol": pl.Utf8, "datetime": pl.Datetime("us"),
**{c: pl.Float64 for c in out_cols}}
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, columns=out_cols),
how="vertical")