69 lines
2.0 KiB
Python
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())
|