mirror of
https://github.com/karpathy/llm.c.git
synced 2026-07-28 20:35:09 -04:00
reduce communication overhead for ZERO stage 1
This commit is contained in:
parent
57f70ea66b
commit
8b57cf6535
1 changed files with 14 additions and 6 deletions
|
|
@ -2295,12 +2295,20 @@ void gpt2_multi_gpu_accumulate(GPT2* model, MultiGpuConfig* multi_gpu_config) {
|
|||
if (multi_gpu_config->num_processes == 1) { return; }
|
||||
// Average all losses.
|
||||
model->accumulated_mean_loss = multi_gpu_cpu_float_mean(model->mean_loss, multi_gpu_config);
|
||||
// Average all gradients.
|
||||
ncclCheck(ncclAllReduce(model->grads_memory, model->grads_memory,
|
||||
model->num_parameters,
|
||||
ncclFloatX, ncclAvg,
|
||||
multi_gpu_config->nccl_comm,
|
||||
0));
|
||||
if(multi_gpu_config->zero_stage == 0) {
|
||||
// no ZERO == standard DDP: Average all gradients.
|
||||
ncclCheck(ncclAllReduce(model->grads_memory, model->grads_memory,
|
||||
model->num_parameters,
|
||||
ncclFloatX, ncclAvg,
|
||||
multi_gpu_config->nccl_comm, 0));
|
||||
} else if (multi_gpu_config->zero_stage == 1) {
|
||||
// ZERO-1: Get average gradient for local shard
|
||||
floatX* local_grads_memory = (floatX*) model->grads_memory + multi_gpu_config->shard_offset;
|
||||
ncclCheck(ncclReduceScatter(model->grads_memory, local_grads_memory,
|
||||
multi_gpu_config->shard_num_parameters,
|
||||
ncclFloatX, ncclAvg,
|
||||
multi_gpu_config->nccl_comm, 0));
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue