"""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())