#!/usr/bin/env python3 """ dtrain - CLI for distributed training on vast.ai Commands: search - Find cheap GPU offers, sorted by price rent - Rent instances by ID or auto-select N cheapest status - Show instance info (ID, GPU, status, SSH command) ps - Show training status (log age + last output line) deploy - Copy script to all instances and run in background run - Run arbitrary command on all instances destroy - Tear down all instances Requires: vastai CLI configured with API key, SSH key registered with vast.ai """ import argparse import json import os import subprocess import sys def run_vast(*args): """Run a vastai CLI command and return the result.""" cmd = ["vastai"] + list(args) result = subprocess.run(cmd, capture_output=True, text=True) if result.returncode != 0: print(f"Error: {result.stderr}", file=sys.stderr) return None return result.stdout # ============ COMMANDS ============ def cmd_search(args): """Find cheap GPU offers.""" query = f"gpu_ram>={args.vram} inet_down>=100" output = run_vast("search", "offers", query, "-o", "dph_total", "--limit", str(args.limit), "--raw") if not output: return 1 offers = json.loads(output) print(f"{'ID':<12} {'GPU':<16} {'$/hr':<8} {'VRAM':<6} {'Location':<20}") print("-" * 70) for o in offers: print(f"{o['id']:<12} {o['gpu_name']:<16} ${o['dph_total']:<7.4f} {o['gpu_ram']/1024:.0f}GB {o['geolocation'][:20]:<20}") return 0 def cmd_rent(args): """Rent instances by ID or pick n cheapest.""" # If --cheap N is set, find the N cheapest offers if args.cheap: query = f"gpu_ram>={args.vram} inet_down>=100" output = run_vast("search", "offers", query, "-o", "dph_total", "--limit", str(args.cheap), "--raw") if not output: return 1 offers = json.loads(output) offer_ids = [o['id'] for o in offers] total = sum(o['dph_total'] for o in offers) print(f"Selected {len(offers)} cheapest:") for o in offers: print(f" {o['id']} - {o['gpu_name']:<16} ${o['dph_total']:.4f}/hr {o['geolocation'][:20]}") print(f"Total: ${total:.4f}/hr") else: offer_ids = args.ids if not offer_ids: print("Error: provide IDs or use --cheap N") return 1 print(f"\nRenting {len(offer_ids)} instance(s)...") for oid in offer_ids: output = run_vast("create", "instance", str(oid), "--image", args.image, "--disk", str(args.disk)) if output: print(output.strip()) print("\nRun 'dtrain status' to check instance status.") return 0 def cmd_status(args): """Show running instances.""" output = run_vast("show", "instances", "--raw") if not output: print("No instances or error fetching.") return 1 instances = json.loads(output) if not instances: print("No running instances.") return 0 print(f"{'ID':<10} {'GPU':<16} {'Status':<12} {'$/hr':<8} {'SSH'}") print("-" * 80) for inst in instances: status = inst.get('actual_status', 'unknown') gpu = inst.get('gpu_name', 'N/A') cost = inst.get('dph_total') or 0 ssh_host = inst.get('ssh_host', '') ssh_port = inst.get('ssh_port', '') ssh_cmd = f"ssh -p {ssh_port} root@{ssh_host}" if ssh_host else "(not ready)" print(f"{inst['id']:<10} {gpu:<16} {status:<12} ${cost:<7.4f} {ssh_cmd}") return 0 def get_instances(): """Get list of running instances.""" output = run_vast("show", "instances", "--raw") if not output: return [] return json.loads(output) def run_on_instance(inst, command, timeout=10): """Run a command on an instance via SSH. Returns (returncode, output) or (None, msg) if not ready.""" ssh_host = inst.get('ssh_host', '') ssh_port = inst.get('ssh_port', '') if not ssh_host: return None, "not ready" # BatchMode=yes prevents password prompts from hanging cmd = ["ssh", "-p", str(ssh_port), "-o", "StrictHostKeyChecking=no", "-o", "ConnectTimeout=5", "-o", "BatchMode=yes", f"root@{ssh_host}", command] try: result = subprocess.run(cmd, capture_output=True, text=True, timeout=timeout) return result.returncode, result.stdout + result.stderr except subprocess.TimeoutExpired: return 1, "timeout" def cmd_deploy(args): """Deploy a script to all instances, verify, and run it.""" instances = get_instances() if not instances: print("No running instances.") return 1 filename = os.path.basename(args.script) dest_path = args.dest.rstrip('/') + '/' + filename ready_instances = [] # Phase 1: Copy script to all instances print("Copying script to instances...") for inst in instances: ssh_host = inst.get('ssh_host', '') ssh_port = inst.get('ssh_port', '') if not ssh_host: print(f" {inst['id']} - not ready, skipping") continue # Copy file dest = f"root@{ssh_host}:{args.dest}" cmd = ["scp", "-P", str(ssh_port), "-o", "StrictHostKeyChecking=no", args.script, dest] print(f" {inst['id']} - copying {args.script}...", end=" ") result = subprocess.run(cmd, capture_output=True, text=True) if result.returncode != 0: print(f"FAILED: {result.stderr.strip()}") continue # Verify ret, out = run_on_instance(inst, f"ls -la {dest_path}") if ret == 0: print("OK") ready_instances.append((inst, dest_path)) else: print(f"FAILED (verify): {out.strip()}") if not ready_instances: print("No instances ready for execution.") return 1 if args.no_run: return 0 # Phase 2: Run script on all instances (in background) print(f"\nStarting {filename} on {len(ready_instances)} instance(s)...") for inst, path in ready_instances: logfile = filename.replace('.py', '.log') run_cmd = f"cd {args.dest} && nohup python {filename} > {logfile} 2>&1 &" ret, out = run_on_instance(inst, run_cmd) if ret is None: print(f" {inst['id']} - not ready") elif ret == 0: print(f" {inst['id']} - started (log: {args.dest}{logfile})") else: print(f" {inst['id']} - FAILED: {out.strip()}") # Summary total = len(instances) skipped = total - len(ready_instances) print(f"\nSummary: {len(ready_instances)}/{total} instances deployed") if skipped: print(f" {skipped} skipped/failed") return 0 def cmd_ps(args): """Show training status by reading log files on each node.""" instances = get_instances() if not instances: print("No running instances.") return 1 # Print header immediately print(f"{'ID':<10} {'GPU':<14} {'Status':<10} {'Log age':<10} {'Last output'}") print("-" * 100) for inst in instances: iid = inst['id'] gpu = inst.get('gpu_name', 'N/A')[:14] status = inst.get('actual_status', 'unknown') if not inst.get('ssh_host'): print(f"{iid:<10} {gpu:<14} {status:<10} {'-':<10} (no ssh)") continue # Show we're checking this node print(f"{iid:<10} {gpu:<14} {status:<10} ", end="", flush=True) # Find most recent .log file, output format: "seconds_since_modified|last_line" log_cmd = 'log=$(ls -t /workspace/*.log 2>/dev/null | head -1); if [ -n "$log" ]; then age=$(($(date +%s) - $(stat -c %Y "$log"))); echo "$age|$(tail -1 "$log" | head -c 60)"; else echo "NO_LOG|"; fi' ret, out = run_on_instance(inst, log_cmd) # Check for SSH/connection errors error_phrases = ['connection refused', 'no route to host', 'connection timed out', 'permission denied', 'host key verification failed', 'network is unreachable', 'timeout'] out_lower = (out or '').lower() is_conn_error = any(err in out_lower for err in error_phrases) if is_conn_error: err_msg = out.strip().split('\n')[0][:40] if out else 'SSH error' print(f"{'-':<10} ({err_msg})") continue # Parse "age|last_line" - skip vast.ai SSH banner lines by looking for our format result_line = None for line in (out or '').strip().split('\n'): if '|' in line and (line[0].isdigit() or line.startswith('NO_LOG')): result_line = line break if result_line: age_str, last_line = result_line.split('|', 1) if age_str == 'NO_LOG': print(f"{'-':<10} (no log file)") else: # Convert raw seconds to friendly format if age_str.endswith('s'): age_str = age_str # already formatted elif age_str.isdigit(): secs = int(age_str) if secs < 60: age_str = f"{secs}s ago" elif secs < 3600: age_str = f"{secs//60}m ago" else: age_str = f"{secs//3600}h ago" print(f"{age_str:<10} {last_line}") else: print(f"{'-':<10} (no output)") return 0 def cmd_run(args): """Run a command on all instances.""" instances = get_instances() if not instances: print("No running instances.") return 1 command = " ".join(args.command) # Track results for summary results = {'ok': 0, 'failed': 0, 'not_ready': 0} for inst in instances: iid = inst['id'] gpu = inst.get('gpu_name', 'N/A') status = inst.get('actual_status', 'unknown') ssh_host = inst.get('ssh_host', '') # Header with instance info header = f"[{iid}] {gpu} ({status})" if not ssh_host: print(f"{header} -- SKIPPED: {status}, waiting for SSH") results['not_ready'] += 1 continue ret, out = run_on_instance(inst, command) if ret is None: print(f"{header} -- SKIPPED: SSH not available") results['not_ready'] += 1 elif ret == 0: print(f"{header} -- OK") if out.strip(): # Indent output for readability for line in out.strip().split('\n'): print(f" {line}") results['ok'] += 1 else: print(f"{header} -- FAILED (exit {ret})") if out.strip(): for line in out.strip().split('\n'): print(f" {line}") results['failed'] += 1 # Summary total = len(instances) print(f"\n--- {results['ok']}/{total} ok", end="") if results['failed']: print(f", {results['failed']} failed", end="") if results['not_ready']: print(f", {results['not_ready']} not ready", end="") print(" ---") return 0 if results['failed'] == 0 else 1 def cmd_destroy(args): """Destroy instances.""" output = run_vast("show", "instances", "--raw") if not output: return 1 instances = json.loads(output) if not instances: print("No instances to destroy.") return 0 print(f"Destroying {len(instances)} instance(s)...") for inst in instances: print(f" {inst['id']} - {inst.get('gpu_name', 'N/A')}") run_vast("destroy", "instance", str(inst['id'])) print("Done.") return 0 # ============ MAIN ============ def main(): parser = argparse.ArgumentParser(prog="dtrain", description="CLI for distributed training on vast.ai") sub = parser.add_subparsers(dest="cmd") # search p = sub.add_parser("search", help="Find cheap GPU offers") p.add_argument("--vram", type=int, default=8, help="Min VRAM in GB (default: 8)") p.add_argument("--limit", type=int, default=10, help="Max results (default: 10)") p.set_defaults(func=cmd_search) # rent p = sub.add_parser("rent", help="Rent instances by ID or pick cheapest") p.add_argument("ids", type=int, nargs="*", help="Offer ID(s) to rent") p.add_argument("-n", "--cheap", type=int, help="Rent N cheapest offers instead") p.add_argument("--vram", type=int, default=8, help="Min VRAM for --cheap (default: 8)") p.add_argument("--image", default="pytorch/pytorch:latest", help="Docker image") p.add_argument("--disk", type=int, default=20, help="Disk GB (default: 20)") p.set_defaults(func=cmd_rent) # status p = sub.add_parser("status", help="Show running instances") p.set_defaults(func=cmd_status) # ps p = sub.add_parser("ps", help="Show training status (log age + last line)") p.set_defaults(func=cmd_ps) # deploy p = sub.add_parser("deploy", help="Deploy script to all instances and run it") p.add_argument("--script", required=True, help="Script to deploy") p.add_argument("--dest", default="/workspace/", help="Destination path (default: /workspace/)") p.add_argument("--no-run", action="store_true", help="Only copy, don't run the script") p.set_defaults(func=cmd_deploy) # run p = sub.add_parser("run", help="Run command on all instances") p.add_argument("command", nargs="+", help="Command to run") p.set_defaults(func=cmd_run) # destroy p = sub.add_parser("destroy", help="Destroy all instances") p.set_defaults(func=cmd_destroy) args = parser.parse_args() if not args.cmd: parser.print_help() return 0 return args.func(args) if __name__ == "__main__": sys.exit(main())