nvidia/C-RADIOv4-H
8527k
1from typing import Optional2 3import torch.distributed as dist4 5 6def get_rank(group: Optional[dist.ProcessGroup] = None):7 return dist.get_rank(group) if dist.is_initialized() else 08 9 10def get_world_size(group: Optional[dist.ProcessGroup] = None):11 return dist.get_world_size(group) if dist.is_initialized() else 112 13 14def barrier(group: Optional[dist.ProcessGroup] = None):15 if dist.is_initialized():16 dist.barrier(group)17 18 19class rank_gate:20 '''21 Execute the function on rank 0 first, followed by all other ranks. Useful when caches may need to be populated in a distributed environment.22 '''23 def __init__(self, func = None):24 self.func = func25 26 def __call__(self, *args, **kwargs):27 rank = get_rank()28 if rank == 0:29 result = self.func(*args, **kwargs)30 barrier()31 if rank > 0:32 result = self.func(*args, **kwargs)33 return result34 35 def __enter__(self, *args, **kwargs):36 if get_rank() > 0:37 barrier()38 39 def __exit__(self, *args, **kwargs):40 if get_rank() == 0:41 barrier()42 