e330e130d4
- 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>
261 lines
11 KiB
Python
261 lines
11 KiB
Python
# 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
|