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

Function mp_split_model_receive

SwissArmyTransformer/sat/mpu/operation.py:76–91  ·  view source on GitHub ↗
(model, use_node_group=True)

Source from the content-addressed store, hash-verified

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()
78 src = get_node_src_rank() if use_node_group else get_model_parallel_src_rank()
79 def iter_repartition(module):
80 for name, sub_module in module.named_children():
81 if isinstance(sub_module, VocabParallelEmbedding):
82 torch.distributed.recv(sub_module.weight.data, src)
83 elif isinstance(sub_module, (ColumnParallelLinear, RowParallelLinear)):
84 torch.distributed.recv(sub_module.weight.data, src)
85 if sub_module.bias is not None and sub_module.bias.numel() != 0:
86 torch.distributed.recv(sub_module.bias.data, src)
87 else:
88 for n, p in sub_module.named_parameters(recurse=False):
89 torch.distributed.broadcast(p.data, src, group=group)
90 iter_repartition(sub_module)
91 iter_repartition(model)
92
93def mp_merge_model_rank0(model, model_full):
94 assert get_model_parallel_world_size() == torch.distributed.get_world_size(), "Merging model is only supported for model_parallel_size == world_size!"

Callers 2

from_pretrainedMethod · 0.90
from_pretrainedMethod · 0.90

Calls 5

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

Tested by

no test coverage detected