(self, new_model_parallel_size=None, full_weight=None)
| 281 | del self.original_bias |
| 282 | |
| 283 | def partition(self, new_model_parallel_size=None, full_weight=None): |
| 284 | assert self.output_size_per_partition == self.output_size or full_weight is not None |
| 285 | flag = 1 |
| 286 | if full_weight is None: |
| 287 | full_weight = self.weight |
| 288 | flag = 2 |
| 289 | if new_model_parallel_size is None: |
| 290 | new_model_parallel_size = get_model_parallel_world_size() |
| 291 | output_size_per_partition = divide(self.output_size, new_model_parallel_size) |
| 292 | new_weights = [] |
| 293 | new_biases = [] |
| 294 | |
| 295 | mp_size = new_model_parallel_size |
| 296 | # weight is arranged as [stride0...stride1...stride2] * [input_size], extract non-contiguous parts |
| 297 | strides = [1]*self.stride if isinstance(self.stride, int) else self.stride # int means equal number of qkv, or ratios |
| 298 | assert full_weight.shape[0] % sum(strides) == 0, 'cannot divide weight evenly' |
| 299 | factor = full_weight.shape[0] // sum(strides) |
| 300 | # decompose weight according to strides |
| 301 | strided_weights, _acm = [], 0 |
| 302 | for i in range(len(strides)): |
| 303 | strided_weights.append(full_weight[_acm:_acm+factor*strides[i], :].detach()) |
| 304 | _acm += factor*strides[i] |
| 305 | |
| 306 | if flag == 2 and self.bias is not None and self.bias.numel() != 0: |
| 307 | # decompose bias according to strides |
| 308 | strided_biases, _acm = [], 0 |
| 309 | for i in range(len(strides)): |
| 310 | strided_biases.append(self.bias[_acm:_acm+factor*strides[i]].detach()) |
| 311 | _acm += factor*strides[i] |
| 312 | |
| 313 | for rank in range(new_model_parallel_size): |
| 314 | mp_rank = rank |
| 315 | new_weight = torch.cat([ |
| 316 | strided_weight[ |
| 317 | (strided_weight.shape[0]//mp_size)*mp_rank: |
| 318 | (strided_weight.shape[0]//mp_size)*(mp_rank+1) |
| 319 | ] |
| 320 | for strided_weight in strided_weights |
| 321 | ], dim=0).contiguous().view(output_size_per_partition, self.input_size) |
| 322 | new_weights.append(torch.clone(new_weight).detach()) |
| 323 | if flag == 2 and self.bias is not None and self.bias.numel() != 0: |
| 324 | new_bias = torch.cat([ |
| 325 | strided_bias[ |
| 326 | (strided_bias.shape[0]//mp_size)*mp_rank: |
| 327 | (strided_bias.shape[0]//mp_size)*(mp_rank+1) |
| 328 | ] |
| 329 | for strided_bias in strided_biases |
| 330 | ], dim=0).contiguous().view(output_size_per_partition) |
| 331 | new_biases.append(torch.clone(new_bias).detach()) |
| 332 | if flag == 1: |
| 333 | return new_weights |
| 334 | else: |
| 335 | return new_weights, new_biases |
| 336 | |
| 337 | def merge(self, new_weights, new_biases): |
| 338 | strides = [1]*self.stride if isinstance(self.stride, int) else self.stride # int means equal number of qkv, or ratios |
nothing calls this directly
no test coverage detected