(self)
| 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 |
no test coverage detected