330 lines
15 KiB
Python
330 lines
15 KiB
Python
# sanguo_factor/fundamental_domains.py
|
||
"""P1-B 四新域事件层: dividend/gdhs/top_holders/valuation_baostock
|
||
(自 fundamental_adapter.py 拆出,纯结构重构).
|
||
|
||
四域缺任一 → 相关特征列全 NaN + 一次性 warning(优雅降级,本地可跑通);
|
||
_load_extra_domains 按引用列瘦身加载(未引用的域不读盘).
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import os
|
||
import warnings
|
||
from datetime import date, datetime
|
||
|
||
import polars as pl
|
||
|
||
from .fundamental_schema import _SYM
|
||
from .fundamental_statements import _norm_dates, _vt_to_file_code
|
||
|
||
# ==================== P1-B 四新域事件层 ====================
|
||
|
||
# 缺域告警去重(每 (域, 路径) 一次;优雅降级 = 特征列全 NaN 不崩,
|
||
# 本地开发无 NAS 数据也能跑通管线)
|
||
_missing_warned: set[tuple[str, str]] = set()
|
||
|
||
# dividend PIT 日期回退链(首列空往后回退;全空行丢弃)
|
||
_DIVIDEND_CHAIN = ["dividPlanAnnounceDate", "dividPreNoticeDate",
|
||
"dividAgmPumDate", "dividPlanDate"]
|
||
|
||
# gdhs 截止日列名(NAS 实测「股东户数统计截止日-本次」;兼容任务书连字符变体)
|
||
_GDHS_CUTOFF_CANDIDATES = ["股东户数统计截止日-本次", "股东户数-统计截止日-本次"]
|
||
|
||
# vb 2026 缺口补口窗口(bs 日喂起点 2026-08-13 → 01-01~08-12 由 static/valuation 补)
|
||
_VB_GAP = (date(2026, 1, 1), date(2026, 8, 12))
|
||
|
||
|
||
def _warn_domain_missing(key: str, path: str) -> None:
|
||
if (key, path) not in _missing_warned:
|
||
_missing_warned.add((key, path))
|
||
warnings.warn(
|
||
f"[fundamental_adapter] 数据域 {key} 缺失({path}) → 相关特征列全 NaN"
|
||
f"(本地开发无 NAS 数据可忽略;NAS 真跑前先确认路径)")
|
||
|
||
|
||
def _code6_to_vt(code: pl.Expr) -> pl.Expr:
|
||
"""6 位代码 → vt_symbol(60→SSE;北交 92/43/82/83→BJSE;其余→SZSE;
|
||
与 forecast 映射同款,BJ 不落入 SZSE)."""
|
||
is_bj = (code.str.starts_with("92") | code.str.starts_with("43")
|
||
| code.str.starts_with("82") | code.str.starts_with("83"))
|
||
return (pl.when(code.str.starts_with("60")).then(code + pl.lit(".SSE"))
|
||
.when(is_bj).then(code + pl.lit(".BJSE"))
|
||
.otherwise(code + pl.lit(".SZSE")))
|
||
|
||
|
||
def _load_dividend_cum(codes: list[str], data_dir: str) -> pl.DataFrame:
|
||
"""dividend 按股文件 → (vt_symbol, eff, _send_cum) 累计事件流.
|
||
|
||
事件 = 一次分红;eff = 回退链首个非空公告日(全空行丢弃);值 =
|
||
dividStocksPs(每股送转合计)。按股日序 cum_sum 供 E11 的 365 天滚动窗
|
||
两端口径相减(cum(d) − cum(d−365) = 窗 (d−365, d] 内事件合计)。
|
||
code 列为 baostock 小写格式(sh.600519)与文件名不同 → 以文件名路由
|
||
(与三表读取同模式,列内容不参与 join)。
|
||
"""
|
||
schema = {"vt_symbol": pl.Utf8, "eff": pl.Datetime("us"), "_send_cum": pl.Float64}
|
||
div_dir = os.path.join(data_dir, "dividend")
|
||
if not os.path.isdir(div_dir):
|
||
_warn_domain_missing("dividend", div_dir)
|
||
return pl.DataFrame(schema=schema)
|
||
frames = []
|
||
for vt in codes:
|
||
file_code = _vt_to_file_code(vt)
|
||
if file_code is None:
|
||
continue
|
||
path = os.path.join(div_dir, f"{file_code}_dividend.parquet")
|
||
if not os.path.exists(path):
|
||
continue
|
||
try:
|
||
df = pl.read_parquet(path)
|
||
except Exception:
|
||
continue
|
||
chain = [c for c in _DIVIDEND_CHAIN if c in df.columns]
|
||
if not chain or "dividStocksPs" not in df.columns:
|
||
continue
|
||
# 日期列字符串('' 为空;兼容 'YYYY-MM-DD 00:00:00')→ coalesce 回退链
|
||
eff = pl.coalesce([
|
||
pl.col(c).cast(pl.Utf8).str.slice(0, 10).str.to_date("%Y-%m-%d", strict=False)
|
||
for c in chain])
|
||
df = df.select(
|
||
pl.lit(vt).alias("vt_symbol"),
|
||
eff.alias("eff"),
|
||
pl.col("dividStocksPs").cast(pl.Float64, strict=False).alias("_send"),
|
||
).filter(pl.col("eff").is_not_null())
|
||
if df.height:
|
||
frames.append(df)
|
||
if not frames:
|
||
return pl.DataFrame(schema=schema)
|
||
ev = pl.concat(frames).sort(["vt_symbol", "eff"])
|
||
return ev.with_columns(
|
||
pl.col("_send").cum_sum().over(_SYM).alias("_send_cum")
|
||
).select(["vt_symbol", "eff", "_send_cum"]).with_columns(
|
||
pl.col("eff").cast(pl.Datetime("us")))
|
||
|
||
|
||
def _load_gdhs_events(codes: list[str], data_dir: str) -> pl.DataFrame:
|
||
"""gdhs 按期文件 → (vt_symbol, eff, gdhs_chg) 事件流.
|
||
|
||
事件表: (代码, 统计截止日)→股东户数;同键多行取最新公告(修订口径);
|
||
gdhs_chg = 相邻事件(按截止日排序)Δln(户数),首事件无上期 → NaN;
|
||
PIT = 公告日期。红线: 不用「股东户数-增减比例」列(每股截止日不规则,
|
||
非统一季环比)。
|
||
"""
|
||
schema = {"vt_symbol": pl.Utf8, "eff": pl.Datetime("us"), "gdhs_chg": pl.Float64}
|
||
gd_dir = os.path.join(data_dir, "gdhs")
|
||
if not os.path.isdir(gd_dir):
|
||
_warn_domain_missing("gdhs", gd_dir)
|
||
return pl.DataFrame(schema=schema)
|
||
code_set = set(codes)
|
||
frames = []
|
||
for fname in sorted(os.listdir(gd_dir)):
|
||
if not fname.endswith(".parquet"):
|
||
continue
|
||
try:
|
||
f = pl.read_parquet(os.path.join(gd_dir, fname))
|
||
except Exception:
|
||
continue
|
||
cutoff = next((c for c in _GDHS_CUTOFF_CANDIDATES if c in f.columns), None)
|
||
if cutoff is None or not all(
|
||
c in f.columns for c in ("代码", "股东户数-本次", "公告日期")):
|
||
continue
|
||
code = pl.col("代码").cast(pl.Utf8).str.strip_chars().str.zfill(6)
|
||
f = f.select(
|
||
_code6_to_vt(code).alias("vt_symbol"),
|
||
pl.col("股东户数-本次").cast(pl.Float64, strict=False).alias("_holders"),
|
||
pl.col(cutoff).cast(pl.Date, strict=False).alias("_cutoff"),
|
||
pl.col("公告日期").cast(pl.Date, strict=False).alias("_ann"),
|
||
).filter(pl.col("vt_symbol").is_in(code_set) & pl.col("_ann").is_not_null()
|
||
& pl.col("_cutoff").is_not_null() & pl.col("_holders").is_not_null())
|
||
if f.height:
|
||
frames.append(f)
|
||
if not frames:
|
||
return pl.DataFrame(schema=schema)
|
||
ev = (pl.concat(frames)
|
||
.sort(["vt_symbol", "_cutoff", "_ann"])
|
||
.group_by(["vt_symbol", "_cutoff"]).agg( # 同 (股, 截止日) 取最新公告
|
||
pl.col("_holders").last(), pl.col("_ann").last())
|
||
.sort(["vt_symbol", "_cutoff"]))
|
||
return ev.with_columns(
|
||
pl.col("_holders").shift(1).over(_SYM).alias("_prev")
|
||
).with_columns(
|
||
pl.when((pl.col("_prev") > 0) & (pl.col("_holders") > 0))
|
||
.then((pl.col("_holders") / pl.col("_prev")).log())
|
||
.otherwise(None).alias("gdhs_chg")
|
||
).select(
|
||
pl.col("vt_symbol"),
|
||
pl.col("_ann").cast(pl.Datetime("us")).alias("eff"),
|
||
pl.col("gdhs_chg"),
|
||
).sort("eff")
|
||
|
||
|
||
def _load_topholder_events(codes: list[str], data_dir: str) -> pl.DataFrame:
|
||
"""top_holders 聚合产物 → (vt_symbol, eff, topholder_chg) 事件流.
|
||
|
||
聚合产物 = scripts/factor_research/preaggregate_top_holders.py 输出
|
||
(file_code, period, hold_pct, n_holders),置于 static 兄弟目录
|
||
factor_cache/ 下(绝不写 static 树)。期→PIT 无法定披露日 → 法定披露
|
||
截止近似(保守侧): Q1→04-30 / H1→08-31 / Q3→10-31 / 年报→次年04-30。
|
||
topholder_chg = 最近期十大合计占比 − 上期(首期无上期 → NaN)。
|
||
"""
|
||
schema = {"vt_symbol": pl.Utf8, "eff": pl.Datetime("us"), "topholder_chg": pl.Float64}
|
||
agg_path = os.path.join(os.path.dirname(os.path.abspath(data_dir)),
|
||
"factor_cache", "top_holders_agg.parquet")
|
||
if not os.path.exists(agg_path):
|
||
_warn_domain_missing("top_holders_agg", agg_path)
|
||
return pl.DataFrame(schema=schema)
|
||
try:
|
||
agg = pl.read_parquet(agg_path)
|
||
except Exception:
|
||
_warn_domain_missing("top_holders_agg", agg_path)
|
||
return pl.DataFrame(schema=schema)
|
||
if not all(c in agg.columns for c in ("file_code", "period", "hold_pct")):
|
||
return pl.DataFrame(schema=schema)
|
||
parts = pl.col("file_code").str.split(".")
|
||
ev = agg.with_columns(
|
||
parts.list.get(0).alias("_c"), parts.list.get(1).alias("_x"),
|
||
pl.col("period").cast(pl.Date, strict=False),
|
||
pl.col("hold_pct").cast(pl.Float64, strict=False),
|
||
).with_columns(
|
||
pl.when(pl.col("_x") == "SH").then(pl.col("_c") + pl.lit(".SSE"))
|
||
.otherwise(pl.col("_c") + pl.lit(".SZSE")).alias("vt_symbol")
|
||
).filter(
|
||
pl.col("vt_symbol").is_in(set(codes)) & pl.col("period").is_not_null()
|
||
& pl.col("hold_pct").is_not_null()
|
||
).sort(["vt_symbol", "period"])
|
||
m = pl.col("period").dt.month()
|
||
y = pl.col("period").dt.year()
|
||
deadline = (
|
||
pl.when(m == 3).then(pl.date(y, 4, 30))
|
||
.when(m == 6).then(pl.date(y, 8, 31))
|
||
.when(m == 9).then(pl.date(y, 10, 31))
|
||
.otherwise(pl.date(y + 1, 4, 30))
|
||
)
|
||
return ev.with_columns(
|
||
pl.col("hold_pct").shift(1).over(_SYM).alias("_prev")
|
||
).with_columns(
|
||
(pl.col("hold_pct") - pl.col("_prev")).alias("topholder_chg")
|
||
).select(
|
||
pl.col("vt_symbol"),
|
||
deadline.cast(pl.Datetime("us")).alias("eff"),
|
||
pl.col("topholder_chg"),
|
||
).sort("eff")
|
||
|
||
|
||
def _inv_guard(col: str) -> pl.Expr:
|
||
"""1/x 守卫: x 缺失/为 0 → NaN(亏损 baostock pe=NaN 天然 NaN)."""
|
||
return (pl.when(pl.col(col).is_not_null() & (pl.col(col) != 0))
|
||
.then(1.0 / pl.col(col)).otherwise(None))
|
||
|
||
|
||
def _load_vb_daily(codes: list[str], data_dir: str,
|
||
start: str, end: str) -> pl.DataFrame:
|
||
"""valuation_baostock 按年文件(+static/valuation 补 2026 缺口) →
|
||
(vt_symbol, eff, ep_vb, bp_vb) 日频流.
|
||
|
||
vb 在 static 的兄弟目录(data/valuation_baostock/{year}.parquet),
|
||
date 为字符串须转 Date;ep=1/peTTM、bp=1/pbMRQ。2026-01-01~08-12
|
||
缺口(bs 日喂 08-13 起)用 static/valuation(ak em 中文列:
|
||
PE(TTM)/市净率)倒数补——两源亏损口径不同(baostock NaN/ak 可为负),
|
||
倒数后均按原始符号保留,分域验证时注意。
|
||
"""
|
||
schema = {"vt_symbol": pl.Utf8, "eff": pl.Datetime("us"),
|
||
"ep_vb": pl.Float64, "bp_vb": pl.Float64}
|
||
root = os.path.dirname(os.path.abspath(data_dir))
|
||
vb_dir = os.path.join(root, "valuation_baostock")
|
||
if not os.path.isdir(vb_dir):
|
||
_warn_domain_missing("valuation_baostock", vb_dir)
|
||
return pl.DataFrame(schema=schema)
|
||
code_set = set(codes)
|
||
d_lo = datetime.strptime(start, "%Y-%m-%d").date()
|
||
d_hi = datetime.strptime(end, "%Y-%m-%d").date()
|
||
frames = []
|
||
for year in range(d_lo.year, d_hi.year + 1):
|
||
path = os.path.join(vb_dir, f"{year}.parquet")
|
||
if not os.path.exists(path):
|
||
continue
|
||
try:
|
||
f = pl.read_parquet(path)
|
||
except Exception:
|
||
continue
|
||
if not all(c in f.columns for c in
|
||
("symbol", "exchange", "date", "peTTM", "pbMRQ")):
|
||
continue
|
||
raw_date = pl.col("date")
|
||
d = (raw_date.cast(pl.Utf8).str.slice(0, 10).str.to_date("%Y-%m-%d", strict=False)
|
||
if f.schema["date"] == pl.Utf8 else raw_date.cast(pl.Date, strict=False))
|
||
# exchange 实测为 SH/SZ(NAS 2026-09-08,任务书 SSE/SZSE 变体兼容)
|
||
xsuf = (pl.when(pl.col("exchange") == pl.lit("SH")).then(pl.lit("SSE"))
|
||
.when(pl.col("exchange") == pl.lit("SZ")).then(pl.lit("SZSE"))
|
||
.otherwise(pl.col("exchange")))
|
||
f = f.select(
|
||
(pl.col("symbol").cast(pl.Utf8).str.strip_chars() + pl.lit(".")
|
||
+ xsuf.cast(pl.Utf8).str.strip_chars()).alias("vt_symbol"),
|
||
d.alias("eff"),
|
||
pl.col("peTTM").cast(pl.Float64, strict=False).alias("_pe"),
|
||
pl.col("pbMRQ").cast(pl.Float64, strict=False).alias("_pb"),
|
||
).filter(pl.col("vt_symbol").is_in(code_set) & pl.col("eff").is_not_null()
|
||
& (pl.col("eff") >= d_lo) & (pl.col("eff") <= d_hi))
|
||
if f.height:
|
||
frames.append(f)
|
||
# 2026 缺口补口: static/valuation 按股中文列
|
||
gap_lo, gap_hi = max(_VB_GAP[0], d_lo), min(_VB_GAP[1], d_hi)
|
||
if gap_lo <= gap_hi:
|
||
val_dir = os.path.join(data_dir, "valuation")
|
||
if not os.path.isdir(val_dir):
|
||
_warn_domain_missing("static/valuation(2026 补口)", val_dir)
|
||
else:
|
||
for vt in codes:
|
||
file_code = _vt_to_file_code(vt)
|
||
if file_code is None:
|
||
continue
|
||
path = os.path.join(val_dir, f"{file_code}_valuation.parquet")
|
||
if not os.path.exists(path):
|
||
continue
|
||
try:
|
||
f = pl.read_parquet(path)
|
||
except Exception:
|
||
continue
|
||
have = [c for c in ("数据日期", "PE(TTM)", "市净率") if c in f.columns]
|
||
if "数据日期" not in have:
|
||
continue
|
||
f = _norm_dates(f.select(have), ["数据日期"])
|
||
pe = (pl.col("PE(TTM)").cast(pl.Float64, strict=False)
|
||
if "PE(TTM)" in have else pl.lit(None, pl.Float64))
|
||
pb = (pl.col("市净率").cast(pl.Float64, strict=False)
|
||
if "市净率" in have else pl.lit(None, pl.Float64))
|
||
f = f.select(
|
||
pl.lit(vt).alias("vt_symbol"),
|
||
pl.col("数据日期").alias("eff"),
|
||
pe.alias("_pe"), pb.alias("_pb"),
|
||
).filter(pl.col("eff").is_not_null()
|
||
& (pl.col("eff") >= gap_lo) & (pl.col("eff") <= gap_hi))
|
||
if f.height:
|
||
frames.append(f)
|
||
if not frames:
|
||
return pl.DataFrame(schema=schema)
|
||
out = (pl.concat(frames).sort(["vt_symbol", "eff"])
|
||
.unique(subset=["vt_symbol", "eff"], keep="last")
|
||
.with_columns(
|
||
_inv_guard("_pe").alias("ep_vb"),
|
||
_inv_guard("_pb").alias("bp_vb"),
|
||
pl.col("eff").cast(pl.Datetime("us"))))
|
||
return out.select(["vt_symbol", "eff", "ep_vb", "bp_vb"]).sort("eff")
|
||
|
||
|
||
_EMPTY_EVENTS = pl.DataFrame(schema={"vt_symbol": pl.Utf8, "eff": pl.Datetime("us")})
|
||
|
||
|
||
def _load_extra_domains(codes: list[str], data_dir: str,
|
||
start: str, end: str, out_cols: list[str]) -> dict:
|
||
"""P1-B 四新域按引用列加载(引用瘦身的域级延伸: 未引用的域不读盘)."""
|
||
cols = set(out_cols)
|
||
return {
|
||
"dividend": (_load_dividend_cum(codes, data_dir)
|
||
if "send_total_12m" in cols else _EMPTY_EVENTS),
|
||
"gdhs": (_load_gdhs_events(codes, data_dir)
|
||
if "gdhs_chg" in cols else _EMPTY_EVENTS),
|
||
"topholder": (_load_topholder_events(codes, data_dir)
|
||
if "topholder_chg" in cols else _EMPTY_EVENTS),
|
||
"vb": (_load_vb_daily(codes, data_dir, start, end)
|
||
if "ep_vb" in cols or "bp_vb" in cols else _EMPTY_EVENTS),
|
||
}
|