diff --git a/train_gpt2.cu b/train_gpt2.cu index f9cd4e2..94da02e 100644 --- a/train_gpt2.cu +++ b/train_gpt2.cu @@ -114,6 +114,9 @@ class NvtxRange { #define MAX_1024_THREADS_BLOCKS 1 #endif +// WarpSize is not a compile time constant, this allows the compiler to optimize +#define WARP_SIZE 32U + // cuBLAS workspace. Hardcoding to 32MiB but only Hopper needs 32, for others 4 is OK const size_t cublaslt_workspace_size = 32 * 1024 * 1024; void* cublaslt_workspace = NULL; @@ -203,10 +206,10 @@ template __device__ float blockReduce(float val, bool final_sync=false, float out_of_bounds=0.0f) { // two reductions of up to 1024 threads: // 1) inside warp (shuffle), 2) cross-warp (shared memory), 3) inside warp (shuffle) - __shared__ float shared_val[32]; - const int lane_id = threadIdx.x % 32; - const int warp_id = threadIdx.x / 32; - const int num_warps = blockDim.x / 32; + __shared__ float shared_val[WARP_SIZE]; + const int lane_id = threadIdx.x % WARP_SIZE; + const int warp_id = threadIdx.x / WARP_SIZE; + const int num_warps = blockDim.x / WARP_SIZE; float warp_val = warp_reduction(val); if (lane_id == 0) { shared_val[warp_id] = warp_val; } @@ -578,10 +581,9 @@ __global__ void encoder_backward_kernel(floatX* dwte, floatX* dwpe, __global__ void layernorm_forward_kernel3(floatX* __restrict__ out, floatX* __restrict__ mean, floatX* __restrict__ rstd, const floatX* __restrict__ inp, const floatX* __restrict__ weight, const floatX* __restrict__ bias, int N, int C) { - const int warp_size = 32; - int lane_id = threadIdx.x % warp_size; - int warp_id = threadIdx.x / warp_size; - int num_warps = blockDim.x / warp_size; + int lane_id = threadIdx.x % WARP_SIZE; + int warp_id = threadIdx.x / WARP_SIZE; + int num_warps = blockDim.x / WARP_SIZE; int idx = blockIdx.x * num_warps + warp_id; if(idx >= N) { return; } // guard @@ -591,7 +593,7 @@ __global__ void layernorm_forward_kernel3(floatX* __restrict__ out, floatX* __re // mean float sum = 0.0f; - for (int i = lane_id; i < C; i += warp_size) { + for (int i = lane_id; i < C; i += WARP_SIZE) { sum += (float)x[i]; } sum = warpReduceSum(sum); @@ -602,7 +604,7 @@ __global__ void layernorm_forward_kernel3(floatX* __restrict__ out, floatX* __re // rstd sum = 0.0f; - for (int i = lane_id; i < C; i += warp_size) { + for (int i = lane_id; i < C; i += WARP_SIZE) { float diff = (float)x[i] - m; sum += diff * diff; } @@ -614,7 +616,7 @@ __global__ void layernorm_forward_kernel3(floatX* __restrict__ out, floatX* __re // final normalization and scaling by weight/bias floatX* o = out + idx * C; - for (int c = lane_id; c < C; c += warp_size) { + for (int c = lane_id; c < C; c += WARP_SIZE) { // load and store using the .cs "streaming" hint to the compiler, // indicating that this data will not be reused soon, and can be streamed through the caches // this allows the threads to get more cache-hits for the (shared) weight and bias parameters @@ -627,8 +629,7 @@ __global__ void fused_residual_forward_kernel5(floatX* residual, floatX* normed, const floatX* inp1, const floatX* inp2, const floatX* weight, const floatX* bias, int N, int C) { - constexpr const int WarpSize = 32; - assert(blockDim.x == WarpSize); + assert(blockDim.x == WARP_SIZE); // load weights and biases into shared memory // do this before we allow any threads to exit! @@ -639,8 +640,8 @@ __global__ void fused_residual_forward_kernel5(floatX* residual, floatX* normed, x128* s_bias = reinterpret_cast(params) + (C / x128::size); x128* s_res = reinterpret_cast(params) + ((2 + threadIdx.y) * C / x128::size); - int sidx = (threadIdx.x + WarpSize * threadIdx.y) * x128::size; - for(int i = sidx; i < C; i += blockDim.y * WarpSize * x128::size) { + int sidx = (threadIdx.x + WARP_SIZE * threadIdx.y) * x128::size; + for(int i = sidx; i < C; i += blockDim.y * WARP_SIZE * x128::size) { s_weight[i/x128::size] = load128(weight + i); s_bias[i/x128::size] = load128(bias + i); } @@ -657,7 +658,7 @@ __global__ void fused_residual_forward_kernel5(floatX* residual, floatX* normed, const float eps = 1e-5f; float sum = 0.0f; - for(int c = threadIdx.x * x128::size; c < C; c += WarpSize * x128::size) { + for(int c = threadIdx.x * x128::size; c < C; c += WARP_SIZE * x128::size) { const x128 in1 = load128cs(inp1 + c); const x128 in2 = load128cs(inp2 + c); x128 out; @@ -673,7 +674,7 @@ __global__ void fused_residual_forward_kernel5(floatX* residual, floatX* normed, float m = sum / C; float v = 0.f; - for(int c = threadIdx.x * x128::size; c < C; c += WarpSize * x128::size) { + for(int c = threadIdx.x * x128::size; c < C; c += WARP_SIZE * x128::size) { const x128 res = s_res[c / x128::size]; for(int k = 0; k < x128::size; ++k) { v += ((float)res[k] - m) * ((float)res[k] - m); @@ -683,7 +684,7 @@ __global__ void fused_residual_forward_kernel5(floatX* residual, floatX* normed, v = warpReduceSum(v) / C; float s = rsqrtf(v + eps); - for(int c = threadIdx.x * x128::size; c < C; c += WarpSize * x128::size) { + for(int c = threadIdx.x * x128::size; c < C; c += WARP_SIZE * x128::size) { const x128 res = s_res[c / x128::size]; const x128 w = s_weight[c / x128::size]; const x128 b = s_bias[c / x128::size]; @@ -898,7 +899,7 @@ template __global__ void matmul_backward_bias_kernel9(OutFloat* dbias, const floatX* dout, int B, int T, int OC, std::bool_constant) { constexpr const int bdx = 4; - constexpr const int bdy = 32 / bdx; + constexpr const int bdy = WARP_SIZE / bdx; assert(blockDim.x == bdx); assert(blockDim.y == bdy); @@ -929,7 +930,7 @@ __global__ void matmul_backward_bias_kernel9(OutFloat* dbias, const floatX* dout } } - __shared__ float sub_results[x128::size][32][bdy]; + __shared__ float sub_results[x128::size][WARP_SIZE][bdy]; // reduce within-warp results for (int k = 0; k < x128::size; k++) { @@ -988,12 +989,12 @@ __global__ void __launch_bounds__(512, 3) // todo - any warnings on Turing with const floatX* mean, const floatX* rstd, int B, int T, int C) { extern __shared__ float shared[]; // size = 2 * C + 1 - int warpId = threadIdx.x / warpSize; // warp index within a block - int warpsInBlock = blockDim.x / warpSize; //number of warps in block + int warpId = threadIdx.x / WARP_SIZE; // warp index within a block + int warpsInBlock = blockDim.x / WARP_SIZE; //number of warps in block int baseIdx = blockIdx.x * warpsInBlock + warpId; - int warpThreadIdx = threadIdx.x % warpSize; // Thread index within the warp + int warpThreadIdx = threadIdx.x % WARP_SIZE; // Thread index within the warp int warpsInGrid = gridDim.x * warpsInBlock; - int C_per_iteration = warpSize * x128::size; + int C_per_iteration = WARP_SIZE * x128::size; int iterations_C = C / C_per_iteration; // the first half of shared memory is bias, second is weight @@ -1021,7 +1022,7 @@ __global__ void __launch_bounds__(512, 3) // todo - any warnings on Turing with // first: two reduce operations float dnorm_mean = 0.0f; float dnorm_norm_mean = 0.0f; - for (int i = warpThreadIdx * x128::size; i < C; i += warpSize * x128::size) { + for (int i = warpThreadIdx * x128::size; i < C; i += WARP_SIZE * x128::size) { x128 dout128_i = load128(dout_bt + i); x128 inp128_i = load128(inp_bt + i); x128 weight128_i = load128(weight + i); @@ -1053,9 +1054,9 @@ __global__ void __launch_bounds__(512, 3) // todo - any warnings on Turing with float norm_bti = ((float)inp128[x] - mean_bt) * rstd_bt; float dnorm_i = (float)weight128[x] * dout_i; // gradient contribution to bias (using shared memory friendly index) - atomicAdd(&dbias_shared[shared_index + x*warpSize], dout_i); + atomicAdd(&dbias_shared[shared_index + x*WARP_SIZE], dout_i); // gradient contribution to weight (using shared memory friendly index) - atomicAdd(&dweight_shared[shared_index + x*warpSize], norm_bti * dout_i); + atomicAdd(&dweight_shared[shared_index + x*WARP_SIZE], norm_bti * dout_i); // gradient contribution to input float dval = 0.0f; dval += dnorm_i; // term 1 @@ -1095,8 +1096,8 @@ __global__ void __launch_bounds__(512, 3) // todo - any warnings on Turing with x128 dbias128 = load128(dbias + global_index); x128 dweight128 = load128(dweight + global_index); for (int x = 0; x < x128::size; x++) { - float s_db = scratch_dbias[shared_index + x*warpSize]; - float s_dw = scratch_dweight[shared_index + x*warpSize]; + float s_db = scratch_dbias[shared_index + x*WARP_SIZE]; + float s_dw = scratch_dweight[shared_index + x*WARP_SIZE]; dbias128[x] = (floatX)(s_db + (float)dbias128[x]); dweight128[x] = (floatX)(s_dw + (float)dweight128[x]); } @@ -1373,7 +1374,7 @@ void layernorm_forward(floatX* out, floatX* mean, floatX* rstd, NVTX_RANGE_FN(); const int block_size = 512; const int N = B * T; - const int grid_size = CEIL_DIV(N * 32, block_size); + const int grid_size = CEIL_DIV(N * WARP_SIZE, block_size); layernorm_forward_kernel3<<>>(out, mean, rstd, inp, weight, bias, N, C); cudaCheck(cudaGetLastError()); } @@ -1518,7 +1519,7 @@ void fused_residual_forward5(floatX* residual, floatX* normed, floatX* mean, flo const floatX* weight, const floatX* bias, int N, int C) { const int block_size = 256; - int block_y = block_size / 32; + int block_y = block_size / WARP_SIZE; const int grid_size = CEIL_DIV(N, block_y); size_t smem = (2 + block_y) * C * sizeof(floatX); @@ -1528,7 +1529,7 @@ void fused_residual_forward5(floatX* residual, floatX* normed, floatX* mean, flo auto status = cudaFuncSetAttribute(fused_residual_forward_kernel5, cudaFuncAttributeMaxDynamicSharedMemorySize, smem); cudaGetLastError(); if(status == cudaSuccess) { - fused_residual_forward_kernel5<<>>(residual, normed, mean, rstd, inp1, inp2, + fused_residual_forward_kernel5<<>>(residual, normed, mean, rstd, inp1, inp2, weight, bias, N, C); } else { residual_forward(residual, inp1, inp2, N*C); @@ -1568,7 +1569,7 @@ void matmul_backward(floatX* dinp, floatX* dweight, floatX* dbias, const int block_size = deviceProp.maxThreadsPerMultiProcessor == 1536 ? 768 : 1024; - dim3 block_dim = {4, 8, (unsigned)block_size/32}; + dim3 block_dim = {4, 8, (unsigned)block_size/WARP_SIZE}; const int OC_per_warp = block_dim.y * x128::size; // 64 at BF16 const int grid_size_x = CEIL_DIV(OC, OC_per_warp); // e.g. 12 horizontal blocks for 768 OCs at BF16 const int grid_size_y = max(1, deviceProp.maxThreadsPerMultiProcessor * deviceProp.multiProcessorCount / (block_size * grid_size_x)); // full GPU!