feat(factor): 批量评估断点续跑——run_id复用+已落库因子跳过+空续跑免load直收尾;跨session CI容器重启只损失在途1因子,重挂即续 [vps]

This commit is contained in:
2026-08-25 20:56:22 +08:00
parent 3304ff46b2
commit 894c547297
3 changed files with 48 additions and 9 deletions
+27 -8
View File
@@ -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()),
+2 -1
View File
@@ -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']}")
+19
View File
@@ -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