From 894c54729707f9feda63f5b9755f8c5af7da4d6f Mon Sep 17 00:00:00 2001 From: claude_dev Date: Tue, 25 Aug 2026 20:56:22 +0800 Subject: [PATCH] =?UTF-8?q?feat(factor):=20=E6=89=B9=E9=87=8F=E8=AF=84?= =?UTF-8?q?=E4=BC=B0=E6=96=AD=E7=82=B9=E7=BB=AD=E8=B7=91=E2=80=94=E2=80=94?= =?UTF-8?q?run=5Fid=E5=A4=8D=E7=94=A8+=E5=B7=B2=E8=90=BD=E5=BA=93=E5=9B=A0?= =?UTF-8?q?=E5=AD=90=E8=B7=B3=E8=BF=87+=E7=A9=BA=E7=BB=AD=E8=B7=91?= =?UTF-8?q?=E5=85=8Dload=E7=9B=B4=E6=94=B6=E5=B0=BE;=E8=B7=A8session=20CI?= =?UTF-8?q?=E5=AE=B9=E5=99=A8=E9=87=8D=E5=90=AF=E5=8F=AA=E6=8D=9F=E5=A4=B1?= =?UTF-8?q?=E5=9C=A8=E9=80=941=E5=9B=A0=E5=AD=90,=E9=87=8D=E6=8C=82?= =?UTF-8?q?=E5=8D=B3=E7=BB=AD=20[vps]?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- sanguo_factor/batch_eval.py | 35 ++++++++++++++++++++++------- scripts/factor_research/run_eval.py | 3 ++- tests/factor/test_batch_eval.py | 19 ++++++++++++++++ 3 files changed, 48 insertions(+), 9 deletions(-) diff --git a/sanguo_factor/batch_eval.py b/sanguo_factor/batch_eval.py index af7acc8..13c8c9b 100644 --- a/sanguo_factor/batch_eval.py +++ b/sanguo_factor/batch_eval.py @@ -43,6 +43,7 @@ def run_batch_eval( cfg=None, progress_cb=None, vnpy_db_override: str | None = None, + run_id: str | None = None, ) -> dict: """跑一轮批量评估,结果增量写入 eval_db,返回摘要.""" from vnpy.alpha.dataset.utility import calculate_by_expression @@ -82,11 +83,29 @@ def run_batch_eval( universe_label = universe if symbols is None else "custom" eval_store.init_db(eval_db) - run_id = eval_store.create_run( - eval_db, label=label, universe=universe_label, - symbols_count=bars["vt_symbol"].n_unique(), factors_total=len(factor_names), - start=start, end=end, params={"limit": limit, "symbols": symbols[:20] if symbols else None}, # 截断防 4400 只全量塞 JSON - ) + if run_id is None: + run_id = eval_store.create_run( + eval_db, label=label, universe=universe_label, + symbols_count=bars["vt_symbol"].n_unique(), factors_total=len(factor_names), + start=start, end=end, params={"limit": limit, "symbols": symbols[:20] if symbols else None}, # 截断防 4400 只全量塞 JSON + ) + done_before = set() + else: + # 断点续跑:复用既有 run 行,跳过已落库因子(容器被 CI 重启后无损接续) + from sanguo_factor.eval_store import get_rows as _get_rows + done_before = {r["factor"] for r in _get_rows(eval_db, run_id)} + factor_names = [f for f in factor_names if f not in done_before] + if not factor_names: + # 空续跑:全部已完成时直接收尾不炸 + eval_store.finish_run(eval_db, run_id, "done", factors_done=len(done_before)) + return { + "run_id": run_id, + "factors_total": len(done_before), + "factors_done": len(done_before), + "errors": [], + "elapsed_sec": round(time.time() - t0, 1), + "symbols_count": int(bars["vt_symbol"].n_unique()), + } errors: list[str] = [] done = 0 @@ -103,11 +122,11 @@ def run_batch_eval( if progress_cb: progress_cb(done, len(factor_names), name) - eval_store.finish_run(eval_db, run_id, "done", factors_done=done) + eval_store.finish_run(eval_db, run_id, "done", factors_done=len(done_before) + done) return { "run_id": run_id, - "factors_total": len(factor_names), - "factors_done": done, + "factors_total": len(done_before) + len(factor_names), + "factors_done": len(done_before) + done, "errors": errors, "elapsed_sec": round(time.time() - t0, 1), "symbols_count": int(bars["vt_symbol"].n_unique()), diff --git a/scripts/factor_research/run_eval.py b/scripts/factor_research/run_eval.py index 20daa6d..9ee08fd 100644 --- a/scripts/factor_research/run_eval.py +++ b/scripts/factor_research/run_eval.py @@ -27,6 +27,7 @@ def main() -> int: ap.add_argument("--limit", type=int, default=None, help="随机抽样 N 只(种子42)") ap.add_argument("--label", default="batch1") ap.add_argument("--db", default=None, help="factor_eval.db 路径;默认 default_eval_db_path()") + ap.add_argument("--run-id", default=None, help="断点续跑:复用既有 run_id,跳过已落库因子") ap.add_argument("--list-factors", default="", metavar="CATEGORY", help="列出类目因子后退出") args = ap.parse_args() @@ -57,7 +58,7 @@ def main() -> int: print(f"[eval] {done}/{total} {current}", flush=True) out = run_batch_eval(factor_names, args.start, args.end, db_path, label=args.label, - symbols=symbols, limit=args.limit, cfg=None, progress_cb=_cb) + symbols=symbols, limit=args.limit, cfg=None, progress_cb=_cb, run_id=args.run_id) print(f"[eval] run_id={out['run_id']} done={out['factors_done']} " f"errors={len(out['errors'])} elapsed={out['elapsed_sec']}s symbols={out['symbols_count']}") diff --git a/tests/factor/test_batch_eval.py b/tests/factor/test_batch_eval.py index 0e9d665..fa091ac 100644 --- a/tests/factor/test_batch_eval.py +++ b/tests/factor/test_batch_eval.py @@ -72,3 +72,22 @@ def test_bad_factor_recorded_not_fatal(db, tmp_path): 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