From 58df2629aecd5437e10fd52fa6f7bc8fc0488626 Mon Sep 17 00:00:00 2001 From: ademeure Date: Mon, 22 Apr 2024 18:09:18 +0100 Subject: [PATCH] Gradients now working in BF16 mode! (still need adam etc...) --- test_gpt2.cu | 44 +++++++-- train_gpt2.cu | 265 +++++++++++++++++++++++++++++++++++++------------- 2 files changed, 233 insertions(+), 76 deletions(-) diff --git a/test_gpt2.cu b/test_gpt2.cu index 93c2e1b..eef74b2 100644 --- a/test_gpt2.cu +++ b/test_gpt2.cu @@ -7,10 +7,10 @@ int check_tensor(float *a, float *b, int n, const char* label) { int ok = 1; printf("%s\n", label); for (int i = 0; i < n; i++) { - if (fabsf(a[i] - b[i]) <= 1e-2) { - if (i < print_upto) { printf("OK "); } + if (fabsf(a[i] - b[i]) <= 1e-1) { + if (i < print_upto) { printf("OK %d ", i); } } else { - if (i < print_upto) { printf("NOT OK "); } + if (i < print_upto) { printf("NOT OK %d ", i); } ok = 0; } if (i < print_upto) { printf("%f %f\n", a[i], b[i]); } @@ -70,9 +70,10 @@ int main(int argc, char *argv[]) { ParameterTensors expected_grads; // will be read from file (from PyTorch) ParameterTensors calculated_grads; // will be calculated by us - float* expected_grads_memory = malloc_and_point_parameters(&expected_grads, model.param_sizes, 0); - float* calculated_grads_memory = malloc_and_point_parameters(&calculated_grads, model.param_sizes, 0); - + float* expected_grads_memory = malloc_and_point_parameters(&expected_grads, model.param_elements, model.param_sizeof, 0); + float* calculated_grads_memory = malloc_and_point_parameters(&calculated_grads, model.param_elements, model.param_sizeof, 0); + float* converted_grads_memory = (float*)mallocCheck(model.num_parameters * sizeof(float)); + // inputs and expected outputs, only used for error checking int* x = (int*)mallocCheck(B * T * sizeof(int)); int* y = (int*)mallocCheck(B * T * sizeof(int)); @@ -97,11 +98,18 @@ int main(int argc, char *argv[]) { float* logits_cpu = (float*)mallocCheck(B * T * V * sizeof(float)); cudaMemcpy(logits_cpu, model.acts.output, B * T * V * sizeof(float), cudaMemcpyDeviceToHost); int logits_ok = 1; + + // FP16 and lower require very high tolerance + float accuracy_threshold = 1e-2; + #if CUBLAS_LOWP != CUDA_R_32F + accuracy_threshold = 10; + #endif + for (int i=0; i= 1e-2) { + if (fabsf(expected_logits[i] - logits_cpu[i]) >= accuracy_threshold) { printf("MISMATCH AT INDEX %d: ", i); printf("%f %f\n", expected_logits[i],logits_cpu[i]); logits_ok = 0; @@ -172,8 +180,27 @@ int main(int argc, char *argv[]) { // check_tensor(calculated_grads.wpe, expected_grads.wpe, maxT * C, "wpe"); // compare the gradients ona the parameters all at once + + // Convert gradients back to FP32 for comparison cudaMemcpy(calculated_grads_memory, model.grads_memory, model.num_parameters * sizeof(float), cudaMemcpyDeviceToHost); - check_tensor(calculated_grads_memory, expected_grads_memory, model.num_parameters, "grads"); + + char* calculated_iterator = (char*)calculated_grads_memory; + float* converted_iterator = (float*)converted_grads_memory; + + for (size_t i = 0; i < NUM_PARAMETER_TENSORS; i++) { + if (model.param_sizeof[i] == sizeof(float)) { + memcpy(converted_iterator, calculated_iterator, model.param_elements[i] * sizeof(float)); + } else { + // TODO: Currently only support float or floatX (cannot mix and match FP16/BF16 etc...) + assert(model.param_sizeof[i] == sizeof(floatX)); + for (size_t j = 0; j < model.param_elements[i]; j++) { + converted_iterator[j] = ((floatX*)calculated_iterator)[j]; + } + } + calculated_iterator += model.param_elements[i] * model.param_sizeof[i]; + converted_iterator += model.param_elements[i]; + } + check_tensor(converted_grads_memory, expected_grads_memory, model.num_parameters, "grads"); } gpt2_update(&model, 1e-4f, 0.9f, 0.999f, 1e-8f, 0.01f, step+1); @@ -217,6 +244,7 @@ int main(int argc, char *argv[]) { free(expected_loss); free(expected_grads_memory); free(calculated_grads_memory); + free(converted_grads_memory); gpt2_free(&model); cudaCheck(cudaFree(cublaslt_workspace)); cublasCheck(cublasDestroy(cublas_handle)); diff --git a/train_gpt2.cu b/train_gpt2.cu index 5deb0fd..a86c1c9 100644 --- a/train_gpt2.cu +++ b/train_gpt2.cu @@ -29,7 +29,7 @@ the layernorms are connected to the residuals so we += in layernorm backward. // ---------------------------------------------------------------------------- // CUDA precision settings -typedef float floatX; // linear layer weights and activations +typedef __nv_bfloat16 floatX; // linear layer weights and activations // CUDA_R_32F or CUDA_R_16BF or CUDA_R_16F #define CUBLAS_LOWP CUDA_R_16BF // if CUDA_R_16F then CUBLAS_COMPUTE_32F or CUBLAS_COMPUTE_16F @@ -72,6 +72,21 @@ cublasLtHandle_t cublaslt_handle; namespace cg = cooperative_groups; +// ---------------------------------------------------------------------------- +template +__global__ void print_tensorT(T* tensor, int N) { + for (int i = 0; i < N; i++) { + printf("%f\n", (float)tensor[i]); + } +} + +// Function that calls print_tensorT, sychronises, and prints message +template +void print_tensor(T* tensor, int N=1000, const char* message="-----------------") { + print_tensorT<<<1, 1>>>(tensor, N); + cudaCheck(cudaDeviceSynchronize()); + printf("\n\n\n %s \n\n\n\n\n\n", message); +} // ---------------------------------------------------------------------------- // fread convenience utils, with nice handling of error checking using macros // simple replace fopen, fread, fclose with fopenCheck, freadCheck, fcloseCheck @@ -186,7 +201,7 @@ __global__ void encoder_forward_kernel2(float* out, // really bad naive kernel with atomicAdd __global__ void encoder_backward_kernel(float* dwte, float* dwpe, - const float* dout, const int* inp, + const floatX* dout, const int* inp, int B, int T, int C) { int idx = blockIdx.x * blockDim.x + threadIdx.x; int N = B * T * C; @@ -199,12 +214,12 @@ __global__ void encoder_backward_kernel(float* dwte, float* dwpe, int ix = inp[b * T + t]; - const float* dout_btc = dout + b * T * C + t * C + c; + const floatX* dout_btc = dout + b * T * C + t * C + c; float* dwte_ix = dwte + ix * C + c; float* dwpe_tc = dwpe + t * C + c; - atomicAdd(dwte_ix, *dout_btc); - atomicAdd(dwpe_tc, *dout_btc); + atomicAdd(dwte_ix, (float)*dout_btc); + atomicAdd(dwpe_tc, (float)*dout_btc); } } @@ -326,7 +341,7 @@ __global__ void permute_kernel(float* q, float* k, float* v, } } -__global__ void permute_kernel_backward(float* dinp, +__global__ void permute_kernel_backward(floatX* dinp, const float* dq, const float* dk, const float* dv, int B, int N, int NH, int d) { int idx = blockIdx.x * blockDim.x + threadIdx.x; @@ -339,9 +354,9 @@ __global__ void permute_kernel_backward(float* dinp, int d_ = rest % d; int inp_idx = (b * N * 3 * NH * d) + (n * 3 * NH * d) + (0 * NH * d) + (nh_ * d) + d_; - dinp[inp_idx] = dq[idx]; - dinp[inp_idx + NH * d] = dk[idx]; - dinp[inp_idx + 2 * (NH * d)] = dv[idx]; + dinp[inp_idx] = (floatX)dq[idx]; + dinp[inp_idx + NH * d] = (floatX)dk[idx]; + dinp[inp_idx + 2 * (NH * d)] = (floatX)dv[idx]; } } @@ -361,7 +376,7 @@ __global__ void unpermute_kernel(float* inp, floatX *out, int B, int N, int NH, } } -__global__ void unpermute_kernel_backward(float* dinp, const float *dout, int B, int N, int NH, int d) { +__global__ void unpermute_kernel_backward(float* dinp, const floatX *dout, int B, int N, int NH, int d) { int idx = blockIdx.x * blockDim.x + threadIdx.x; if (idx < B * NH * N * d) { int b = idx / (NH * N * d); @@ -371,7 +386,7 @@ __global__ void unpermute_kernel_backward(float* dinp, const float *dout, int B, int n = rest / d; int d_ = rest % d; int other_idx = (b * NH * N * d) + (n * NH * d) + (nh_ * d) + d_; - dinp[idx] = dout[other_idx]; + dinp[idx] = (float)dout[other_idx]; } } @@ -469,17 +484,17 @@ __global__ void gelu_forward_kernel(floatX* out, const floatX* inp, int N) { } } -__global__ void gelu_backward_kernel(float* dinp, const float* inp, const float* dout, const int N) { +__global__ void gelu_backward_kernel(floatX* dinp, const floatX* inp, const floatX* dout, const int N) { int i = blockIdx.x * blockDim.x + threadIdx.x; if (i < N) { - float x = inp[i]; + float x = (float)inp[i]; float cube = 0.044715f * x * x * x; float tanh_arg = GELU_SCALING_FACTOR * (x + cube); float tanh_out = tanhf(tanh_arg); float coshf_out = coshf(tanh_arg); float sech_out = 1.0f / (coshf_out * coshf_out); float local_grad = 0.5f * (1.0f + tanh_out) + x * 0.5f * sech_out * GELU_SCALING_FACTOR * (1.0f + 3.0f * 0.044715f * x * x); - dinp[i] = local_grad * dout[i]; + dinp[i] = (floatX)(local_grad * (float)dout[i]); } } @@ -614,17 +629,42 @@ __global__ void matmul_backward_bias_kernel2(float* dbias, const float* dout, in // first, thread coarsening to sum reduce the problem size from B*T to 32 float sum = 0.0f; for(int i = warp.thread_rank(); i < BT; i += warp.size()) { - sum += dout[i * OC + idx]; + sum += (float)dout[i * OC + idx]; } // now do a warp-level reduce to get the sum across the 32 threads in this warp sum = cg::reduce(warp, sum, cg::plus{}); // write the result to output (global memory) if(warp.thread_rank() == 0) { - dbias[idx] += sum; + dbias[idx] = (float)((float)dbias[idx] + sum); } } -__global__ void layernorm_backward_kernel(float* dinp, float* dweight, float* dbias, +// cooperative groups solution, one warp per output channel +__global__ void matmul_backward_bias_kernel3_lowp(floatX* dbias, const floatX* dout, int B, int T, int OC) { + // dout is (B, T, OC), dbias is (OC) + // e.g. if block_size = 128, then we have 4 warps per block, each in charge of one output channel + namespace cg = cooperative_groups; + cg::thread_block block = cg::this_thread_block(); + cg::thread_block_tile<32> warp = cg::tiled_partition<32>(block); + // meta_group_size is the number of warps in a block (e.g. 4), meta_group_rank is the warp index (0,1,2,3) + int idx = blockIdx.x * warp.meta_group_size() + warp.meta_group_rank(); + if(idx >= OC) { return; } + int BT = B * T; // number of elements to reduce in total, per channel + // first, thread coarsening to sum reduce the problem size from B*T to 32 + float sum = 0.0f; + for(int i = warp.thread_rank(); i < BT; i += warp.size()) { + sum += (float)dout[i * OC + idx]; + } + // now do a warp-level reduce to get the sum across the 32 threads in this warp + sum = cg::reduce(warp, sum, cg::plus{}); + // write the result to output (global memory) + if(warp.thread_rank() == 0) { + dbias[idx] = (floatX)((float)dbias[idx] + sum); + } +} + + +__global__ void layernorm_backward_kernel(floatX* dinp, float* dweight, float* dbias, const float* dout, const float* inp, const float* weight, const float* mean, const float* rstd, int B, int T, int C) { cg::thread_block block = cg::this_thread_block(); @@ -640,7 +680,7 @@ __global__ void layernorm_backward_kernel(float* dinp, float* dweight, float* db const float* dout_bt = dout + b * T * C + t * C; const float* inp_bt = inp + b * T * C + t * C; - float* dinp_bt = dinp + b * T * C + t * C; + floatX* dinp_bt = dinp + b * T * C + t * C; float mean_bt = mean[b * T + t]; float rstd_bt = rstd[b * T + t]; @@ -674,10 +714,65 @@ __global__ void layernorm_backward_kernel(float* dinp, float* dweight, float* db dval -= dnorm_mean; // term 2 dval -= norm_bti * dnorm_norm_mean; // term 3 dval *= rstd_bt; // final scale - dinp_bt[i] += dval; + dinp_bt[i] = (floatX)((float)dinp_bt[i] + dval); } } +__global__ void layernorm_backward_kernel_mixedp(floatX* dinp, float* dweight, float* dbias, + const floatX* dout, const float* inp, const float* weight, const float* mean, const float* rstd, + int B, int T, int C) { + cg::thread_block block = cg::this_thread_block(); + cg::thread_block_tile<32> warp = cg::tiled_partition<32>(block); + int idx = blockIdx.x * warp.meta_group_size() + warp.meta_group_rank(); + int N = B * T; + if(idx >= N) { + return; + } + + int b = idx / T; + int t = idx % T; + + const floatX* dout_bt = dout + b * T * C + t * C; + const float* inp_bt = inp + b * T * C + t * C; + floatX* dinp_bt = dinp + b * T * C + t * C; + float mean_bt = mean[b * T + t]; + float rstd_bt = rstd[b * T + t]; + + // first: two reduce operations + float dnorm_mean = 0.0f; + float dnorm_norm_mean = 0.0f; + for (int i = warp.thread_rank(); i < C; i += warp.size()) { + float norm_bti = (inp_bt[i] - mean_bt) * rstd_bt; + float dnorm_i = weight[i] * (float)dout_bt[i]; + dnorm_mean += dnorm_i; + dnorm_norm_mean += dnorm_i * norm_bti; + } + + dnorm_mean = cg::reduce(warp, dnorm_mean, cg::plus{}); + dnorm_norm_mean = cg::reduce(warp, dnorm_norm_mean, cg::plus{}); + + dnorm_mean = dnorm_mean / C; + dnorm_norm_mean = dnorm_norm_mean / C; + + // now iterate again and accumulate all the gradients + for (int i = warp.thread_rank(); i < C; i += warp.size()) { + float norm_bti = (inp_bt[i] - mean_bt) * rstd_bt; + float dnorm_i = weight[i] * (float)dout_bt[i]; + // gradient contribution to bias + atomicAdd(&dbias[i], (float)dout_bt[i]); + // gradient contribution to weight + atomicAdd(&dweight[i], norm_bti * (float)dout_bt[i]); + // gradient contribution to input + float dval = 0.0f; + dval += dnorm_i; // term 1 + dval -= dnorm_mean; // term 2 + dval -= norm_bti * dnorm_norm_mean; // term 3 + dval *= rstd_bt; // final scale + dinp_bt[i] = (floatX)((float)dinp_bt[i] + dval); + } +} + + __global__ void softmax_autoregressive_backward_kernel(float* dpreatt, const float* datt, const float* att, int B, int T, int C, float scale) { constexpr const int BlockSize = 256; @@ -855,7 +950,7 @@ void encoder_forward(float* out, } void encoder_backward(float* dwte, float* dwpe, - const float* dout, const int* inp, + const floatX* dout, const int* inp, int B, int T, int C) { const int N = B * T * C; const int block_size = 256; @@ -898,7 +993,7 @@ void matmul_forward_cublas(float* out, // https://docs.nvidia.com/cuda/cublas/#cublasltmatmul // https://github.com/NVIDIA/CUDALibrarySamples/blob/master/cuBLASLt/LtSgemm/sample_cublasLt_LtSgemm.cu void matmul_forward_cublaslt(floatX* out, - floatX* inp, floatX* weight, float* bias, + floatX* inp, floatX* weight, floatX* bias, int B, int T, int C, int OC) { int has_bias = (bias != NULL); @@ -941,7 +1036,7 @@ void matmul_forward_cublaslt(floatX* out, cublasCheck(cublasLtMatrixLayoutCreate(&weightLayout, CUBLAS_LOWP, C, OC, C)); cublasCheck(cublasLtMatrixLayoutCreate(&inputLayout, CUBLAS_LOWP, C, B*T, C)); cublasCheck(cublasLtMatrixLayoutCreate(&outputLayout, CUBLAS_LOWP, OC, B*T, OC)); - cublasCheck(cublasLtMatrixLayoutCreate(&biasLayout, CUDA_R_32F, OC, 1, OC)); + cublasCheck(cublasLtMatrixLayoutCreate(&biasLayout, CUBLAS_LOWP, OC, 1, OC)); // create a preference handle with specified max workspace cublasCheck(cublasLtMatmulPreferenceCreate(&preference)); @@ -1041,7 +1136,7 @@ void gelu_forward(floatX* out, const floatX* inp, int N) { cudaCheck(cudaGetLastError()); } -void gelu_backward(float* dinp, const float* inp, const float* dout, const int N) { +void gelu_backward(floatX* dinp, const floatX* inp, const floatX* dout, const int N) { const int block_size = 128; const int grid_size = CEIL_DIV(N, block_size); gelu_backward_kernel<<>>(dinp, inp, dout, N); @@ -1074,7 +1169,32 @@ void matmul_backward(float* dinp, float* dweight, float* dbias, } } -void layernorm_backward(float* dinp, float* dweight, float* dbias, +void matmul_backward_lowp(floatX* dinp, floatX* dweight, floatX* dbias, + floatX* dout, floatX* inp, floatX* weight, + int B, int T, int C, int OC) { + float one = 1.0f; + float zero = 0.0f; + // backward to input, uses = in the backward pass (set the gradient) + // TODO + + //cublasCheck(cublasSgemm(cublas_handle, CUBLAS_OP_N, CUBLAS_OP_N, C, B*T, OC, &one, weight, C, dout, OC, &zero, dinp, C)); + + cublasCheck(cublasGemmEx(cublas_handle, CUBLAS_OP_N, CUBLAS_OP_N, C, B*T, OC, &one, weight, CUBLAS_LOWP, C, dout, CUBLAS_LOWP, OC, &zero, dinp, CUBLAS_LOWP, C, CUBLAS_LOWP_COMPUTE, CUBLAS_GEMM_DEFAULT_TENSOR_OP)); + + // backward to weight, uses += in the backward pass (accumulate the gradient) + cublasCheck(cublasGemmEx(cublas_handle, CUBLAS_OP_N, CUBLAS_OP_T, C, OC, B*T, &one, inp, CUBLAS_LOWP, C, dout, CUBLAS_LOWP, OC, &one, dweight, CUBLAS_LOWP, C, CUBLAS_LOWP_COMPUTE, CUBLAS_GEMM_DEFAULT_TENSOR_OP)); + + //cublasCheck(cublasSgemm(cublas_handle, CUBLAS_OP_N, CUBLAS_OP_T, C, OC, B*T, &one, inp, C, dout, OC, &one, dweight, C)); + // backward to bias, if given, does a += + if (dbias != NULL) { + const int block_size = 512; + const int grid_size = CEIL_DIV(OC * 32, block_size); + matmul_backward_bias_kernel3_lowp<<>>(dbias, dout, B, T, OC); + cudaCheck(cudaGetLastError()); + } +} + +void layernorm_backward(floatX* dinp, float* dweight, float* dbias, const float* dout, const float* inp, const float* weight, const float* mean, const float* rstd, int B, int T, int C) { const int block_size = 256; @@ -1085,10 +1205,22 @@ void layernorm_backward(float* dinp, float* dweight, float* dbias, cudaCheck(cudaGetLastError()); } +void layernorm_backward_mixedp(floatX* dinp, float* dweight, float* dbias, + const floatX* dout, const float* inp, const float* weight, const float* mean, const float* rstd, + int B, int T, int C) { + const int block_size = 256; + const int N = B * T; + // one warp per token, so we need to divide by 32 here. + const int grid_size = CEIL_DIV(N, block_size / 32); + layernorm_backward_kernel_mixedp<<>>(dinp, dweight, dbias, dout, inp, weight, mean, rstd, B, T, C); + cudaCheck(cudaGetLastError()); +} + + // the sequence of transformations in this compound op is: // 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* scratch, - const float* dout, +void attention_backward(floatX* dinp, float* dqkvr, float* dpreatt, float* datt, float* scratch, + const floatX* dout, const float* qkvr, const float* att, int B, int T, int C, int NH) { const int block_size = 256; @@ -1157,15 +1289,15 @@ typedef struct { float* ln1w; // (L, C) float* ln1b; // (L, C) floatX* qkvw; // (L, 3*C, C) - float* qkvb; // (L, 3*C) + floatX* qkvb; // (L, 3*C) floatX* attprojw; // (L, C, C) - float* attprojb; // (L, C) + floatX* attprojb; // (L, C) float* ln2w; // (L, C) float* ln2b; // (L, C) floatX* fcw; // (L, 4*C, C) - float* fcb; // (L, 4*C) + floatX* fcb; // (L, 4*C) floatX* fcprojw; // (L, C, 4*C) - float* fcprojb; // (L, C) + floatX* fcprojb; // (L, C) float* lnfw; // (C) float* lnfb; // (C) } ParameterTensors; @@ -1400,7 +1532,7 @@ void gpt2_build_from_checkpoint(GPT2 *model, const char* checkpoint_path) { // Setup param_sizeof for FP32 vs FP16 parameters etc. for (size_t i = 0; i < NUM_PARAMETER_TENSORS; i++) { // qkvw, attprojw, fcw, fcprojw are floatX - if(i == 4 || i == 6 || i == 10 || i == 12) { + if((i >= 4 && i <= 7) || (i >= 10 && i <= 13)) { model->param_sizeof[i] = sizeof(floatX); } else { model->param_sizeof[i] = sizeof(float); @@ -1495,7 +1627,8 @@ void gpt2_forward(GPT2 *model, int* inputs, int* targets, int B, int T) { // Setup param_sizeof for FP32 vs FP16 parameters etc. for (size_t i = 0; i < NUM_ACTIVATION_TENSORS; i++) { // ln1/ln2/atty/attproj/fch/fch_gelu/fcproj/qkvr/scratch - if(i == 1 || i == 4 || i == 6 || i == 8 || i == 11 || i == 12 || i == 13 /*|| i == 19 || i == 20*/) { + // not atty(4) and fch(11) for now as reused as scratch, TODO + if(i == 1 || /*i == 4 ||*/ i == 6 || i == 8 || /*i == 11 ||*/ i == 12 || i == 13 /*|| i == 19 || i == 20*/) { model->act_sizeof[i] = sizeof(floatX); } else { model->act_sizeof[i] = sizeof(float); @@ -1543,15 +1676,15 @@ void gpt2_forward(GPT2 *model, int* inputs, int* targets, int B, int T) { float* l_ln1w = params.ln1w + l * C; float* l_ln1b = params.ln1b + l * C; floatX* l_qkvw = params.qkvw + l * 3*C * C; - float* l_qkvb = params.qkvb + l * 3*C; + floatX* l_qkvb = params.qkvb + l * 3*C; floatX* l_attprojw = params.attprojw + l * C * C; - float* l_attprojb = params.attprojb + l * C; + floatX* l_attprojb = params.attprojb + l * C; float* l_ln2w = params.ln2w + l * C; float* l_ln2b = params.ln2b + l * C; floatX* l_fcw = params.fcw + l * 4*C * C; - float* l_fcb = params.fcb + l * 4*C; + floatX* l_fcb = params.fcb + l * 4*C; floatX* l_fcprojw = params.fcprojw + l * C * 4*C; - float* l_fcprojb = params.fcprojb + l * C; + floatX* l_fcprojb = params.fcprojb + l * C; // get the pointers of the activations for this layer floatX* l_ln1 = acts.ln1 + l * B * T * C; @@ -1615,9 +1748,6 @@ void gpt2_zero_grad(GPT2 *model) { } void gpt2_backward(GPT2 *model) { - // TODO - /* - // double check we forwarded previously, with targets if (model->mean_loss == -1.0f) { printf("Error: must forward with targets before backward\n"); @@ -1668,7 +1798,7 @@ void gpt2_backward(GPT2 *model) { matmul_backward(grads_acts.bt4c, grads.wte, NULL, acts.output, 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 + floatX* dresidual = (floatX*)grads_acts.residual3; // the main buffer holding the gradient in the backward pass layernorm_backward(dresidual, grads.lnfw, grads.lnfb, grads_acts.bt4c, residual, params.lnfw, acts.lnf_mean, acts.lnf_rstd, B, T, C); // now backward all the layers @@ -1677,68 +1807,67 @@ void gpt2_backward(GPT2 *model) { // get the pointers of the weights for this layer float* l_ln1w = params.ln1w + l * C; - float* l_qkvw = params.qkvw + l * 3*C * C; - float* l_attprojw = params.attprojw + l * C * C; + floatX* l_qkvw = params.qkvw + l * 3*C * C; + floatX* l_attprojw = params.attprojw + l * C * C; float* l_ln2w = params.ln2w + l * C; - float* l_fcw = params.fcw + l * 4*C * C; - float* l_fcprojw = params.fcprojw + l * C * 4*C; + floatX* l_fcw = params.fcw + l * 4*C * C; + floatX* l_fcprojw = params.fcprojw + l * C * 4*C; // get the pointers of the gradients of the weights for this layer float* dl_ln1w = grads.ln1w + l * C; float* dl_ln1b = grads.ln1b + l * C; - float* dl_qkvw = grads.qkvw + l * 3*C * C; - float* dl_qkvb = grads.qkvb + l * 3*C; - float* dl_attprojw = grads.attprojw + l * C * C; - float* dl_attprojb = grads.attprojb + l * C; + floatX* dl_qkvw = grads.qkvw + l * 3*C * C; + floatX* dl_qkvb = grads.qkvb + l * 3*C; + floatX* dl_attprojw = grads.attprojw + l * C * C; + floatX* dl_attprojb = grads.attprojb + l * C; float* dl_ln2w = grads.ln2w + l * C; float* dl_ln2b = grads.ln2b + l * C; - float* dl_fcw = grads.fcw + l * 4*C * C; - float* dl_fcb = grads.fcb + l * 4*C; - float* dl_fcprojw = grads.fcprojw + l * C * 4*C; - float* dl_fcprojb = grads.fcprojb + l * C; + floatX* dl_fcw = grads.fcw + l * 4*C * C; + floatX* dl_fcb = grads.fcb + l * 4*C; + floatX* dl_fcprojw = grads.fcprojw + l * C * 4*C; + floatX* dl_fcprojb = grads.fcprojb + l * C; // get the pointers of the activations for this layer - float* l_ln1 = acts.ln1 + l * B * T * C; + floatX* l_ln1 = acts.ln1 + l * B * T * C; float* l_ln1_mean = acts.ln1_mean + l * B * T; float* l_ln1_rstd = acts.ln1_rstd + l * B * T; float* l_qkvr = acts.qkvr + l * B * T * 3*C; - float* l_atty = acts.atty + l * B * T * C; + floatX* l_atty = acts.atty + l * B * T * C; float* l_att = acts.att + l * B * NH * T * T; float* l_residual2 = acts.residual2 + l * B * T * C; - float* l_ln2 = acts.ln2 + l * B * T * C; + floatX* l_ln2 = acts.ln2 + l * B * T * C; float* l_ln2_mean = acts.ln2_mean + l * B * T; float* l_ln2_rstd = acts.ln2_rstd + l * B * T; - float* l_fch = acts.fch + l * B * T * 4*C; - float* l_fch_gelu = acts.fch_gelu + l * B * T * 4*C; + floatX* l_fch = acts.fch + l * B * T * 4*C; + floatX* l_fch_gelu = acts.fch_gelu + l * B * T * 4*C; // get the pointers of the gradients of the activations for this layer // notice that there is no l *, because we just have a single copy, and keep // re-using this memory in every Transformer block as we calculate backward pass // we need a B x T x C buffer; thankfully, the forward activation for lnf isn't needed anymore, // so we can co-opt it here. - float* dl_btc = acts.lnf; - float* dl_bt4c = grads_acts.bt4c; + floatX* dl_btc = (floatX*)acts.lnf; + floatX* dl_bt4c = (floatX*)grads_acts.bt4c; float* dl_preatt = grads_acts.preatt; // re-use scratch buffer of the forward pass float* scratch = acts.output; // backprop this layer - matmul_backward(dl_bt4c, dl_fcprojw, dl_fcprojb, dresidual, l_fch_gelu, l_fcprojw, B, T, 4*C, C); + matmul_backward_lowp(dl_bt4c, dl_fcprojw, dl_fcprojb, dresidual, l_fch_gelu, l_fcprojw, B, T, 4*C, C); gelu_backward(dl_bt4c, l_fch, dl_bt4c, B*T*4*C); - matmul_backward(dl_btc, dl_fcw, dl_fcb, dl_bt4c, l_ln2, l_fcw, B, T, C, 4 * C); + matmul_backward_lowp(dl_btc, dl_fcw, dl_fcb, dl_bt4c, l_ln2, l_fcw, B, T, C, 4 * C); // layernorm backward does += to the dresidual, so it correctly accumulates grad from the MLP block above - layernorm_backward(dresidual, dl_ln2w, dl_ln2b, dl_btc, l_residual2, l_ln2w, l_ln2_mean, l_ln2_rstd, B, T, C); - matmul_backward(dl_btc, dl_attprojw, dl_attprojb, dresidual, l_atty, l_attprojw, B, T, C, C); + layernorm_backward_mixedp(dresidual, dl_ln2w, dl_ln2b, dl_btc, l_residual2, l_ln2w, l_ln2_mean, l_ln2_rstd, B, T, C); + matmul_backward_lowp(dl_btc, dl_attprojw, dl_attprojb, dresidual, l_atty, l_attprojw, B, T, C, C); // we more B x T x (4)C buffers. l_atty and l_fch aren't needed anymore at this point, so reuse their memory - float* buffer_a = l_atty; - float* buffer_b = l_fch; // this is B x T x 4C, so even larger than what we need + float* buffer_a = (float*)l_atty; + float* buffer_b = (float*)l_fch; // this is B x T x 4C, so even larger than what we need attention_backward(dl_bt4c, buffer_b, dl_preatt, scratch, buffer_a, dl_btc, l_qkvr, l_att, B, T, C, NH); - matmul_backward(dl_btc, dl_qkvw, dl_qkvb, dl_bt4c, l_ln1, l_qkvw, B, T, C, 3 * C); + matmul_backward_lowp(dl_btc, dl_qkvw, dl_qkvb, dl_bt4c, 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_btc, residual, l_ln1w, l_ln1_mean, l_ln1_rstd, B, T, C); + layernorm_backward_mixedp(dresidual, dl_ln1w, dl_ln1b, dl_btc, residual, l_ln1w, l_ln1_mean, l_ln1_rstd, B, T, C); } encoder_backward(grads.wte, grads.wpe, dresidual, model->inputs, B, T, C); - */ } void gpt2_update(GPT2 *model, float learning_rate, float beta1, float beta2, float eps, float weight_decay, int t) {