v1 distributed data parallel

This commit is contained in:
lucasdelimanogueira 2024-05-24 22:17:22 -03:00
parent 01b93132bc
commit 7a3b0f276b

21
norch/nn/distributed.py Normal file
View file

@ -0,0 +1,21 @@
from .module import *
import norch.distributed as dist
import os
class DistributedDataParallel(Module):
def __init__(self, module):
super().__init__()
self.module = module
def forward(self, *inputs, **kwargs):
return self.module(*inputs, **kwargs)
def backward(self):
self.module.backward()
for module, name, parameter in self.parameters():
dist.allreduce_sum_tensor(parameter)