MCPcopy Create free account
hub / github.com/FlashSampling/FlashSampling / _torchrun_worker

Function _torchrun_worker

src/fused_mm_sampling/tp_info.py:65–80  ·  view source on GitHub ↗
(fn: Callable, args: tuple)

Source from the content-addressed store, hash-verified

63
64def _torchrun_worker(fn: Callable, args: tuple) -> None:
65 rank = int(os.environ["RANK"])
66 local_rank = int(os.environ["LOCAL_RANK"])
67
68 if rank == 0:
69 _print_gpu_topology()
70 print(
71 f"Rank {rank}: torchrun-bound CPUs: {sorted(os.sched_getaffinity(0))}",
72 flush=True,
73 )
74 torch.cuda.set_device(local_rank)
75 backend = "nccl"
76 if rank == 0:
77 print(f"Using distributed backend: '{backend}' (torchrun)")
78
79 dist.init_process_group(backend=backend, init_method="env://", device_id=local_rank)
80 try:
81 fn(*args)
82 finally:
83 dist.destroy_process_group()

Callers 1

run_maybe_distributedFunction · 0.85

Calls 2

_print_gpu_topologyFunction · 0.85
fnFunction · 0.50

Tested by

no test coverage detected