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

Method partition

SwissArmyTransformer/sat/mpu/layers.py:453–476  ·  view source on GitHub ↗
(self, new_model_parallel_size=None, full_weight=None)

Source from the content-addressed store, hash-verified

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
455 flag = 1
456 if full_weight is None:
457 full_weight = self.weight
458 flag = 2
459 if new_model_parallel_size is None:
460 new_model_parallel_size = get_model_parallel_world_size()
461 input_size_per_partition = divide(self.input_size, new_model_parallel_size)
462 new_weights = []
463 new_biases = []
464 for rank in range(new_model_parallel_size):
465 mp_rank = rank
466 weight = torch.clone(
467 full_weight[:, mp_rank*input_size_per_partition
468 :(mp_rank+1)*input_size_per_partition],
469 ).detach()
470 new_weights.append(weight)
471 if flag == 2 and self.bias is not None and self.bias.numel() != 0:
472 new_biases.append(torch.clone(self.bias.data).detach())
473 if flag == 1:
474 return new_weights
475 else:
476 return new_weights, new_biases
477
478 def merge(self, new_weights, new_biases):
479 self.weight.data.copy_(torch.cat(new_weights, 1))

Callers 3

iter_repartitionFunction · 0.45
iter_mergeFunction · 0.45

Calls 3

divideFunction · 0.85
appendMethod · 0.80

Tested by

no test coverage detected