MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / repartition

Method repartition

SwissArmyTransformer/sat/mpu/layers.py:442–451  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

440 return output
441
442 def repartition(self):
443 assert self.input_size_per_partition == self.input_size
444 self.input_size_per_partition = divide(self.input_size, get_model_parallel_world_size())
445 mp_rank = get_model_parallel_rank()
446 self.original_weight = self.weight
447 self.weight = torch.nn.Parameter(torch.clone(
448 self.weight[:, mp_rank*self.input_size_per_partition
449 :(mp_rank+1)*self.input_size_per_partition],
450 ).detach())
451 del self.original_weight
452
453 def partition(self, new_model_parallel_size=None, full_weight=None):
454 assert self.input_size_per_partition == self.input_size or full_weight is not None

Callers 1

iter_repartitionFunction · 0.45

Calls 3

divideFunction · 0.85
get_model_parallel_rankFunction · 0.85

Tested by

no test coverage detected