vastai-utils/sched/__main__.py

123 lines
3.9 KiB
Python
Raw Normal View History

"""Entry point: python -m sched run --script jobs/mnist.py --trace --report report.html"""
from __future__ import annotations
import argparse
import asyncio
import logging
import sys
from .job import Job
from .scheduler import Scheduler
from .transport import SSHTransport
async def _run_mock(sched: Scheduler, job: Job):
"""Run with mock transport: create local nodes, skip provisioning."""
from .scheduler import Node
for rank in range(job.num_nodes):
conn = await sched.transport.connect(f"mock{rank}", 22)
sched.nodes.append(
Node(rank=rank, host=f"mock{rank}", port=22, conn=conn, status="ready")
)
sched._wrap_connections()
await sched.deploy()
job_mod = sched._load_job_module()
state = job_mod.make_model().state_dict()
for r in range(job.rounds):
state = await sched.run_round(state, r)
async def _run_compose(sched: Scheduler, job: Job):
"""Run against docker-compose workers: worker-0:22, worker-1:22, etc."""
from .scheduler import Node
for rank in range(job.num_nodes):
host = f"worker-{rank}"
conn = await sched.transport.connect(host, 22)
sched.nodes.append(
Node(rank=rank, host=host, port=22, conn=conn, status="ready")
)
sched._wrap_connections()
await sched.deploy()
job_mod = sched._load_job_module()
state = job_mod.make_model().state_dict()
for r in range(job.rounds):
state = await sched.run_round(state, r)
def main():
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s %(levelname)s %(name)s: %(message)s",
)
parser = argparse.ArgumentParser(prog="sched")
sub = parser.add_subparsers(dest="cmd")
p = sub.add_parser("run", help="Run a training job")
p.add_argument("--script", required=True, help="Path to job module")
p.add_argument("--nodes", type=int, default=2)
p.add_argument("--vram", type=int, default=8)
p.add_argument("--rounds", type=int, default=10)
p.add_argument("--local-steps", type=int, default=100)
p.add_argument("--lr", type=float, default=0.01)
p.add_argument("--batch-size", type=int, default=64)
p.add_argument("--trace", action="store_true", help="Enable live tracing to stderr")
p.add_argument("--report", metavar="PATH", help="Write HTML report to PATH")
p.add_argument("--mock", action="store_true", help="Use local mock transport (no vast.ai)")
p.add_argument("--compose", action="store_true", help="Connect to docker-compose workers (worker-0:22, worker-1:22, ...)")
p.add_argument("--docker-image", metavar="IMAGE", help="Override docker_image in job config")
args = parser.parse_args()
if not args.cmd:
parser.print_help()
return 1
tracer = None
if args.trace or args.report:
from .trace import Tracer
tracer = Tracer(file=sys.stderr)
job = Job(
script=args.script,
num_nodes=args.nodes,
min_vram=args.vram,
rounds=args.rounds,
local_steps=args.local_steps,
lr=args.lr,
batch_size=args.batch_size,
)
if args.docker_image:
job.docker_image = args.docker_image
if args.mock:
from tests.mock_transport import MockTransport
transport = MockTransport()
sched = Scheduler(job, transport, tracer=tracer)
asyncio.run(_run_mock(sched, job))
elif args.compose:
sched = Scheduler(job, SSHTransport(), tracer=tracer)
asyncio.run(_run_compose(sched, job))
else:
sched = Scheduler(job, SSHTransport(), tracer=tracer)
asyncio.run(sched.run())
if tracer:
tracer.gantt()
tracer.summary()
if args.report:
from .report import generate_html
path = generate_html(tracer, args.report)
print(f"Report: {path}", file=sys.stderr)
return 0
if __name__ == "__main__":
sys.exit(main())