first commit
This commit is contained in:
@@ -0,0 +1,68 @@
|
||||
"""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())
|
||||
Reference in New Issue
Block a user