Created
August 26, 2026 21:34
-
-
Save Apsu/c83c5a51b1636a6c16f876fa73155393 to your computer and use it in GitHub Desktop.
A script to wait on DNS resolution for a group of k8s pods to ensure pytorch communicator starts in sync
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| #!/usr/bin/env python3 | |
| """ | |
| wait.py - Dead simple DNS-based coordination for distributed training | |
| """ | |
| import argparse | |
| import socket | |
| import time | |
| import sys | |
| import os | |
| from datetime import datetime | |
| def log(message: str): | |
| """Simple timestamped logging""" | |
| timestamp = datetime.now().strftime("%H:%M:%S") | |
| print(f"[{timestamp}] {message}") | |
| def hostname_exists(hostname: str) -> bool: | |
| """Check if hostname resolves in DNS""" | |
| try: | |
| socket.gethostbyname(hostname) | |
| return True | |
| except socket.gaierror: | |
| return False | |
| def get_hostname(service_name: str, namespace: str, node_idx: int) -> str: | |
| """Generate hostname for a given node index""" | |
| if "-job" in service_name: | |
| # Kubernetes Job pattern | |
| base_name = service_name.replace("-job", "") | |
| return f"{service_name}-{node_idx}.{base_name}-service.{namespace}.svc.cluster.local" | |
| else: | |
| # StatefulSet pattern | |
| return f"{service_name}-{node_idx}.{service_name}-service.{namespace}.svc.cluster.local" | |
| def main(): | |
| parser = argparse.ArgumentParser(description="DNS-based node coordination") | |
| parser.add_argument("--expected-nodes", type=int, required=True) | |
| parser.add_argument("--service-name", type=str, required=True) | |
| parser.add_argument("--namespace", type=str, required=True) | |
| parser.add_argument("--timeout", type=int, default=600) | |
| parser.add_argument("--sync-delay", type=int, default=30) | |
| args = parser.parse_args() | |
| node_index = int(os.environ.get("JOB_COMPLETION_INDEX", "0")) | |
| print("=" * 50) | |
| print(f"Node {node_index}: Waiting for {args.expected_nodes} nodes via DNS") | |
| print("=" * 50) | |
| start_time = time.time() | |
| # Wait for all nodes to exist in DNS | |
| while True: | |
| if time.time() - start_time > args.timeout: | |
| log(f"ERROR: Timeout after {args.timeout} seconds") | |
| sys.exit(1) | |
| missing = [] | |
| for i in range(args.expected_nodes): | |
| hostname = get_hostname(args.service_name, args.namespace, i) | |
| if not hostname_exists(hostname): | |
| missing.append(i) | |
| if not missing: | |
| log(f"Node {node_index}: All {args.expected_nodes} nodes exist in DNS!") | |
| break | |
| log(f"Node {node_index}: Waiting for nodes {missing}...") | |
| time.sleep(5) | |
| # Fixed sync delay so all nodes start torchrun around the same time | |
| log(f"Node {node_index}: Waiting {args.sync_delay}s for sync...") | |
| time.sleep(args.sync_delay) | |
| log(f"Node {node_index}: Ready to start training!") | |
| return 0 | |
| if __name__ == "__main__": | |
| sys.exit(main()) |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment