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

Method partition

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

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 3

divideFunction · 0.85
appendMethod · 0.80

Tested by

no test coverage detected