PyNorch/norch/distributed/distributed.py

31 lines
873 B
Python
Raw Permalink Normal View History

2024-05-24 15:12:53 -03:00
import os
import ctypes
2024-05-24 18:19:39 -03:00
from norch import Tensor, CTensor
2024-05-24 15:12:53 -03:00
def init_process_group(rank, world_size, backend='nccl'):
Tensor._C.init_process_group.argtypes = [ctypes.c_int, ctypes.c_int]
Tensor._C.init_process_group.restype = None
Tensor._C.init_process_group(rank, world_size)
2024-05-25 10:53:09 -03:00
def get_rank():
return int(os.getenv('OMPI_COMM_WORLD_RANK', 0))
2024-05-25 10:53:09 -03:00
def get_world_size():
return int(os.getenv('OMPI_COMM_WORLD_SIZE', 1))
2024-05-25 10:53:09 -03:00
def broadcast_tensor(tensor, src=0):
2024-05-24 18:19:39 -03:00
Tensor._C.broadcast_tensor.argtypes = [ctypes.POINTER(CTensor), ctypes.c_int]
2024-05-24 18:19:39 -03:00
Tensor._C.broadcast_tensor.restype = None
Tensor._C.broadcast_tensor(tensor.tensor, src)
2024-05-24 18:19:39 -03:00
2024-05-24 18:31:59 -03:00
def allreduce_sum_tensor(tensor):
2024-05-24 18:19:39 -03:00
Tensor._C.allreduce_sum_tensor.argtypes = [ctypes.POINTER(CTensor)]
Tensor._C.allreduce_sum_tensor.restype = None
2024-05-24 18:19:39 -03:00
Tensor._C.allreduce_sum_tensor(tensor.tensor)