From af9ac14b2f0f4c852c8eeb3c4f9b5a2d380fa760 Mon Sep 17 00:00:00 2001 From: Erik Schultheis Date: Sat, 20 Apr 2024 20:18:41 +0300 Subject: [PATCH 1/2] don't track preatt and vaccum across layers --- train_gpt2.cu | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/train_gpt2.cu b/train_gpt2.cu index 64d1dcc..ac1191c 100644 --- a/train_gpt2.cu +++ b/train_gpt2.cu @@ -996,7 +996,7 @@ void layernorm_backward(float* dinp, float* dweight, float* dbias, // inp (B,T,3C) -> qkvr (B,T,3C) -> preatt (B,NH,T,T) -> att (B,NH,T,T) -> vaccum (B,T,C) -> out (B,T,C) void attention_backward(float* dinp, float* dqkvr, float* dpreatt, float* datt, float* dvaccum, const float* dout, - const float* inp, const float* qkvr, const float* preatt, const float* att, const float* vaccum, + const float* inp, const float* qkvr, const float* att, int B, int T, int C, int NH) { const int block_size = 256; int HS = C / NH; // head size @@ -1171,7 +1171,7 @@ void fill_in_activation_sizes(size_t* act_sizes, int B, int T, GPT2Config config act_sizes[3] = L * B * T; // ln1_rstd act_sizes[4] = L * B * T * 3*C; // qkv act_sizes[5] = L * B * T * C; // atty - act_sizes[6] = L * B * NH * T * T; // preatt + act_sizes[6] = B * NH * T * T; // preatt act_sizes[7] = L * B * NH * T * T; // att act_sizes[8] = L * B * T * C; // attproj act_sizes[9] = L * B * T * C; // residual2 @@ -1189,7 +1189,7 @@ void fill_in_activation_sizes(size_t* act_sizes, int B, int T, GPT2Config config act_sizes[21] = B * T * V; // probs act_sizes[22] = B * T; // losses act_sizes[23] = L * B * T * 3*C; // qkvr - act_sizes[24] = L * B * T * C; // v_accum + act_sizes[24] = B * T * C; // v_accum act_sizes[25] = B * T * V; // dlogits (for fused_classifier) } @@ -1392,9 +1392,7 @@ void gpt2_forward(GPT2 *model, int* inputs, int* targets, int B, int T) { float* l_qkv = acts.qkv + l * B * T * 3*C; float* l_qkvr = acts.qkvr + l * B * T * 3*C; float* l_atty = acts.atty + l * B * T * C; - float* l_preatt = acts.preatt + l * B * NH * T * T; float* l_att = acts.att + l * B * NH * T * T; - float* l_v_accum = acts.v_accum + l * B * T * C; float* l_attproj = acts.attproj + l * B * T * C; float* l_residual2 = acts.residual2 + l * B * T * C; float* l_ln2 = acts.ln2 + l * B * T * C; @@ -1404,6 +1402,10 @@ void gpt2_forward(GPT2 *model, int* inputs, int* targets, int B, int T) { float* l_fch_gelu = acts.fch_gelu + l * B * T * 4*C; float* l_fcproj = acts.fcproj + l * B * T * C; float* l_residual3 = acts.residual3 + l * B * T * C; + // these are only needed as scratchpads for the forward pass, but + // need not be stored for backward + float* l_preatt = acts.preatt; + float* l_v_accum = acts.v_accum; // now do the forward pass layernorm_forward(l_ln1, l_ln1_mean, l_ln1_rstd, residual, l_ln1w, l_ln1b, B, T, C); @@ -1545,9 +1547,7 @@ void gpt2_backward(GPT2 *model) { float* l_qkv = acts.qkv + l * B * T * 3*C; float* l_qkvr = acts.qkvr + l * B * T * 3*C; float* l_atty = acts.atty + l * B * T * C; - float* l_preatt = acts.preatt + l * B * NH * T * T; float* l_att = acts.att + l * B * NH * T * T; - float* l_v_accum = acts.v_accum + l * B * T * C; float* l_residual2 = acts.residual2 + l * B * T * C; float* l_ln2 = acts.ln2 + l * B * T * C; float* l_ln2_mean = acts.ln2_mean + l * B * T; @@ -1575,7 +1575,7 @@ void gpt2_backward(GPT2 *model) { // layernorm backward does += to the dresidual, so it correctly accumulates grad from the MLP block above layernorm_backward(dresidual, dl_ln2w, dl_ln2b, dl_ln2, l_residual2, l_ln2w, l_ln2_mean, l_ln2_rstd, B, T, C); matmul_backward(dl_atty, dl_attprojw, dl_attprojb, dresidual, l_atty, l_attprojw, B, T, C, C); - attention_backward(dl_qkv, dl_qkvr, dl_preatt, dl_att, dl_v_accum, dl_atty, l_qkv, l_qkvr, l_preatt, l_att, l_v_accum, B, T, C, NH); + attention_backward(dl_qkv, dl_qkvr, dl_preatt, dl_att, dl_v_accum, dl_atty, l_qkv, l_qkvr, l_att, B, T, C, NH); matmul_backward(dl_ln1, dl_qkvw, dl_qkvb, dl_qkv, l_ln1, l_qkvw, B, T, C, 3*C); // layernorm backward does += to dresidual, so it correctly accumulates gradient for the Attention block above layernorm_backward(dresidual, dl_ln1w, dl_ln1b, dl_ln1, residual, l_ln1w, l_ln1_mean, l_ln1_rstd, B, T, C); From 24dd7289530f8679f878c193e0ac4cec9c9a0ea1 Mon Sep 17 00:00:00 2001 From: Erik Schultheis Date: Sat, 20 Apr 2024 01:20:36 +0300 Subject: [PATCH 2/2] only need one V-sized buffer that can switch from logits to dlogits in the fused backward pass --- test_gpt2.cu | 40 ++++++++++++++++++++++------------------ train_gpt2.cu | 31 ++++++++++++++----------------- 2 files changed, 36 insertions(+), 35 deletions(-) diff --git a/test_gpt2.cu b/test_gpt2.cu index 4f155b1..7d72f89 100644 --- a/test_gpt2.cu +++ b/test_gpt2.cu @@ -90,6 +90,27 @@ int main(int argc, char *argv[]) { // overall OK signal for the test int allok = 1; + // First, do target-free forward pass to validate logits + gpt2_forward(&model, x, NULL, B, T); + // at this point, target should be equal to expected_logits, let's compare + // copy logits to CPU so we can compare them + float* logits_cpu = (float*)mallocCheck(B * T * V * sizeof(float)); + cudaMemcpy(logits_cpu, model.acts.logits, B * T * V * sizeof(float), cudaMemcpyDeviceToHost); + int logits_ok = 1; + for (int i=0; i= 1e-2) { + printf("MISMATCH AT INDEX %d: ", i); + printf("%f %f\n", expected_logits[i],logits_cpu[i]); + logits_ok = 0; + break; + } + } + if(!logits_ok) { printf("NOT "); } + printf("OK (LOGITS)\n"); + // let's do 10 training iterations, following the pytorch code float losses[10]; for (int step = 0; step < 10; step++) { @@ -104,24 +125,7 @@ int main(int argc, char *argv[]) { if (step == 0) { // error checking at step 0 for reference activations - // at this point, target should be equal to expected_logits, let's compare - // copy logits to CPU so we can compare them - float* logits_cpu = (float*)mallocCheck(B * T * V * sizeof(float)); - cudaMemcpy(logits_cpu, model.acts.logits, B * T * V * sizeof(float), cudaMemcpyDeviceToHost); - int logits_ok = 1; - for (int i=0; i= 1e-2) { - printf("MISMATCH AT INDEX %d: ", i); - printf("%f %f\n", expected_logits[i],logits_cpu[i]); - logits_ok = 0; - break; - } - } - if(!logits_ok) { printf("NOT "); } - printf("OK (LOGITS)\n"); + allok = allok && logits_ok; free(logits_cpu); diff --git a/train_gpt2.cu b/train_gpt2.cu index ac1191c..f3f155b 100644 --- a/train_gpt2.cu +++ b/train_gpt2.cu @@ -738,8 +738,9 @@ __device__ SoftmaxParams prepare_softmax_blockwide_nofloat4(cg::thread_block_til } // same as 2 but not using float4 (see dev/cuda/classifier_fused.cu) -__global__ void fused_classifier_kernel3(float* dlogits, float* losses, float* probs, - const float* logits, const float* dlosses, const int* targets, +// will _update_ logits to logit gradients +__global__ void fused_classifier_kernel3(float* logits, float* losses, float* probs, + const float* dlosses, const int* targets, int B, int T, int V, int P) { namespace cg = cooperative_groups; cg::thread_block block = cg::this_thread_block(); @@ -769,10 +770,8 @@ __global__ void fused_classifier_kernel3(float* dlogits, float* losses, float* p if (probs != NULL) { probs[idx * P + i] = prob; } - if (dlogits != NULL) { - float indicator = (i == ix) ? 1.0f : 0.0f; - dlogits[idx * P + i] = (prob - indicator) * dloss; - } + float indicator = (i == ix) ? 1.0f : 0.0f; + logits[idx * P + i] = (prob - indicator) * dloss; } } @@ -1034,13 +1033,14 @@ void attention_backward(float* dinp, float* dqkvr, float* dpreatt, float* datt, cudaCheck(cudaGetLastError()); } -void fused_classifier3(float* dlogits, float* losses, - const float* logits, const float* dlosses, const int* targets, +// replaces logits with logit gradients +void fused_classifier3(float* logits, float* losses, + const float* dlosses, const int* targets, int B, int T, int V, int P) { const int block_size = 1024; const int N = B * T; const int grid_size = N; - fused_classifier_kernel3<<>>(dlogits, losses, NULL, logits, dlosses, targets, B, T, V, P); + fused_classifier_kernel3<<>>(logits, losses, NULL, dlosses, targets, B, T, V, P); cudaCheck(cudaGetLastError()); } @@ -1128,7 +1128,7 @@ float* malloc_and_point_parameters(ParameterTensors* params, size_t* param_sizes return params_memory; } -#define NUM_ACTIVATION_TENSORS 26 +#define NUM_ACTIVATION_TENSORS 25 typedef struct { float* encoded; // (B, T, C) float* ln1; // (L, B, T, C) @@ -1150,14 +1150,13 @@ typedef struct { float* lnf; // (B, T, C) float* lnf_mean; // (B, T) float* lnf_rstd; // (B, T) + // if we have targets, this will be the logit _gradients_. float* logits; // (B, T, V) float* probs; // (B, T, V) float* losses; // (B, T) // adding these two compared to the CPU .c code, needed for attention kernel as buffers float* qkvr; // (L, B, T, 3*C) float* v_accum; // (L, B, T, C) - // dlogits is used in fused_classifier. we backprop into it in the fused fwdbwd kernel for speed - float* dlogits; // (B,T,V) } ActivationTensors; void fill_in_activation_sizes(size_t* act_sizes, int B, int T, GPT2Config config) { @@ -1190,7 +1189,6 @@ void fill_in_activation_sizes(size_t* act_sizes, int B, int T, GPT2Config config act_sizes[22] = B * T; // losses act_sizes[23] = L * B * T * 3*C; // qkvr act_sizes[24] = B * T * C; // v_accum - act_sizes[25] = B * T * V; // dlogits (for fused_classifier) } float* malloc_and_point_activations(ActivationTensors* acts, const size_t* act_sizes) { @@ -1205,7 +1203,7 @@ float* malloc_and_point_activations(ActivationTensors* acts, const size_t* act_s &acts->preatt, &acts->att, &acts->attproj, &acts->residual2, &acts->ln2, &acts->ln2_mean, &acts->ln2_rstd, &acts->fch, &acts->fch_gelu, &acts->fcproj, &acts->residual3, &acts->lnf, &acts->lnf_mean, &acts->lnf_rstd, &acts->logits, &acts->probs, &acts->losses, - &acts->qkvr, &acts->v_accum, &acts->dlogits + &acts->qkvr, &acts->v_accum }; float* acts_memory_iterator = acts_memory; for (size_t i = 0; i < NUM_ACTIVATION_TENSORS; i++) { @@ -1428,7 +1426,7 @@ void gpt2_forward(GPT2 *model, int* inputs, int* targets, int B, int T) { if (targets != NULL) { // fused classifier: does the forward pass and first part of the backward pass // we're passing dlosses = NULL, which will default them to 1.0f/(B*T), i.e. uniform loss - fused_classifier3(acts.dlogits, acts.losses, acts.logits, NULL, model->targets, B, T, V, V); + fused_classifier3(acts.logits, acts.losses, NULL, model->targets, B, T, V, V); // for convenience also evaluate the mean loss (TODO re-think this compute+sync point) // move the (B,T) losses to CPU cudaCheck(cudaMemcpy(model->cpu_losses, acts.losses, B * T * sizeof(float), cudaMemcpyDeviceToHost)); @@ -1479,7 +1477,6 @@ void gpt2_backward(GPT2 *model) { bw_act_sizes[18] = 0; // lnf_mean bw_act_sizes[19] = 0; // lnf_rstd bw_act_sizes[21] = 0; // probs - bw_act_sizes[25] = 0; // dlogits are already in the forward pass // count up and allocate the space model->grads_acts_memory = malloc_and_point_activations(&model->grads_acts, bw_act_sizes); model->num_grad_acts = 0; @@ -1510,7 +1507,7 @@ void gpt2_backward(GPT2 *model) { // technically that is a small, inline backward() pass of calculating // total, final loss as the mean over all losses over all (B,T) positions in the batch // next: backward the classifier matmul - matmul_backward(grads_acts.lnf, grads.wte, NULL, acts.dlogits, acts.lnf, params.wte, B, T, C, V); + matmul_backward(grads_acts.lnf, grads.wte, NULL, acts.logits, acts.lnf, params.wte, B, T, C, V); // backward the final layernorm float* residual = acts.residual3 + (L-1) * B * T * C; // last residual is in residual3 float* dresidual = grads_acts.residual3; // the main buffer holding the gradient in the backward pass