diff --git a/test_gpt2.cu b/test_gpt2.cu index e9fa1ec..f3df565 100644 --- a/test_gpt2.cu +++ b/test_gpt2.cu @@ -145,6 +145,8 @@ int main(int argc, char *argv[]) { cudaMemcpy(calculated_grads.fcb + l * 4*C, model.grads.fcb + l * 4*C, 4*C * sizeof(float), cudaMemcpyDeviceToHost); cudaMemcpy(calculated_grads.ln2w + l * C, model.grads.ln2w + l * C, C * sizeof(float), cudaMemcpyDeviceToHost); cudaMemcpy(calculated_grads.ln2b + l * C, model.grads.ln2b + l * C, C * sizeof(float), cudaMemcpyDeviceToHost); + cudaMemcpy(calculated_grads.attprojw + l * C * C, model.grads.attprojw + l * C * C, C * C * sizeof(float), cudaMemcpyDeviceToHost); + cudaMemcpy(calculated_grads.attprojb + l * C, model.grads.attprojb + l * C, C * sizeof(float), cudaMemcpyDeviceToHost); check_tensor(calculated_grads.lnfb, expected_grads.lnfb, C, "lnfb"); check_tensor(calculated_grads.lnfw, expected_grads.lnfw, C, "lnfw"); @@ -154,6 +156,8 @@ int main(int argc, char *argv[]) { check_tensor(calculated_grads.fcb + l * 4*C, expected_grads.fcb + l * 4*C, 4*C, "fcb"); check_tensor(calculated_grads.ln2w + l * C, expected_grads.ln2w + l * C, C, "ln2w"); check_tensor(calculated_grads.ln2b + l * C, expected_grads.ln2b + l * C, C, "ln2b"); + check_tensor(calculated_grads.attprojw + l * C * C, expected_grads.attprojw + l * C * C, C * C, "attprojw"); + check_tensor(calculated_grads.attprojb + l * C, expected_grads.attprojb + l * C, C, "attprojb"); } } diff --git a/train_gpt2.cu b/train_gpt2.cu index f7e030d..ce76a88 100644 --- a/train_gpt2.cu +++ b/train_gpt2.cu @@ -1342,11 +1342,11 @@ void gpt2_backward(GPT2 *model) { gelu_backward(dl_fch, l_fch, dl_fch_gelu, B*T*4*C); matmul_backward(dl_ln2, dl_fcw, dl_fcb, dl_fch, l_ln2, l_fcw, B, T, C, 4*C); layernorm_backward(dl_residual2, dl_ln2w, dl_ln2b, dl_ln2, l_residual2, l_ln2w, l_ln2_mean, l_ln2_rstd, B, T, C); + residual_backward(dresidual, dl_attproj, dl_residual2, B*T*C); + matmul_backward(dl_atty, dl_attprojw, dl_attprojb, dl_attproj, l_atty, l_attprojw, B, T, C, C); break; // break until we get all the other blocks in place, so we're only backwarding the last layer - // residual_backward(dresidual, dl_attproj, dl_residual2, B*T*C); - // matmul_backward(dl_atty, dl_attprojw, dl_attprojb, dl_attproj, l_atty, l_attprojw, B, T, C, C); // attention_backward(dl_qkv, dl_preatt, dl_att, dl_atty, l_qkv, 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(dresidual, dl_ln1w, dl_ln1b, dl_ln1, residual, l_ln1w, l_ln1_mean, l_ln1_rstd, B, T, C);