From 841c4dec03ab1ba13d08eb9eb3f4ee214b0ee855 Mon Sep 17 00:00:00 2001 From: Erik Schultheis Date: Mon, 27 May 2024 03:58:45 +0300 Subject: [PATCH] update main training file --- train_gpt2.cu | 211 +++++++++++++++++++++++++++++--------------------- 1 file changed, 122 insertions(+), 89 deletions(-) diff --git a/train_gpt2.cu b/train_gpt2.cu index 3f62468..c839dad 100644 --- a/train_gpt2.cu +++ b/train_gpt2.cu @@ -248,6 +248,21 @@ struct alignas(16) Packed128 { static_assert(sizeof(bits) == sizeof(payload), "Size mismatch."); memcpy(&payload, &bits, sizeof(bits)); } + + __device__ static Packed128 constant(ElementType value) { + Packed128 result; + for(int k = 0; k < size; ++k) { + result.payload[k] = value; + } + return result; + } + __device__ static Packed128 zeros() { + return constant(0); + } + __device__ static Packed128 ones() { + return constant(1); + } + __device__ ElementType& operator[](int index) { return payload[index]; } @@ -1050,42 +1065,42 @@ __global__ void reduce_add_sum_kernel(floatX* dst, const float* src, size_t n, s } __global__ void __launch_bounds__(512, 2) // todo - any warnings on Turing with only 1024 threads? - layernorm_backward_kernel9(floatX* dinp, floatX* dweight, floatX* dbias, float* scratch, + layernorm_backward_kernel10(floatX* dinp, floatX* dweight, floatX* dbias, float* scratch, const floatX* dout, const floatX* inp, const floatX* weight, const floatX* mean, const floatX* rstd, int B, int T, int C) { - extern __shared__ float shared[]; // size = 2*C + 2*block_size + 1 - int warpsInBlock = blockDim.x / WARP_SIZE; //number of warps in block + int BLOCK_SIZE = blockDim.x; + int warpsInBlock = BLOCK_SIZE / WARP_SIZE; //number of warps in block + extern __shared__ float shared[]; + int warpId = threadIdx.x / WARP_SIZE; // warp index within a block int baseIdx = blockIdx.x * warpsInBlock + warpId; int warpThreadIdx = threadIdx.x % WARP_SIZE; // Thread index within the warp int warpsInGrid = gridDim.x * warpsInBlock; int C_per_iteration = WARP_SIZE * x128::size; - int iterations_C = CEIL_DIV(C, C_per_iteration); + int iterations_C = CEIL_DIV(C, C_per_iteration); // + 2; // the first half of shared memory is bias, second is weight + size_t rounded_C = CEIL_DIV(C, (32 * x128::size)) * (32 * x128::size); float* dbias_shared = shared; - float* dweight_shared = shared + C; - float* dbias_tmp_shared = shared + 2 * C; - float* dweight_tmp_shared = shared + 2 * C + blockDim.x; + float* dweight_shared = shared + rounded_C; + // warp zero doesn't actually write to the _tmp_shared memory locations, so we don't need to reserve memory + // the obvious solution is to change the addressing below to use (threadId.x-32) as offset, but that causes + // register spills, so instead we mess with the base pointer here, which doesn't increase register usage. + float* dbias_tmp_shared = shared + 2 * rounded_C - WARP_SIZE * f128::size; + float* dweight_tmp_shared = shared + 2 * rounded_C + f128::size * BLOCK_SIZE - 2 * WARP_SIZE * f128::size; // init shared memory to zero - for(int i = threadIdx.x; i < C; i+= blockDim.x){ - dbias_shared[i] = 0.0f; - dweight_shared[i] = 0.0f; + for(int i = threadIdx.x * f128::size; i < rounded_C; i += BLOCK_SIZE * f128::size) { + store128(dbias_shared + i, f128::zeros()); + store128(dweight_shared + i, f128::zeros()); } - unsigned int *tmp_flag = (unsigned int*)(shared + 2*C + 2*blockDim.x); __syncthreads(); - for (int idx = baseIdx; idx < B * T; idx += warpsInGrid) { - int b = idx / T; - int t = idx % T; - - const floatX* dout_bt = dout + b * T * C + t * C; - const floatX* inp_bt = inp + b * T * C + t * C; - floatX* dinp_bt = dinp + b * T * C + t * C; - const float mean_bt = (float)mean[b * T + t]; - const float rstd_bt = (float)rstd[b * T + t]; + for (int bt = baseIdx; bt < B * T; bt += warpsInGrid) { + const floatX* dout_bt = dout + bt * C; + const floatX* inp_bt = inp +bt * C; + floatX* dinp_bt = dinp + bt * C; // first: two reduce operations float dnorm_mean = 0.0f; @@ -1095,69 +1110,85 @@ __global__ void __launch_bounds__(512, 2) // todo - any warnings on Turing with x128 inp128_i = load128(inp_bt + i); x128 weight128_i = load128(weight + i); for (int k = 0; k < x128::size; k++) { - float norm_bti = ((float)inp128_i[k] - mean_bt) * rstd_bt; float dnorm_i = (float)weight128_i[k] * (float)dout128_i[k]; dnorm_mean += dnorm_i; - dnorm_norm_mean += dnorm_i * norm_bti; + dnorm_norm_mean += dnorm_i * (float)inp128_i[k]; } } - dnorm_mean = warpReduceSum(dnorm_mean) / C; - dnorm_norm_mean = warpReduceSum(dnorm_norm_mean) / C; - // now iterate again and accumulate all the gradients - // unfortunately we cannot use the same index for x128 arrays and shared memory - // as atomics can only be 32-bit rather than 128-bit (at least pre-SM90/Hopper) - // so this would result in an 8-way bank conflict, and kill performance - // so instead, we use a shared memory friendly index, and reorder before the final write - for (int i = 0; i < iterations_C; i++) { - int global_index = (warpThreadIdx * x128::size) + (i * C_per_iteration); - int shared_index = warpThreadIdx + (i * C_per_iteration); - if (global_index >= C) { - break; + const float mean_bt = (float)mean[bt]; + const float rstd_bt = (float)rstd[bt]; + dnorm_mean = warpReduceSum(dnorm_mean) / C; + dnorm_norm_mean = warpReduceSum(dnorm_norm_mean) / C * rstd_bt - dnorm_mean * mean_bt; + + for (int c = 0; c < iterations_C; c++) { + int global_index = (warpThreadIdx * x128::size) + (c * C_per_iteration); + + x128 dout128 = x128::zeros(); + x128 inp128 = x128::zeros(); + x128 dinp128 = x128::zeros(); + x128 weight128 = x128::zeros(); + + if(global_index < C) { + dout128 = load128cs(dout_bt + global_index); + inp128 = load128cs(inp_bt + global_index); + dinp128 = load128(dinp_bt + global_index); + weight128 = load128(weight + global_index); } - x128 dout128 = load128cs(dout_bt + global_index); - x128 inp128 = load128cs(inp_bt + global_index); - x128 dinp128 = load128(dinp_bt + global_index); - x128 weight128 = load128(weight + global_index); + for(int o = 0; o < x128::size / f128::size; ++o) { + f128 dbias_f; + f128 dweight_f; + for(int i = 0; i < f128::size; ++i) { + int x = o * f128::size + i; + float dout_i = (float)dout128[x]; + float norm_bti = ((float)inp128[x] - mean_bt) * rstd_bt; + dbias_f[i] = dout_i; + dweight_f[i] = norm_bti * dout_i; - for (int x = 0; x < x128::size; x++) { - float dout_i = (float)dout128[x]; - float norm_bti = ((float)inp128[x] - mean_bt) * rstd_bt; - float dnorm_i = (float)weight128[x] * dout_i; + float dval = 0.0f; + dval += (float) weight128[x] * (float)dout128[x]; // term 1 + dval -= dnorm_mean; // term 2 + dval -= norm_bti * dnorm_norm_mean; // term 3 + dval *= rstd_bt; // final scale + dinp128[x] = (floatX) ((float) dinp128[x] + dval); + } - // sum up the gradients for bias and weight across the entire block - // this is basically a reduction (but only inter-warp, not intra-warp) - // doing it this way allows us to avoid using atomics while using many warps if (warpId != 0) { - dbias_tmp_shared[threadIdx.x] = dout_i; - dweight_tmp_shared[threadIdx.x] = norm_bti * dout_i; + store128(dbias_tmp_shared + threadIdx.x * f128::size, dbias_f); + // this seems to generate a 64-bit store, instead of 128-bit. + // however, forcing 128-bit (e.g., using inline ptx), results in register + // spilling and much worse performance, so we'll keep it like this for now + // but ideally, we could reduce the register pressure a little. + store128(dweight_tmp_shared + threadIdx.x * f128::size, dweight_f); } __syncthreads(); if (warpId == 0) { - float dbias_tmp = dout_i; - float dweight_tmp = norm_bti * dout_i; for (int j = 1; j < warpsInBlock; j++) { - dbias_tmp += dbias_tmp_shared[threadIdx.x + j * WARP_SIZE]; - dweight_tmp += dweight_tmp_shared[threadIdx.x + j * WARP_SIZE]; + f128 dbias_tmp = load128(dbias_tmp_shared + f128::size * (threadIdx.x + j * WARP_SIZE)); + f128 dweight_tmp = load128(dweight_tmp_shared + f128::size * (threadIdx.x + j * WARP_SIZE)); + for(int i = 0; i < f128::size; ++i) { + dbias_f[i] += dbias_tmp[i]; + dweight_f[i] += dweight_tmp[i]; + } } - // gradient contribution to bias (using shared memory friendly index) - dbias_shared[shared_index + x*WARP_SIZE] += dbias_tmp; - // gradient contribution to weight (using shared memory friendly index) - dweight_shared[shared_index + x*WARP_SIZE] += dweight_tmp; } __syncthreads(); - - // 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 - dinp128[x] = (floatX)((float)dinp128[x] + dval); + if (warpId == 0) { + f128 db_old = load128(dbias_shared + global_index + f128::size * o); + f128 dw_old = load128(dweight_shared + global_index + f128::size * o); + for(int i = 0; i < f128::size; ++i) { + dbias_f[i] += db_old[i]; + dweight_f[i] += dw_old[i]; + } + store128(dbias_shared + global_index + f128::size * o, dbias_f); + store128(dweight_shared + global_index + f128::size * o, dweight_f); + } + } + if(global_index < C) { + // cache in L2 as this is read by the next kernel, but bypass L1 to minimise thrashing + store128cg(dinp_bt + global_index, dinp128); } - // cache in L2 as this is read by the next kernel, but bypass L1 to minimise thrashing - store128cg(dinp_bt + global_index, dinp128); } } __syncthreads(); @@ -1169,24 +1200,25 @@ __global__ void __launch_bounds__(512, 2) // todo - any warnings on Turing with scratch += 32; float* scratch_dbias = scratch; float* scratch_dweight = scratch + C; - for(int i = threadIdx.x; i < C; i+= blockDim.x) { + for(int i = threadIdx.x * f128::size; i < C; i += BLOCK_SIZE * f128::size) { // Write to global memory in the same "shared memory banking friendly" order - scratch_dbias[i + 2*C*blockIdx.x] = dbias_shared[i]; - scratch_dweight[i + 2*C*blockIdx.x] = dweight_shared[i]; + store128(scratch_dbias + i + 2*C*blockIdx.x, load128(dbias_shared + i)); + store128(scratch_dweight + i + 2*C*blockIdx.x, load128(dweight_shared + i)); } - - // todo - everything below could become a separate kernel for better performance with maybe less code - // not enough parallelism even inside that single SM... do we need another level of reduction?! __syncthreads(); + // that portion of shared memory is no longer used, so we can repurpose it for the scratch flag. + unsigned int *tmp_flag = (unsigned int*)(shared + 2*rounded_C); if (threadIdx.x == 0) { *tmp_flag = atomicInc(scratchFlag, gridDim.x); } __syncthreads(); if (*tmp_flag == gridDim.x-1) { // Reduction of the partial sums by the final block - for(int i = threadIdx.x * f128::size; i < C; i+= blockDim.x * f128::size) { - f128 dbias_accum(make_int4(0, 0, 0, 0)); - f128 dweight_accum(make_int4(0, 0, 0, 0)); + // todo - there isn't enough parallelism even inside that single SM... + // ==> so could maybe split into another kernel with YET ANOTHER level of reduction?! + for(int i = threadIdx.x * f128::size; i < C; i += BLOCK_SIZE * f128::size) { + f128 dbias_accum = f128::zeros(); + f128 dweight_accum = f128::zeros(); for (int read_block_idx = 0; read_block_idx < gridDim.x; read_block_idx++) { int offset = i + 2*C*read_block_idx; @@ -1202,24 +1234,25 @@ __global__ void __launch_bounds__(512, 2) // todo - any warnings on Turing with } __syncthreads(); - // reorder from atomic/shared memory-friendly index to real global memory index - // and convert from float/FP32 to floatX/BF16 for the final write - // this is separate also because it cannot use as many warps as the above (f128 vs x128) + // convert from float/FP32 to floatX/BF16 for the final write + // this is separate because it cannot use as many warps as the above (f128 vs x128) // todo - if we split this code into another kernel, we could maybe do it at the same time? - for (int i = warpId; i < iterations_C; i += warpsInBlock) { - int global_index = (warpThreadIdx * x128::size) + (i * C_per_iteration); - int shared_index = warpThreadIdx + (i * C_per_iteration); + for (int c = warpId; c < iterations_C; c += warpsInBlock) { + int global_index = (warpThreadIdx * x128::size) + (c * C_per_iteration); if (global_index >= C) { break; } x128 dbias128 = load128(dbias + global_index); x128 dweight128 = load128(dweight + global_index); - for (int x = 0; x < x128::size; x++) { - float s_db = dbias_shared[shared_index + x*WARP_SIZE]; - float s_dw = dweight_shared[shared_index + x*WARP_SIZE]; - dbias128[x] = (floatX)(s_db + (float)dbias128[x]); - dweight128[x] = (floatX)(s_dw + (float)dweight128[x]); + for(int o = 0; o < x128::size / f128::size; ++o) { + f128 s_db = load128(dbias_shared + global_index + o * f128::size); + f128 s_dw = load128(dweight_shared + global_index + o * f128::size); + for(int i = 0; i < f128::size; ++i) { + int x = o * f128::size + i; + dbias128[x] = (floatX)(s_db[i] + (float)dbias128[x]); + dweight128[x] = (floatX)(s_dw[i] + (float)dweight128[x]); + } } store128(dbias + global_index, dbias128); store128(dweight + global_index, dweight128); @@ -1788,14 +1821,14 @@ void layernorm_backward(floatX* dinp, floatX* dweight, floatX* dbias, float* scr const int block_size = 512; const int blocks_per_sm = 2; // supported on every architecture and less cache thrashing than 3 const int grid_size = blocks_per_sm * deviceProp.multiProcessorCount; - size_t shared_mem_size = (2*C + 2*block_size + 1) * sizeof(float); // see kernel + size_t rounded_C = CEIL_DIV(C, (32 * x128::size)) * (32 * x128::size); + size_t shared_mem_size = (2 * rounded_C + 2 * (block_size - 32) * f128::size) * sizeof(float); - cudaMemset(scratch, 0, 1 * sizeof(float)); // only need to reset the flag to 0 - layernorm_backward_kernel9<<>>(dinp, dweight, dbias, scratch, dout, inp, weight, mean, rstd, B, T, C); + cudaCheck(cudaMemset(scratch, 0, 1 * sizeof(float))); // only need to reset the flag to 0 + layernorm_backward_kernel10<<>>(dinp, dweight, dbias, scratch, 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(floatX* dinp, floatX* dqkvr, floatX* dpreatt, floatX* datt, floatX* scratch,