From afe969e307ce0ce49cbcebc0349abddca0a96098 Mon Sep 17 00:00:00 2001 From: claude_dev Date: Tue, 25 Aug 2026 00:17:47 +0800 Subject: [PATCH] =?UTF-8?q?feat(factor):=20=E8=AF=84=E4=BC=B0=E7=BB=93?= =?UTF-8?q?=E6=9E=9C=E5=AD=98=E5=82=A8=20factor=5Feval.db=E2=80=94?= =?UTF-8?q?=E2=80=94runs/results=E4=B8=A4=E8=A1=A8,metrics=E6=8C=89?= =?UTF-8?q?=E5=91=A8=E6=9C=9FJSON,SANGUO=5FFACTOR=5FEVAL=5FDB=E5=8F=AF?= =?UTF-8?q?=E8=A6=86=E7=9B=96=20[vps]?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- sanguo_factor/eval_store.py | 151 ++++++++++++++++++++++++++++++++ tests/factor/test_eval_store.py | 57 ++++++++++++ 2 files changed, 208 insertions(+) create mode 100644 sanguo_factor/eval_store.py create mode 100644 tests/factor/test_eval_store.py diff --git a/sanguo_factor/eval_store.py b/sanguo_factor/eval_store.py new file mode 100644 index 0000000..0aa7430 --- /dev/null +++ b/sanguo_factor/eval_store.py @@ -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"])} diff --git a/tests/factor/test_eval_store.py b/tests/factor/test_eval_store.py new file mode 100644 index 0000000..0a377c8 --- /dev/null +++ b/tests/factor/test_eval_store.py @@ -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