(model, use_node_group=True)
| 74 | iter_repartition(model, model_full) |
| 75 | |
| 76 | def 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 | |
| 93 | def 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!" |
no test coverage detected