Files
sanguo_vnpy_v2/tests/factor/test_batch_eval.py
T

127 lines
5.6 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_batch_eval.py
"""批量引擎端到端:合成小库 → 2因子计算 → 指标落盘可查."""
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.batch_eval import run_batch_eval
from sanguo_factor import eval_store
_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)
"""
@pytest.fixture(scope="module")
def db(tmp_path_factory):
"""3 只 × ~400 交易日(2017-06~2018-12)合成库:一只单调涨/一只震荡/一只反着走."""
# Mount alpha datasets first to ensure factors are available
alpha_datasets.mount_all()
p = tmp_path_factory.mktemp("b")
db = str(p / "qt.db")
conn = sqlite3.connect(db)
conn.execute(_DDL)
import pandas as pd
days = pd.bdate_range("2017-06-01", "2018-12-28")
for i, day in enumerate(days):
d = day.strftime("%Y-%m-%d")
rows = [
("600000", "SSE", f"{d} 00:00:00", "d", 100.0, 1_000_000.0, 0,
10 + i * 0.01, 10 + i * 0.01, 10 + i * 0.01, 10 + i * 0.01), # 单调涨
("000001", "SZSE", f"{d} 00:00:00", "d", 100.0, 500_000.0, 0,
5, 5, 5, 5), # 平盘(截面另一端)
("300001", "SZSE", f"{d} 00:00:00", "d", 200.0, 900_000.0, 0,
20 - i * 0.005, 20 - i * 0.005, 20 - i * 0.005, 20 - i * 0.005), # 缓跌
]
conn.executemany("INSERT INTO dbbardata VALUES(?,?,?,?,?,?,?,?,?,?,?)", rows)
conn.commit(); conn.close()
return db
def test_run_batch_eval_end_to_end(db, tmp_path, monkeypatch):
eval_db = str(tmp_path / "factor_eval.db")
out = run_batch_eval(
factor_names=["ma_20", "roc_5"], start="2018-01-01", end="2018-06-30",
eval_db=eval_db, label="t", symbols=None, cfg=None, vnpy_db_override=db,
)
assert out["factors_done"] == 2
assert out["errors"] == []
rows = eval_store.get_rows(eval_db, out["run_id"])
assert {r["factor"] for r in rows} == {"ma_20", "roc_5"}
# ma_20 = ts_mean(close,20)/close:涨股该值持续低(均价低于现价) 跌股高 → 与次日收益负相关为主
m = eval_store.get_detail(eval_db, out["run_id"], "ma_20")["metrics"]
assert "1" in m and "5" in m and "10" in m and "turnover" in m
assert isinstance(m["1"]["ic_mean"], float)
def test_bad_factor_recorded_not_fatal(db, tmp_path):
eval_db = str(tmp_path / "factor_eval.db")
out = run_batch_eval(
factor_names=["不存在的因子", "kmid"], start="2018-01-01", end="2018-03-31",
eval_db=eval_db, label="t2", cfg=None, vnpy_db_override=db,
)
assert out["factors_done"] == 2
assert len(out["errors"]) == 1
detail = eval_store.get_detail(eval_db, out["run_id"], "不存在的因子")
assert "error" in detail["metrics"]
def test_resume_skips_done_factors(db, tmp_path):
eval_db = str(tmp_path / "f2.db")
out1 = run_batch_eval(factor_names=["ma_20", "roc_5"], start="2018-01-01", end="2018-06-30",
eval_db=eval_db, label="t", cfg=None, vnpy_db_override=db)
from sanguo_factor import eval_store
before = eval_store.get_rows(eval_db, out1["run_id"])
out2 = run_batch_eval(factor_names=["ma_20", "roc_5", "kmid"], start="2018-01-01", end="2018-06-30",
eval_db=eval_db, label="t", cfg=None, vnpy_db_override=db, run_id=out1["run_id"])
after = eval_store.get_rows(eval_db, out1["run_id"])
assert out2["run_id"] == out1["run_id"]
assert len(after) == 3 and len(before) == 2 # 新增 kmid
rows = eval_store.get_rows(eval_db, out1["run_id"])
assert {r["factor"] for r in rows} == {"ma_20", "roc_5", "kmid"}
# 空续跑:全部已完成时直接收尾不炸
out3 = run_batch_eval(factor_names=["ma_20"], start="2018-01-01", end="2018-06-30",
eval_db=eval_db, label="t", cfg=None, vnpy_db_override=db, run_id=out1["run_id"])
assert out3["factors_done"] >= 3
def test_factor_values_export(db, tmp_path):
"""factor_values_out: 每因子一份宽表 parquet(datetime×vt_symbol,选股回测契约)."""
from pathlib import Path
import pandas as pd
eval_db = str(tmp_path / "factor_eval.db")
out_dir = str(tmp_path / "values")
out = run_batch_eval(
factor_names=["ma_20"], start="2018-01-01", end="2018-06-30",
eval_db=eval_db, label="t-export", cfg=None, vnpy_db_override=db,
factor_values_out=out_dir,
)
assert out["errors"] == []
p = Path(out_dir) / "ma_20.parquet"
assert p.exists()
F = pd.read_parquet(p)
assert set(F.columns) == {"600000.SSE", "000001.SZSE", "300001.SZSE"}
# 与指标同窗: 导出索引全部落在评估窗口内,且存在有效截面值
assert F.index.min() >= pd.Timestamp("2018-01-01")
assert F.index.max() <= pd.Timestamp("2018-06-30")
assert F.notna().any().any()
def test_factor_values_none_keeps_default(db, tmp_path):
"""factor_values_out=None(默认): 不产生任何导出文件(既有评估批行为不变)."""
from pathlib import Path
eval_db = str(tmp_path / "factor_eval.db")
run_batch_eval(
factor_names=["ma_20"], start="2018-01-01", end="2018-03-31",
eval_db=eval_db, label="t-noexport", cfg=None, vnpy_db_override=db,
)
assert not list(Path(tmp_path).glob("**/*.parquet"))