224 lines
10 KiB
Python
224 lines
10 KiB
Python
# 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(d−365) = 窗 (d−365, 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")
|