PyNorch/norch/nn/parallel.py
2024-06-03 15:15:35 -03:00

50 lines
1.2 KiB
Python

from .module import *
import norch.distributed as dist
import os
import norch
class DistributedDataParallel(Module):
def __init__(self, module):
super().__init__()
self.module = module
self.broadcast_parameters()
self.register_grads_hooks()
def forward(self, *inputs, **kwargs):
return self.module(*inputs, **kwargs)
def broadcast_parameters(self):
"""
Broadcast parameters of device 0 to all devices
"""
for _, _, parameter in self.parameters():
dist.broadcast_tensor(parameter)
@staticmethod
def allreduce_grads_hook(grad):
"""
Everytime a gradient is assign to some value, it calculates mean of this gradient among all devices
"""
avg_grad = grad
if isinstance(grad, norch.Tensor):
dist.allreduce_sum_tensor(grad)
avg_grad = grad / dist.get_world_size()
return avg_grad
def register_grads_hooks(self):
"""
Everytime a gradient is assign it calls this allreduce hook
"""
for _, _, parameter in self.parameters():
parameter.register_hook(self.allreduce_grads_hook)