Files
sanguo_vnpy_v2/tests/factor/test_challenger_lgbm.py
T

362 lines
17 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.
"""LGBM challenger 三路对拍(spec §4.4 2026-10-10):数据层."""
import json
from pathlib import Path
import numpy as np
import pandas as pd
import pytest
from sanguo_factor.challenger_lgbm import (
build_label,
cs_rank_norm,
equal_weight_signal,
icir_signal,
load_pool,
lgbm_walk_forward,
run_challenge,
score_signal,
)
def test_load_pool_twelve_sources():
pool = load_pool()
assert len(pool) == 12
assert pool["vma_60"]["direction"] == "+"
assert pool["vol_ma5"]["direction"] == "-"
assert 0.0 < pool["alpha16"]["weight"] < 0.2
def test_build_label_gap_proof():
idx = pd.date_range("2025-01-01", periods=4, freq="D")
close = pd.DataFrame({"a": [10.0, 11.0, 12.0, 13.0], "b": [20.0, 20.0, 21.0, 22.0]}, index=idx)
lab = build_label(close)
# label[t]=close[t+2]/close[t+1]-1:t0 行=12/11-1
assert lab.iloc[0]["a"] == pytest.approx(12.0 / 11.0 - 1)
# 末两行无可成交区间=NaN
assert lab.iloc[-1].isna().all() and lab.iloc[-2].isna().all()
def test_cs_rank_norm_nan_passthrough_and_centered():
df = pd.DataFrame({"a": [1.0, 2.0, 3.0, np.nan],
"b": [4.0, 3.0, 2.0, 1.0]})
out = cs_rank_norm(df)
# 行2截面 [a=3.0, b=2.0]: a 为最大, rank pct 1.0 - 0.5 = 0.5
# (plan 原断言 0.75 与其自身注释"rank pct 1.0 - 0.5"矛盾, 按权威实现修正)
assert out.loc[2, "a"] == pytest.approx(0.5)
assert np.isnan(out.loc[3, "a"]) # NaN 透传不占截面
row0 = out.loc[0, ["a", "b"]].tolist()
assert row0 == [pytest.approx(0.0), pytest.approx(0.5)] # 截面最小/最大
def _toy_values(idx, cols):
rng = np.random.default_rng(7)
return {n: pd.DataFrame(rng.normal(size=(len(idx), len(cols))), index=idx, columns=cols)
for n in ("f_pos", "f_neg")}
def test_equal_weight_direction_adjusted():
idx = pd.date_range("2025-06-01", periods=3, freq="D")
cols = ["a", "b"]
values = {"f_pos": pd.DataFrame(1.0, index=idx, columns=cols),
"f_neg": pd.DataFrame(1.0, index=idx, columns=cols)}
pool = {"f_pos": {"direction": "+"}, "f_neg": {"direction": "-"}}
sig = equal_weight_signal(values, pool)
# +1 与 -1 等权平均=0
assert (sig == 0.0).all().all()
def test_icir_signal_uses_archive_weights():
idx = pd.date_range("2025-06-01", periods=2, freq="D")
cols = ["a"]
values = {"f_pos": pd.DataFrame(2.0, index=idx, columns=cols),
"f_neg": pd.DataFrame(2.0, index=idx, columns=cols)}
pool = {"f_pos": {"direction": "+", "weight": 0.75},
"f_neg": {"direction": "-", "weight": 0.25}}
sig = icir_signal(values, pool)
# 常值 2.0 单列截面的 CSRankNorm=0.5,期望=0.5*(+0.75-0.25)
# (plan 原期望 2.0*(0.75-0.25) 未算入秩归一,按权威实现修正;
# DataFrame 与 pytest.approx 标量不兼容,用 np.allclose)
assert np.allclose(sig.values, 0.5 * (0.75 - 0.25))
def test_lgbm_walk_forward_holdout_isolated():
"""样外月绝不进训练(test 隔离铁律):给训练月与样外月截然不同的
因子-收益关系,样外预测应反映训练期学到的关系而非记忆样外."""
import pandas as pd
from sanguo_factor.challenger_lgbm import build_label, lgbm_walk_forward
idx = pd.date_range("2025-01-01", "2025-03-31", freq="B")
cols = [f"s{i}" for i in range(5)]
rng = np.random.default_rng(42)
values = {"s0": pd.DataFrame(rng.normal(size=(len(idx), 5)), index=idx, columns=cols)}
label = build_label(pd.DataFrame(100 + rng.normal(scale=0.5, size=(len(idx), 5)),
index=idx, columns=cols))
oos = idx[idx >= "2025-03-01"]
# toy 数据(~215 样本)喂生产超参(lambda_l1=205+min_child_samples=100)会零分裂
# →gain=0;走函数自带的 params 覆盖入口传轻量超参验证 split/gain 通路
light_params = {"objective": "mse", "learning_rate": 0.1, "num_leaves": 4,
"min_child_samples": 5, "lambda_l1": 0.0, "lambda_l2": 1.0,
"seed": 42, "num_threads": 4, "verbose": -1}
pred, imp = lgbm_walk_forward(values, label, last_month="2025-03", params=light_params)
assert list(pred.index) == [d for d in oos if d in label.index and label.loc[d].notna().any()]
assert set(imp.keys()) == {"s0"} and imp["s0"] > 0
def test_score_signal_perfect_and_flat():
# 5 列:实现有最小截面阈值(len(pair)<5 skip),plan 原稿 3 列进不了评分
idx = pd.date_range("2025-03-03", periods=10, freq="B")
cols = ["a", "b", "c", "d", "e"]
sig = pd.DataFrame(np.linspace(-1, 1, 50).reshape(10, 5), index=idx, columns=cols)
lab = sig * 1.0 # 完美信号
s = score_signal(sig, lab)
assert s["ic_mean"] == pytest.approx(1.0, abs=1e-6)
assert s["q5q1"] > 0
flat = pd.DataFrame(0.0, index=idx, columns=cols)
s0 = score_signal(flat, lab)
assert s0["days"] == 10 # 进评分的天数(IC=NaN 但天在)
def test_run_challenge_end_to_end(tmp_path, monkeypatch):
"""端到端:合成 2 因子×40 日数据落宽表 parquet+合成 close db 太重——
本用例走 values_dir 真文件+vnpy_db 用 monkeypatch 替换拉取函数."""
from sanguo_factor import challenger_lgbm as cl
idx = pd.date_range("2025-01-01", "2025-03-31", freq="B")
cols = ["a", "b"]
rng = np.random.default_rng(3)
values = {n: pd.DataFrame(rng.normal(size=(len(idx), 2)), index=idx, columns=cols)
for n in ("s0", "s1")}
vdir = tmp_path / "factor_values"
vdir.mkdir()
for n, v in values.items():
v.to_parquet(vdir / f"{n}.parquet")
close = pd.DataFrame(100 + np.cumsum(rng.normal(scale=0.4, size=(len(idx), 2)), axis=0),
index=idx, columns=cols)
# s0/s1 不在 quant12 档案内,plan 原稿缺此 patch 必 FileNotFoundError
monkeypatch.setattr(cl, "load_pool", lambda: {
"s0": {"direction": "+", "weight": 0.6},
"s1": {"direction": "-", "weight": 0.4}})
# 真实 loader(universe.load_universe_bars)必传 start/end,签名 4 参
monkeypatch.setattr(cl, "_load_close_wide", lambda db, c, s, e: close)
out = cl.run_challenge(str(vdir), "fake.db", "2025-03-31", str(tmp_path))
doc = json.loads(Path(out).read_text())
assert doc["as_of"] == "2025-03-31"
assert set(doc["signals"]) == {"equal", "icir", "lgbm"}
for k, v in doc["signals"].items():
assert {"ic_mean", "icir", "q5q1", "days"} <= set(v)
assert doc["pool"]["names"] == ["s0", "s1"]
assert len(doc["feature_importance"]) == 2
assert doc["split"]["oos_month"] == "2025-03"
def test_cli_smoke(tmp_path, monkeypatch, capsys):
from sanguo_factor import challenger_lgbm as cl
idx = pd.date_range("2025-01-01", "2025-03-31", freq="B")
close = pd.DataFrame(100 + np.cumsum(np.random.default_rng(1).normal(0.4, size=(len(idx), 2)), axis=0),
index=idx, columns=["a", "b"])
monkeypatch.setattr(cl, "_load_close_wide", lambda db, c, s, e: close)
monkeypatch.setattr(cl, "load_pool", lambda: {
"s0": {"direction": "+", "weight": 0.6},
"s1": {"direction": "-", "weight": 0.4}})
vdir = tmp_path / "vals"; vdir.mkdir()
for n in ("s0", "s1"):
pd.DataFrame(np.random.default_rng(2).normal(size=(len(idx), 2)),
index=idx, columns=["a", "b"]).to_parquet(vdir / f"{n}.parquet")
rc = cl.main(["--values-dir", str(vdir), "--as-of", "2025-03-31",
"--out-dir", str(tmp_path), "--vnpy-db", "fake.db"])
assert rc == 0
assert (tmp_path / "challenger_lgbm" / "nas_2025-03-31.json").exists()
# ———— codex 复审班次二修批(2026-10-10):H1/H7 embargo+训练窗裁拆 ————
def test_split_masks_embargo_and_floor():
"""H-1/H→M-7:训练窗=[样外月前推 11 个自然月首,样外首日前 2 交易日).
embargo 2 交易的数学保证:训练样本 t 的 label 取 close[t+1],close[t+2]
(build_label shift(-2)/(--1)),故 max(train_pos)+2 < first_oos_pos
⇔ 样外月收盘价不参与任何训练样本的 label 计算.
"""
from sanguo_factor.challenger_lgbm import _split_masks
idx = pd.date_range("2024-01-01", "2025-03-31", freq="B")
train_mask, oos_mask = _split_masks(idx, "2025-03")
oos_pos = np.flatnonzero(oos_mask)
tr_pos = np.flatnonzero(train_mask)
assert oos_pos[0] == int(np.argmax(oos_mask))
# 样外月收盘绝不进训练 label:max 训练位 +2 严格早于样外首日
assert tr_pos.max() + 2 < oos_pos[0]
assert tr_pos.max() == oos_pos[0] - 3 # embargo=2 恰好剔够
assert not train_mask[oos_pos].any() # 训练/样外不重叠
# 训练下界=2025-03 前推 11 个自然月之首(2024-04-01)
assert idx[tr_pos].min() >= pd.Timestamp("2024-04-01")
# floor 之上只剔样外月+embargo 两天
above = np.flatnonzero(np.asarray(idx >= pd.Timestamp("2024-04-01")))
assert len(above) == len(tr_pos) + len(oos_pos) + 2
# 无样外月:训练也空(防全量误训)
t0, o0 = _split_masks(idx, "2026-01")
assert not o0.any() and not t0.any()
# ———— H-2:三路统一可评样本 mask ————
def test_unified_scoring_mask_same_sample_set(tmp_path, monkeypatch):
"""某源缺 2 股:未统一时 equal/icir 按 6 股评分、lgbm 只剩 4 股(<5
跳过)→days 分叉;统一 mask 后三路同一 date×stock 集合→days 一致."""
from sanguo_factor import challenger_lgbm as cl
idx = pd.date_range("2025-01-01", "2025-03-31", freq="B")
cols = ["a", "b", "c", "d", "e", "f"]
rng = np.random.default_rng(11)
s0 = pd.DataFrame(rng.normal(size=(len(idx), 6)), index=idx, columns=cols)
s0[["b", "c"]] = np.nan # 源 s0 缺 b/c 两股
s1 = pd.DataFrame(rng.normal(size=(len(idx), 6)), index=idx, columns=cols)
vdir = tmp_path / "vals"; vdir.mkdir()
s0.to_parquet(vdir / "s0.parquet")
s1.to_parquet(vdir / "s1.parquet")
close = pd.DataFrame(100 + np.cumsum(rng.normal(scale=0.4, size=(len(idx), 6)), axis=0),
index=idx, columns=cols)
monkeypatch.setattr(cl, "load_pool", lambda: {
"s0": {"direction": "+", "weight": 0.6},
"s1": {"direction": "-", "weight": 0.4}})
monkeypatch.setattr(cl, "_load_close_wide", lambda db, c, s, e: close)
out = cl.run_challenge(str(vdir), "fake.db", "2025-03-31", str(tmp_path))
doc = json.loads(Path(out).read_text())
days = {k: v["days"] for k, v in doc["signals"].items()}
assert len(set(days.values())) == 1 # 三路同一评分集
assert set(days.values()) == {0} # 统一后每日常截面=4 股<5,全跳
# ———— H-3/M-5:缺源与坏件走 skip 语义 ————
def test_missing_sources_skip_no_artifact(tmp_path, monkeypatch, capsys):
from sanguo_factor import challenger_lgbm as cl
idx = pd.date_range("2025-01-01", "2025-03-31", freq="B")
vdir = tmp_path / "vals"; vdir.mkdir()
pd.DataFrame(np.zeros((len(idx), 2)), index=idx,
columns=["a", "b"]).to_parquet(vdir / "s0.parquet")
monkeypatch.setattr(cl, "load_pool", lambda: {
"s0": {"direction": "+", "weight": 0.5},
"s1": {"direction": "+", "weight": 0.5}})
rc = cl.main(["--values-dir", str(vdir), "--as-of", "2025-03-31",
"--out-dir", str(tmp_path), "--vnpy-db", "fake.db"])
assert rc == 0
assert not (tmp_path / "challenger_lgbm").exists() # 不产半吊子件
out = capsys.readouterr().out
assert "missing sources" in out and "expected=2" in out and "loaded=1" in out
assert "s1" in out
def test_bad_parquet_structured_skip(tmp_path, monkeypatch, capsys):
"""坏件(非 parquet 字节)/空件(0 行)结构化剔除不裸抛,汇入缺源 skip."""
from sanguo_factor import challenger_lgbm as cl
idx = pd.date_range("2025-01-01", "2025-03-31", freq="B")
vdir = tmp_path / "vals"; vdir.mkdir()
pd.DataFrame(np.zeros((len(idx), 2)), index=idx,
columns=["a", "b"]).to_parquet(vdir / "s0.parquet")
(vdir / "s1.parquet").write_bytes(b"definitely not a parquet")
pd.DataFrame().to_parquet(vdir / "s2.parquet") # 0 行空件
monkeypatch.setattr(cl, "load_pool", lambda: {
n: {"direction": "+", "weight": 1.0 / 3} for n in ("s0", "s1", "s2")})
rc = cl.main(["--values-dir", str(vdir), "--as-of", "2025-03-31",
"--out-dir", str(tmp_path), "--vnpy-db", "fake.db"])
assert rc == 0
assert not (tmp_path / "challenger_lgbm").exists()
out = capsys.readouterr().out
assert "expected=3" in out and "loaded=1" in out
assert "s1" in out and "s2" in out
# ———— H-4:_load_close_wide 真实路径(monkeypatch universe.load_universe_bars) ————
def test_load_close_wide_real_path(monkeypatch):
import polars as pl
import sanguo_factor.universe as uni
from sanguo_factor import challenger_lgbm as cl
calls = {}
def fake_loader(db, start, end, symbols=None, limit=None):
calls.update(db=db, start=start, end=end, symbols=list(symbols or []))
return pl.DataFrame({
"vt_symbol": ["a.SSE", "a.SSE", "b.SSE", "b.SSE"],
"datetime": [pd.Timestamp("2025-01-02"), pd.Timestamp("2025-01-03")] * 2,
"close": [10.0, 11.0, 20.0, 21.0],
})
monkeypatch.setattr(uni, "load_universe_bars", fake_loader)
wide = cl._load_close_wide("fake.db", pd.Index(["a.SSE", "b.SSE"]),
"2025-01-01", "2025-01-31")
assert calls["db"] == "fake.db"
assert calls["start"] == "2025-01-01" and calls["end"] == "2025-01-31"
assert calls["symbols"] == ["a.SSE", "b.SSE"]
assert isinstance(wide.index, pd.DatetimeIndex)
assert list(wide.columns) == ["a.SSE", "b.SSE"]
assert wide.loc[pd.Timestamp("2025-01-03"), "a.SSE"] == pytest.approx(11.0)
# 反向:loader 空帧 → 空宽表不炸
monkeypatch.setattr(uni, "load_universe_bars", lambda *a, **k: pl.DataFrame(
schema={"vt_symbol": pl.Utf8, "datetime": pl.Datetime, "close": pl.Float64}))
wide0 = cl._load_close_wide("fake.db", pd.Index([]), "2025-01-01", "2025-01-31")
assert wide0.empty
# ———— M-6:cols_ref=池内 union ————
def test_cols_ref_is_union_of_pool_columns(tmp_path, monkeypatch):
from sanguo_factor import challenger_lgbm as cl
idx = pd.date_range("2025-01-01", "2025-03-31", freq="B")
rng = np.random.default_rng(5)
s0 = pd.DataFrame(rng.normal(size=(len(idx), 2)), index=idx, columns=["a", "b"])
s1 = pd.DataFrame(rng.normal(size=(len(idx), 2)), index=idx, columns=["b", "c"])
vdir = tmp_path / "vals"; vdir.mkdir()
s0.to_parquet(vdir / "s0.parquet")
s1.to_parquet(vdir / "s1.parquet")
captured = {}
def fake_close(db, cols_ref, start, end):
captured["cols"] = list(cols_ref)
return pd.DataFrame(100 + rng.normal(scale=0.4, size=(len(idx), len(cols_ref))),
index=idx, columns=list(cols_ref))
monkeypatch.setattr(cl, "load_pool", lambda: {
"s0": {"direction": "+", "weight": 0.5},
"s1": {"direction": "-", "weight": 0.5}})
monkeypatch.setattr(cl, "_load_close_wide", fake_close)
cl.run_challenge(str(vdir), "fake.db", "2025-03-31", str(tmp_path))
assert captured["cols"] == ["a", "b", "c"]
# ———— M-8:池冗余按日循环 ————
def test_redundancy_flags_per_day_unit():
from sanguo_factor.challenger_lgbm import _redundancy_flags
idx = pd.date_range("2025-01-01", periods=60, freq="B")
cols = [f"s{i}" for i in range(20)]
rng = np.random.default_rng(9)
x = pd.DataFrame(rng.normal(size=(60, 20)), index=idx, columns=cols)
z = pd.DataFrame(rng.normal(size=(60, 20)), index=idx, columns=cols)
flagged = _redundancy_flags({"x": x, "y": x.copy(), "z": z})
pairs = {(f["a"], f["b"]) for f in flagged}
assert ("x", "y") in pairs # 完全同源必标
assert ("x", "z") not in pairs and ("y", "z") not in pairs # 独立源不标
xy = next(f for f in flagged if (f["a"], f["b"]) == ("x", "y"))
assert xy["corr"] == pytest.approx(1.0, abs=1e-6)
def test_unsorted_or_dupe_index_parquet_rejected(tmp_path, monkeypatch, capsys):
"""codex 复审尾批:乱序/重复日期件 structured reject 并入缺源 skip.
embargo 的 shift(-2) 数学以位置序为地基,乱序/重复日期件必须拒."""
from sanguo_factor import challenger_lgbm as cl
idx = pd.date_range("2025-01-01", "2025-03-31", freq="B")
vdir = tmp_path / "vals"; vdir.mkdir()
good = pd.DataFrame(np.zeros((len(idx), 2)), index=idx, columns=["a", "b"])
good.to_parquet(vdir / "s0.parquet")
good.iloc[::-1].to_parquet(vdir / "s1.parquet") # 乱序:日期倒排
pd.concat([good.iloc[:1], good]).to_parquet(vdir / "s2.parquet") # 重复首日
monkeypatch.setattr(cl, "load_pool", lambda: {
n: {"direction": "+", "weight": 1.0 / 3} for n in ("s0", "s1", "s2")})
rc = cl.main(["--values-dir", str(vdir), "--as-of", "2025-03-31",
"--out-dir", str(tmp_path), "--vnpy-db", "fake.db"])
assert rc == 0
assert not (tmp_path / "challenger_lgbm").exists()
out = capsys.readouterr().out
assert "expected=3" in out and "loaded=1" in out
assert "s1" in out and "s2" in out
assert "乱序" in out