diff --git a/train_gpt2.cu b/train_gpt2.cu index a8fe6a9..04095e4 100644 --- a/train_gpt2.cu +++ b/train_gpt2.cu @@ -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 }