This function loads partitions from rank 0. It takes less memory when world size is large.
(model, model_full, use_node_group=True)
| 44 | from .initialize import get_model_parallel_group, get_model_parallel_src_rank, get_model_parallel_world_size |
| 45 | |
| 46 | def 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 | |
| 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() |
no test coverage detected