v1 distributed data parallel
This commit is contained in:
parent
01b93132bc
commit
7a3b0f276b
1 changed files with 21 additions and 0 deletions
21
norch/nn/distributed.py
Normal file
21
norch/nn/distributed.py
Normal 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)
|
||||
|
||||
|
||||
|
||||
|
||||
Loading…
Add table
Add a link
Reference in a new issue