Files
2026-09-19 12:54:45 +08:00

69 lines
2.0 KiB
Python

"""Resume experiment runs filtered by model with custom concurrency.
Usage: python run_by_model.py <experiment_id> <model_id> <concurrency>
Reuses ExperimentRunner internals; safe to re-run (stale running runs are
reset to pending by the caller beforehand via runner.run_experiment normally —
here we handle pending/failed only, plus stale running without result).
"""
import asyncio
import sys
import uuid
from app.db import SessionLocal, init_db
from app.experiments.runner import ExperimentRunner
from app.models import ExperimentRun
async def main() -> None:
experiment_id = uuid.UUID(sys.argv[1])
model_id = sys.argv[2]
concurrency = int(sys.argv[3]) if len(sys.argv) > 3 else 4
init_db()
db = SessionLocal()
runner = ExperimentRunner(db)
# Reclaim runs orphaned in "running" by previously killed processes.
stale = (
db.query(ExperimentRun)
.filter_by(status="running", model_id=model_id)
.filter(~ExperimentRun.result.has())
.all()
)
for run in stale:
run.status = "pending"
db.commit()
runs = (
db.query(ExperimentRun)
.filter(
ExperimentRun.experiment_id == experiment_id,
ExperimentRun.model_id == model_id,
ExperimentRun.status.in_(["pending", "failed"]),
)
.all()
)
print(f"{model_id}: {len(runs)} runs to execute, concurrency={concurrency}")
semaphore = asyncio.Semaphore(concurrency)
async def execute(run):
async with semaphore:
return await runner._execute_run(run)
tasks = [asyncio.create_task(execute(run)) for run in runs]
done = 0
for coro in asyncio.as_completed(tasks):
result = await coro
done += 1
if result["status"] != "done":
print("WARN:", result)
if done % 10 == 0:
print(f"progress: {done}/{len(tasks)}")
print(f"{model_id}: completed {done}/{len(tasks)}")
if __name__ == "__main__":
asyncio.run(main())