From 7a3b0f276b3ea48a94d459a854c0b21db57be313 Mon Sep 17 00:00:00 2001 From: lucasdelimanogueira Date: Fri, 24 May 2024 22:17:22 -0300 Subject: [PATCH] v1 distributed data parallel --- norch/nn/distributed.py | 21 +++++++++++++++++++++ 1 file changed, 21 insertions(+) create mode 100644 norch/nn/distributed.py diff --git a/norch/nn/distributed.py b/norch/nn/distributed.py new file mode 100644 index 0000000..a9a5de5 --- /dev/null +++ b/norch/nn/distributed.py @@ -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) + + + +