From 8b57cf65355c453d394231bfca478ef2d270bda5 Mon Sep 17 00:00:00 2001 From: Erik Schultheis Date: Thu, 16 May 2024 14:12:29 +0300 Subject: [PATCH] reduce communication overhead for ZERO stage 1 --- train_gpt2.cu | 20 ++++++++++++++------ 1 file changed, 14 insertions(+), 6 deletions(-) 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 }