(self, new_model_parallel_size=None, full_weight=None)
| 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)) |
no test coverage detected