Files
sanguo_vnpy_v2/tests/factor/test_composite_library.py
T
claude_dev e330e130d4 chore(factor): P1二轮回评修缮——v1.1口径残留对齐+导入示例纠偏+replace_strict [nas]
- composite_library docstring/测试注释: 7源/19源/fund7/all19 残留 → 6源/18源
  (7dfcaca 剔 topholder 后文档未同步,防后续按旧口径引用)
- import_runs_to_main_db 示例 --dst data_backup→data(09-08 导错库坑,
  示例即陷阱原文,加⚠️一行防再踩)
- forecast 预告类型映射 replace→replace_strict(消 398 条 DeprecationWarning,
  参数语义等价,231 测全绿)
- P1 晨报§五: 补 09-09 只读探针复核(首期 20210930/5431文件每期)+§19.10
  修订文本入档(master 侧一行改,本分支无该节)

Co-Authored-By: Claude Code <noreply@anthropic.com>
2026-09-09 12:12:00 +08:00

261 lines
11 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.
# tests/factor/test_composite_library.py
"""合成层 v1(Combiner): 3 合成因子注册 + 引擎解析 + 方向传播 + 数值域 + 端到端."""
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 numpy as np
import pandas as pd
import polars as pl
import pytest
from sanguo_factor import composite_library
from sanguo_factor.composite_library import (
QUANT_SOURCES, FUND_SOURCES, build_composite_expression, rank_term,
)
from sanguo_factor.registry import list_factors, get_factor
from vnpy.alpha.dataset.utility import calculate_by_expression
@pytest.fixture(autouse=True)
def _ensure_composite_registered():
"""其它测试模块清空 _REGISTRY 后,这里幂等重挂源+合成,保证本模块与顺序无关."""
composite_library._register_all()
# ==================== 注册与源表契约 ====================
def test_three_composites_registered():
facs = list_factors("composite")
assert {f["name"] for f in facs} == {"composite_quant12", "composite_fund6",
"composite_all18"}
def test_source_tables_cover_18_independent_sources():
assert len(QUANT_SOURCES) == 12
assert len(FUND_SOURCES) == 6
names = [s[0] for s in QUANT_SOURCES] + [s[0] for s in FUND_SOURCES]
assert len(set(names)) == 18 # 无重复源
for name, _dir in QUANT_SOURCES:
f = get_factor(name)
assert f is not None and f["category"] in ("alpha101", "alpha158", "builtin"), name
for name, _dir in FUND_SOURCES:
f = get_factor(name)
assert f is not None and f["category"] == "fundamental", name
# 方向只允许 +/-
assert {d for _, d in QUANT_SOURCES + FUND_SOURCES} <= {"+", "-"}
def test_quant_sources_direction_table():
"""负向源 6 个(vol_ma5/wvma_20/klow/cord_5/kup/alpha81),其余正向."""
dirs = dict(QUANT_SOURCES)
neg = {n for n, d in dirs.items() if d == "-"}
assert neg == {"vol_ma5", "wvma_20", "klow", "cord_5", "kup", "alpha81"}
def test_registration_idempotent():
import importlib
importlib.reload(composite_library)
assert len(list_factors("composite")) == 3
# ==================== 表达式构造契约 ====================
def test_quant12_wraps_every_source_in_cs_rank_with_direction():
"""负向源形态钉死: (-1) * cs_rank((expr))(负号在 rank 外,口径见 library 注释)."""
expr = get_factor("composite_quant12")["expression"]
assert expr.startswith("(") and expr.endswith("/ 12")
for name, direction in QUANT_SOURCES:
assert rank_term(get_factor(name)["expression"], direction) in expr, \
f"{name} 项缺失或形态不符"
assert expr.count("(-1) * cs_rank((") == 6 # 6 个负向源
def test_fund6_embeds_registered_expressions_without_double_rank():
"""财务源注册表达式已 cs_rank 定向 → 直接内嵌,不得再包一层 cs_rank."""
expr = get_factor("composite_fund6")["expression"]
assert expr.endswith("/ 6")
for name, _direction in FUND_SOURCES:
assert f"({get_factor(name)['expression']})" in expr, f"{name} 未原样内嵌"
assert "cs_rank((cs_rank(" not in expr # 无双重 rank 包裹
def test_all18_mixes_both_kinds():
expr = get_factor("composite_all18")["expression"]
assert expr.endswith("/ 18")
# 12 量价项包 cs_rank + 6 财务项内嵌
for name, direction in QUANT_SOURCES:
assert rank_term(get_factor(name)["expression"], direction) in expr, name
for name, _direction in FUND_SOURCES:
assert f"({get_factor(name)['expression']})" in expr, name
# ==================== cs_rank 截面语义锁定 ====================
def test_cs_rank_is_per_day_cross_section():
"""引擎 cs_rank 语义 = 按 datetime 分组截面 rank(非全表 rank)——合成口径的前提."""
df = pl.DataFrame({
"datetime": ["2023-01-02"] * 3 + ["2023-01-03"] * 3,
"vt_symbol": ["A", "B", "C"] * 2,
"close": [10.0, 5.0, 20.0, 30.0, 5.0, 10.0],
}).with_columns(pl.col("datetime").str.to_datetime())
out = calculate_by_expression(df, "cs_rank(close)")
vals = out["data"].to_list()
# day1: 10<20 → rank(10)=2, rank(5)=1, rank(20)=3;day2: 30=最大 → 3
assert vals == [2.0, 1.0, 3.0, 3.0, 1.0, 2.0]
# ==================== 方向传播与数值域(真实引擎) ====================
_STOCKS5 = ["600000.SSE", "000001.SZSE", "300001.SZSE", "600004.SSE", "000333.SZSE"]
_FUND_COLS = ["equity", "share_capital", "nsi", "gp_ttm", "gdhs_chg",
"growth_scissors", "sue_np"]
def _engine_df(n_days: int = 130) -> pl.DataFrame:
"""5 股 × n_days 合成长表: 覆盖 12 量价源全部引用列 + 7 财务特征列."""
rng = np.random.default_rng(42)
days = pd.bdate_range("2023-01-02", periods=n_days)
rows = []
px = {s: 8.0 * (i + 1) for i, s in enumerate(_STOCKS5)}
bias = {s: rng.normal(0, 0.5) for s in _STOCKS5}
for i, day in enumerate(days):
for j, s in enumerate(_STOCKS5):
ret = rng.normal(0, 0.02)
px[s] = max(px[s] * (1 + ret), 0.5)
open_ = px[s] * (1 + rng.normal(0, 0.005))
high = max(open_, px[s]) * (1 + abs(rng.normal(0, 0.004)))
low = min(open_, px[s]) * (1 - abs(rng.normal(0, 0.004)))
volume = float(np.exp(rng.normal(10, 0.4)))
vwap = (high + low + px[s]) / 3.0
row = {"datetime": day, "vt_symbol": s, "open": open_, "high": high,
"low": low, "close": px[s], "volume": volume, "vwap": vwap}
for c in _FUND_COLS:
row[c] = bias[s] + rng.normal(0, 1.0)
rows.append(row)
return pl.DataFrame(rows).with_columns(pl.col("datetime").cast(pl.Datetime("us")))
def _eval_sorted(df: pl.DataFrame, expression: str) -> pd.DataFrame:
out = calculate_by_expression(df, expression).to_pandas()
return out.sort_values(["datetime", "vt_symbol"]).reset_index(drop=True)
def test_negative_sign_cancels_same_source():
"""同源一正一负 → 合成恒 ≈ 0(负号经 rank 外前置 (-1) * 正确传递)."""
df = _engine_df(40)
expr = build_composite_expression([rank_term("close", "+"), rank_term("close", "-")])
out = _eval_sorted(df, expr)
assert np.allclose(out["data"].dropna(), 0.0, atol=1e-9)
assert out["data"].notna().all()
def test_same_direction_equals_single_source():
"""同源同向 → 合成 = 单源 rank(等权均值退化)."""
df = _engine_df(40)
expr = build_composite_expression([rank_term("close", "+"), rank_term("close", "+")])
got = _eval_sorted(df, expr)["data"]
want = _eval_sorted(df, "cs_rank(close)")["data"]
np.testing.assert_allclose(got, want)
@pytest.mark.parametrize("composite,quant_dirs,fund_dirs", [
("composite_quant12", dict(QUANT_SOURCES), {}),
("composite_fund6", {}, dict(FUND_SOURCES)),
("composite_all18", dict(QUANT_SOURCES), dict(FUND_SOURCES)),
])
def test_composite_equals_mean_of_directed_terms(composite, quant_dirs, fund_dirs):
"""数值域: 合成值 == 逐源定向求均值(引擎逐项评估 vs 整体表达式全等)."""
df = _engine_df(130)
got = _eval_sorted(df, get_factor(composite)["expression"])["data"]
parts = []
for name, direction in quant_dirs.items():
# 量价源先截面化(与 build 同口径)再定向
s = _eval_sorted(df, f"cs_rank(({get_factor(name)['expression']}))")["data"]
parts.append(s if direction == "+" else -s)
for name, direction in fund_dirs.items():
# 财务源注册表达式已定向,原样评估
s = _eval_sorted(df, get_factor(name)["expression"])["data"]
parts.append(s if direction == "+" else -s)
want = sum(parts) / len(parts)
assert (got.isna() == want.isna()).all(), "null 位置应一致(逐源 null 传染)"
valid = ~got.isna()
assert valid.sum() > 0, "长窗合成应有非空值(预热期外)"
np.testing.assert_allclose(got[valid], want[valid], atol=1e-9)
# 数值域: 合成 ∈ [min(项均值), max(项均值)] 逐日截面收紧到 [−N, N] 秩域
n_stocks = len(_STOCKS5)
assert got[valid].min() >= -n_stocks - 1e-9
assert got[valid].max() <= n_stocks + 1e-9
# ==================== 端到端(batch_eval 复合门) ====================
_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)
"""
_EOD_STOCKS = [("600000", "SSE"), ("000001", "SZSE"), ("300001", "SZSE"),
("600004", "SSE"), ("000333", "SZSE")]
@pytest.fixture(scope="module")
def db(tmp_path_factory):
"""5 只 × 2022-06~2024-01 合成日线.快振荡价(相位错开 → 截面排序频繁交叉),
OHLCV 逐日扰动,turnover=wap×volume(vwap≠close).退化形态会杀源因子:
O=C → alpha2 全并列 NaN;慢振荡/趋势 → 5 bar 窗内 cs_rank(high) 恒定
→ alpha16 的 ts_cov 恒 null(3 只慢振荡实测 0/243 天可算)."""
composite_library._register_all()
rng = np.random.default_rng(7)
p = tmp_path_factory.mktemp("cb")
db = str(p / "qt.db")
conn = sqlite3.connect(db)
conn.execute(_DDL)
days = pd.bdate_range("2022-06-01", "2024-01-05")
phase = {s: 1.3 * j for j, (s, _ex) in enumerate(_EOD_STOCKS)}
for i, day in enumerate(days):
d = day.strftime("%Y-%m-%d")
for sym, ex in _EOD_STOCKS:
px = 10.0 + 3.0 * np.sin(i / 2.2 + phase[sym]) + rng.normal(0, 0.25)
open_ = px * (1 + rng.normal(0, 0.005))
high = max(open_, px) * (1 + abs(rng.normal(0, 0.004)))
low = min(open_, px) * (1 - abs(rng.normal(0, 0.004)))
volume = float(np.exp(rng.normal(4.6, 0.4)))
wap = (high + low + px) / 3.0
conn.execute("INSERT INTO dbbardata VALUES(?,?,?,?,?,?,?,?,?,?,?)",
(sym, ex, f"{d} 00:00:00", "d", volume, wap * volume, 0,
open_, high, low, px))
conn.commit()
conn.close()
return db
def test_composite_end_to_end(db, synthetic_static, tmp_path):
"""category=composite 同样触发财务特征 join(复合门),三因子无错误评估."""
from sanguo_factor.batch_eval import run_batch_eval
from sanguo_factor import eval_store
eval_db = str(tmp_path / "comp_eval.db")
out = run_batch_eval(
factor_names=["composite_quant12", "composite_fund6", "composite_all18"],
start="2023-02-01", end="2023-12-31",
eval_db=eval_db, label="comp_t", cfg=None, vnpy_db_override=db,
fund_data_dir=synthetic_static,
)
assert out["factors_done"] == 3
assert out["errors"] == []
assert {r["factor"] for r in eval_store.get_rows(eval_db, out["run_id"],
category="composite")} == \
{"composite_quant12", "composite_fund6", "composite_all18"}
for name in ("composite_quant12", "composite_fund6", "composite_all18"):
m = eval_store.get_detail(eval_db, out["run_id"], name)["metrics"]
assert "error" not in m, f"{name}: {m.get('error')}"
# 量价源全天候覆盖 → quant12 截面 IC 样本点非零;
# fund6/all18 在合成静态域 gdhs_chg 等事件流源覆盖受限 → count 可为 0(非错误)
m = eval_store.get_detail(eval_db, out["run_id"], "composite_quant12")["metrics"]
assert m["1"]["count"] > 0