362 lines
17 KiB
Python
362 lines
17 KiB
Python
"""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
|