94 lines
4.3 KiB
Python
94 lines
4.3 KiB
Python
# 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
|