feat(factor): batch_eval接入财务因子——特征列join+端到端IC(财务批P0 WP3) [nas]
- run_batch_eval 增 fund_data_dir 参数;仅因子列表含 category=fundamental 时 build_fundamental_features(bars 的 codes+trading_dates) join 进 alpha_df (列白名单经 join 扩展,量价批零开销零改动路径) - __init__ 补挂 fundamental_library(import sanguo_factor 全家可见) - 端到端3测: 财务4因子IC落库/量价+财务混批/纯量价无静态域回归 - 修顺序依赖: 其它测试清 _REGISTRY 只重挂alpha → 本批两模块 autouse 幂等重注册 [nas] Co-Authored-By: Claude Code <noreply@anthropic.com>
This commit is contained in:
@@ -1,3 +1,4 @@
|
||||
"""Sanguo factor module for vnpy alpha strategies."""
|
||||
from . import library # noqa: F401 (triggers _register_all to register built-in factors)
|
||||
from . import alpha_datasets # noqa: F401 挂载 Alpha101/158(导入即注册)
|
||||
from . import fundamental_library # noqa: F401 财务因子 P0 批(导入即注册)
|
||||
|
||||
@@ -22,6 +22,8 @@ from .registry import get_factor
|
||||
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
|
||||
|
||||
|
||||
def _forward_return_matrices(close_wide: pd.DataFrame, periods=(1, 5, 10)) -> dict[int, pd.DataFrame]:
|
||||
@@ -43,9 +45,15 @@ def run_batch_eval(
|
||||
cfg=None,
|
||||
progress_cb=None,
|
||||
vnpy_db_override: str | None = None,
|
||||
fund_data_dir: str | None = None,
|
||||
run_id: str | None = None,
|
||||
) -> dict:
|
||||
"""跑一轮批量评估,结果增量写入 eval_db,返回摘要."""
|
||||
"""跑一轮批量评估,结果增量写入 eval_db,返回摘要.
|
||||
|
||||
fund_data_dir: 财务静态域根目录(None → cfg.data_paths["static_dir"] →
|
||||
NAS 默认 /volume1/stock/sanguo_vnpy_v2/data/static);仅当因子列表含
|
||||
category="fundamental" 时才读取并 join(量价批零开销)。
|
||||
"""
|
||||
from vnpy.alpha.dataset.utility import calculate_by_expression
|
||||
|
||||
# Register fast polars operators (idempotent)
|
||||
@@ -63,6 +71,17 @@ def run_batch_eval(
|
||||
|
||||
alpha_df = bars.select(["vt_symbol", "datetime", "open", "high", "low", "close",
|
||||
"volume", "turnover", "vwap"])
|
||||
# 财务因子批: 特征列 join 进 alpha_df(表达式引擎按列名直接消费;bars 释放前完成)
|
||||
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")
|
||||
# Pre-compute per-symbol warmup cutoff dates (bar_idx >= WARMUP_BARS 的首日)
|
||||
# 用于替代 per-factor hash join,改为 pivot 后 pandas 广播掩码置 NaN
|
||||
cutoffs = (
|
||||
|
||||
@@ -0,0 +1,101 @@
|
||||
# tests/factor/test_fundamental_batch.py
|
||||
"""财务因子批端到端: 合成行情库 + 合成静态域 → batch_eval 出 IC 指标."""
|
||||
import sqlite3
|
||||
import sys, os
|
||||
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")))
|
||||
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", "vnpy_v4.4.0")))
|
||||
|
||||
import pytest
|
||||
|
||||
from sanguo_factor import alpha_datasets # 挂载量价因子
|
||||
from sanguo_factor import fundamental_library # noqa: F401 挂载财务因子(import 即注册)
|
||||
from sanguo_factor.batch_eval import run_batch_eval
|
||||
from sanguo_factor import eval_store
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _ensure_fundamental_registered():
|
||||
"""其它测试模块清空 _REGISTRY 后只重挂 alpha/builtin(顺序依赖前科),
|
||||
这里逐测试幂等重注册财务因子,保证本模块与顺序无关."""
|
||||
fundamental_library._register_all()
|
||||
|
||||
|
||||
_DDL = """
|
||||
CREATE TABLE dbbardata(
|
||||
symbol TEXT, exchange TEXT, datetime TEXT, interval TEXT,
|
||||
volume REAL, turnover REAL, open_interest REAL,
|
||||
open_price REAL, high_price REAL, low_price REAL, close_price REAL)
|
||||
"""
|
||||
|
||||
_STOCKS = [("600000", "SSE"), ("000001", "SZSE"), ("300001", "SZSE")]
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def db(tmp_path_factory):
|
||||
"""3 只 × 2022-06~2024-01 合成日线(与合成财务域同代码;趋势构造截面差异)."""
|
||||
alpha_datasets.mount_all()
|
||||
p = tmp_path_factory.mktemp("fb")
|
||||
db = str(p / "qt.db")
|
||||
conn = sqlite3.connect(db)
|
||||
conn.execute(_DDL)
|
||||
import pandas as pd
|
||||
days = pd.bdate_range("2022-06-01", "2024-01-05")
|
||||
drift = {"600000": 0.0012, "000001": -0.0006, "300001": 0.0004}
|
||||
base = {"600000": 10.0, "000001": 5.0, "300001": 20.0}
|
||||
for i, day in enumerate(days):
|
||||
d = day.strftime("%Y-%m-%d")
|
||||
rows = []
|
||||
for sym, ex in _STOCKS:
|
||||
px = base[sym] * (1.0 + drift[sym]) ** i
|
||||
rows.append((sym, ex, f"{d} 00:00:00", "d", 100.0, px * 100.0, 0,
|
||||
px, px, px, px))
|
||||
conn.executemany("INSERT INTO dbbardata VALUES(?,?,?,?,?,?,?,?,?,?,?)", rows)
|
||||
conn.commit()
|
||||
conn.close()
|
||||
return db
|
||||
|
||||
|
||||
def test_fundamental_end_to_end(db, synthetic_static, tmp_path):
|
||||
eval_db = str(tmp_path / "fund_eval.db")
|
||||
out = run_batch_eval(
|
||||
factor_names=["fund_roe_ttm", "fund_tacc", "fund_ep_ttm",
|
||||
"fund_forecast_type"],
|
||||
start="2023-02-01", end="2023-12-31",
|
||||
eval_db=eval_db, label="fund_t", cfg=None, vnpy_db_override=db,
|
||||
fund_data_dir=synthetic_static,
|
||||
)
|
||||
assert out["factors_done"] == 4
|
||||
assert out["errors"] == []
|
||||
for name in ("fund_roe_ttm", "fund_tacc", "fund_ep_ttm", "fund_forecast_type"):
|
||||
m = eval_store.get_detail(eval_db, out["run_id"], name)["metrics"]
|
||||
assert "error" not in m, f"{name}: {m.get('error')}"
|
||||
# 截面 3 只 → 逐日 rank IC 可算(预告因子 2023-07-15 后才有效, count 少但非零)
|
||||
assert isinstance(m["1"]["ic_mean"], float) or m["1"]["ic_mean"] is None
|
||||
assert m["1"]["count"] >= 0
|
||||
|
||||
|
||||
def test_mixed_batch_price_plus_fundamental(db, synthetic_static, tmp_path):
|
||||
"""量价+财务同批混跑: 白名单 join 只在含财务因子时发生,互不干扰."""
|
||||
eval_db = str(tmp_path / "mix_eval.db")
|
||||
out = run_batch_eval(
|
||||
factor_names=["ma_20", "fund_sue_np"],
|
||||
start="2023-02-01", end="2023-06-30",
|
||||
eval_db=eval_db, label="mix_t", cfg=None, vnpy_db_override=db,
|
||||
fund_data_dir=synthetic_static,
|
||||
)
|
||||
assert out["errors"] == []
|
||||
rows = eval_store.get_rows(eval_db, out["run_id"])
|
||||
assert {r["factor"] for r in rows} == {"ma_20", "fund_sue_np"}
|
||||
# SUE 在 2023-06-30 前可见的报告期均无完整 8 期窗 → 全 NaN → eliminated 非 error
|
||||
m = eval_store.get_detail(eval_db, out["run_id"], "fund_sue_np")["metrics"]
|
||||
assert "error" not in m
|
||||
|
||||
|
||||
def test_pure_price_batch_untouched(db, tmp_path):
|
||||
"""纯量价批不传静态域照常跑(回归: 未新增强依赖)."""
|
||||
eval_db = str(tmp_path / "px_eval.db")
|
||||
out = run_batch_eval(
|
||||
factor_names=["roc_5"], start="2023-02-01", end="2023-03-31",
|
||||
eval_db=eval_db, label="px_t", cfg=None, vnpy_db_override=db,
|
||||
)
|
||||
assert out["errors"] == []
|
||||
assert out["factors_done"] == 1
|
||||
Reference in New Issue
Block a user