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

Function mp_split_model_rank0

SwissArmyTransformer/sat/mpu/operation.py:46–74  ·  view source on GitHub ↗

This function loads partitions from rank 0. It takes less memory when world size is large.

(model, model_full, use_node_group=True)

Source from the content-addressed store, hash-verified

44from .initialize import get_model_parallel_group, get_model_parallel_src_rank, get_model_parallel_world_size
45
46def mp_split_model_rank0(model, model_full, use_node_group=True):
47 """
48 This function loads partitions from rank 0.
49 It takes less memory when world size is large.
50 """
51 group = get_node_group() if use_node_group else get_model_parallel_group()
52 src = get_node_src_rank() if use_node_group else get_model_parallel_src_rank()
53 local_world_size = get_node_world_size() if use_node_group else get_model_parallel_world_size()
54 def iter_repartition(new_model, module):
55 for (new_name, sub_new_model), (name, sub_module) in zip(new_model.named_children(), module.named_children()):
56 if isinstance(sub_module, (ColumnParallelLinear, RowParallelLinear, VocabParallelEmbedding)):
57 new_weights, new_biases = sub_module.partition()
58 for i in range(local_world_size):
59 if i == 0:
60 sub_new_model.weight.data.copy_(new_weights[src%len(new_weights)])
61 else:
62 torch.distributed.send(new_weights[(src+i)%len(new_weights)].cuda(), src+i)
63 if new_biases:
64 for i in range(local_world_size):
65 if i == 0:
66 sub_new_model.bias.data.copy_(new_biases[src%len(new_weights)])
67 else:
68 torch.distributed.send(new_biases[(src+i)%len(new_biases)].cuda(), src+i)
69 else:
70 for (nn, np), (n, p) in zip(sub_new_model.named_parameters(recurse=False), sub_module.named_parameters(recurse=False)):
71 np.data.copy_(torch.clone(p.data).detach())
72 torch.distributed.broadcast(np.data, src, group=group)
73 iter_repartition(sub_new_model, sub_module)
74 iter_repartition(model, model_full)
75
76def mp_split_model_receive(model, use_node_group=True):
77 group = get_node_group() if use_node_group else get_model_parallel_group()

Callers 2

from_pretrainedMethod · 0.90
from_pretrainedMethod · 0.90

Calls 7

get_node_groupFunction · 0.85
get_model_parallel_groupFunction · 0.85
get_node_src_rankFunction · 0.85
get_node_world_sizeFunction · 0.85
iter_repartitionFunction · 0.85

Tested by

no test coverage detected