reduce communication overhead for ZERO stage 1

This commit is contained in:
Erik Schultheis 2024-05-16 14:12:29 +03:00
parent 57f70ea66b
commit 8b57cf6535

View file

@ -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
}