feat(factor): 评估结果存储 factor_eval.db——runs/results两表,metrics按周期JSON,SANGUO_FACTOR_EVAL_DB可覆盖 [vps]

This commit is contained in:
2026-08-25 00:17:47 +08:00
parent 89a8dd1487
commit afe969e307
2 changed files with 208 additions and 0 deletions
+151
View File
@@ -0,0 +1,151 @@
# sanguo_factor/eval_store.py
"""批量评估结果落盘 factor_eval.db(eval_runs + eval_results).
metrics 按周期 {"1": {...}, "5": {...}, "10": {...}, "turnover": float} 存 JSON,
schema 平坦、加周期零迁移。
"""
import json
import os
import sqlite3
import uuid
from datetime import datetime
_SCHEMA = """
CREATE TABLE IF NOT EXISTS eval_runs(
run_id TEXT PRIMARY KEY,
label TEXT NOT NULL,
universe TEXT NOT NULL,
symbols_count INTEGER NOT NULL,
factors_total INTEGER NOT NULL,
start TEXT NOT NULL,
end TEXT NOT NULL,
status TEXT NOT NULL,
created_at TEXT NOT NULL,
finished_at TEXT,
factors_done INTEGER NOT NULL DEFAULT 0,
params_json TEXT NOT NULL
);
CREATE TABLE IF NOT EXISTS eval_results(
run_id TEXT NOT NULL,
factor TEXT NOT NULL,
category TEXT NOT NULL,
expression TEXT NOT NULL,
metrics_json TEXT NOT NULL,
PRIMARY KEY(run_id, factor)
);
"""
def default_eval_db_path() -> str:
"""SANGUO_FACTOR_EVAL_DB 覆盖;默认与 backtest_results.db 同目录."""
env = os.environ.get("SANGUO_FACTOR_EVAL_DB")
if env:
return env
from sanguo_data.config import load_config, find_config_path
cfg = load_config(find_config_path())
vnpy_db = cfg.data_paths.get("vnpy_db", "")
return os.path.join(os.path.dirname(os.path.abspath(vnpy_db)), "factor_eval.db")
def _conn(path: str) -> sqlite3.Connection:
conn = sqlite3.connect(path, timeout=30)
conn.execute("PRAGMA busy_timeout=30000")
return conn
def init_db(path: str) -> None:
os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True)
conn = _conn(path)
try:
conn.executescript(_SCHEMA)
conn.commit()
finally:
conn.close()
def create_run(path, label, universe, symbols_count, factors_total, start, end, params) -> str:
run_id = f"ev_{datetime.now():%Y%m%d_%H%M%S}_{uuid.uuid4().hex[:4]}"
conn = _conn(path)
try:
conn.execute(
"INSERT INTO eval_runs VALUES(?,?,?,?,?,?,?,?,?,?,0,?)",
(run_id, label, universe, symbols_count, factors_total, start, end,
"running", datetime.now().isoformat(timespec='seconds'), None,
json.dumps(params, ensure_ascii=False)),
)
conn.commit()
finally:
conn.close()
return run_id
def save_results(path, run_id, rows: list[dict]) -> None:
conn = _conn(path)
try:
conn.executemany(
"INSERT OR REPLACE INTO eval_results VALUES(?,?,?,?,?)",
[(run_id, r["factor"], r["category"], r["expression"],
json.dumps(r["metrics"], ensure_ascii=False)) for r in rows],
)
cur = conn.execute(
"SELECT COUNT(*) FROM eval_results WHERE run_id=?", (run_id,))
done = cur.fetchone()[0]
conn.execute("UPDATE eval_runs SET factors_done=? WHERE run_id=?", (done, run_id))
conn.commit()
finally:
conn.close()
def finish_run(path, run_id, status, factors_done) -> None:
conn = _conn(path)
try:
conn.execute(
"UPDATE eval_runs SET status=?, finished_at=?, factors_done=? WHERE run_id=?",
(status, datetime.now().isoformat(timespec='seconds'), factors_done, run_id),
)
conn.commit()
finally:
conn.close()
def list_runs(path) -> list[dict]:
conn = _conn(path)
conn.row_factory = sqlite3.Row
try:
rows = conn.execute(
"SELECT * FROM eval_runs ORDER BY created_at DESC").fetchall()
finally:
conn.close()
return [dict(r) for r in rows]
def get_rows(path, run_id, category=None, search=None) -> list[dict]:
q = "SELECT * FROM eval_results WHERE run_id=?"
args: list = [run_id]
if category:
q += " AND category=?"
args.append(category)
if search:
q += " AND (factor LIKE ? OR expression LIKE ?)"
args.extend([f"%{search}%", f"%{search}%"])
conn = _conn(path)
conn.row_factory = sqlite3.Row
try:
rows = conn.execute(q, args).fetchall()
finally:
conn.close()
return [{**dict(r), "metrics": json.loads(r["metrics_json"])} for r in rows]
def get_detail(path, run_id, factor) -> dict | None:
conn = _conn(path)
conn.row_factory = sqlite3.Row
try:
r = conn.execute(
"SELECT * FROM eval_results WHERE run_id=? AND factor=?",
(run_id, factor)).fetchone()
finally:
conn.close()
if r is None:
return None
return {**dict(r), "metrics": json.loads(r["metrics_json"])}
+57
View File
@@ -0,0 +1,57 @@
# tests/factor/test_eval_store.py
"""eval_store:建库/建run/UPSERT落盘/查询过滤/finish."""
import sys, os
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..")))
from sanguo_factor import eval_store
def _rows():
return [
{"factor": "alpha2", "category": "alpha101", "expression": "(-1)*ts_corr(...)",
"metrics": {"1": {"ic_mean": 0.05}, "turnover": 0.3}},
{"factor": "kmid", "category": "alpha158", "expression": "(close-open)/open",
"metrics": {"1": {"ic_mean": -0.02}, "turnover": 0.6}},
]
def test_roundtrip(tmp_path):
db = str(tmp_path / "factor_eval.db")
eval_store.init_db(db)
run_id = eval_store.create_run(db, label="冒烟", universe="custom", symbols_count=50,
factors_total=2, start="2024-01-01", end="2024-12-31",
params={"periods": [1, 5, 10]})
eval_store.save_results(db, run_id, _rows())
eval_store.finish_run(db, run_id, "done", factors_done=2)
runs = eval_store.list_runs(db)
assert len(runs) == 1 and runs[0]["run_id"] == run_id and runs[0]["status"] == "done"
rows = eval_store.get_rows(db, run_id)
assert len(rows) == 2
assert rows[0]["metrics"]["1"]["ic_mean"] == 0.05
assert eval_store.get_detail(db, run_id, "kmid")["category"] == "alpha158"
assert eval_store.get_detail(db, run_id, "nope") is None
def test_filter_by_category_and_search(tmp_path):
db = str(tmp_path / "factor_eval.db")
eval_store.init_db(db)
run_id = eval_store.create_run(db, label="x", universe="all_a", symbols_count=1,
factors_total=2, start="", end="", params={})
eval_store.save_results(db, run_id, _rows())
assert len(eval_store.get_rows(db, run_id, category="alpha101")) == 1
assert len(eval_store.get_rows(db, run_id, search="corr")) == 1 # 搜表达式
assert len(eval_store.get_rows(db, run_id, search="kmid")) == 1 # 搜因子名
def test_upsert_same_run(tmp_path):
db = str(tmp_path / "factor_eval.db")
eval_store.init_db(db)
run_id = eval_store.create_run(db, label="x", universe="all_a", symbols_count=1,
factors_total=1, start="", end="", params={})
eval_store.save_results(db, run_id, [_rows()[0]])
eval_store.save_results(db, run_id, [{**_rows()[0], "metrics": {"1": {"ic_mean": 0.09}}}] * 1)
rows = eval_store.get_rows(db, run_id)
assert len(rows) == 1 and rows[0]["metrics"]["1"]["ic_mean"] == 0.09