Skip to content

Instantly share code, notes, and snippets.

@Apsu
Created August 26, 2026 21:34
Show Gist options
  • Select an option

  • Save Apsu/c83c5a51b1636a6c16f876fa73155393 to your computer and use it in GitHub Desktop.

Select an option

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
#!/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