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