Files
sanguo_vnpy_v2/tests/factor/test_batch_eval.py
T

94 lines
4.3 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