diff --git a/test_gpt2.cu b/test_gpt2.cu index 32bea71..6871b22 100644 --- a/test_gpt2.cu +++ b/test_gpt2.cu @@ -107,6 +107,12 @@ int main(int argc, char *argv[]) { cublasCheck(cublasSetMathMode(cublas_handle, cublas_math_mode)); cudaCheck(cudaMalloc(&cublaslt_workspace, cublaslt_workspace_size)); + #ifdef ENABLE_CUDNN + checkCudnnErr(cudnnCreate(&cudnn_handle)); + cudaCheck(cudaMalloc(&cudnn_workspace, cudnn_workspace_size)); + printf("INIT CUDNN %d\n", cudnn_workspace_size); + #endif + // build the GPT-2 model from a checkpoint GPT2 model; gpt2_build_from_checkpoint(&model, "gpt2_124M_bf16.bin"); @@ -172,7 +178,7 @@ int main(int argc, char *argv[]) { float logit_accuracy_threshold = 1e-2f; float loss_diff_threshold = 0.05f; #if defined(ENABLE_BF16) || defined(ENABLE_F16) - logit_accuracy_threshold = 15.0f; + logit_accuracy_threshold = 25.0f; // 15.0f was too low even without cuDNN?! :( #endif // compare the output logits from the forward pass diff --git a/train_gpt2.cu b/train_gpt2.cu index 8451480..4ff30fd 100644 --- a/train_gpt2.cu +++ b/train_gpt2.cu @@ -25,6 +25,7 @@ Also we're using TinyStories here for example as it is a bigger dataset Example launch using bfloat16 on 4 GPUs, same as above: mpirun -np 4 ./train_gpt2cu -b 8 -v 200 -s 200 -i data/TinyStories */ +#define ENABLE_CUDNN // can be enabled via nvcc "-DENABLE_CUDNN" #include #include @@ -103,6 +104,23 @@ const ncclDataType_t ncclFloatX = ncclBfloat16; #endif #endif +#ifdef ENABLE_CUDNN +#include +namespace fe = cudnn_frontend; +#if CUBLAS_LOWP == CUDA_R_16BF +#define CUDNN_16BIT fe::DataType_t::BFLOAT16 +#else +#define CUDNN_16BIT fe::DataType_t::HALF +#endif + +static cudnnHandle_t cudnn_handle; +static size_t cudnn_workspace_size = 0; // dynamically allocated as needed (up to 256MiB!) +static void* cudnn_workspace = NULL; + +#define checkCudaErr(err) assert((int)err == 0); +#define checkCudnnErr(err) assert((int)err == 0); +#endif // ENABLE_CUDNN + // ---------------------------------------------------------------------------- // CUDA utils @@ -429,6 +447,261 @@ void printf0(const char *format, ...) { } } +// ---------------------------------------------------------------------------- +// cuDNN path +#ifdef ENABLE_CUDNN + +using graph_tensors_fwd = std::tuple, + std::shared_ptr, // Q, + std::shared_ptr, // K, + std::shared_ptr, // V, + std::shared_ptr, // Attn_scale, + std::shared_ptr, // O + std::shared_ptr>; // Stats + +using graph_tensors_bwd = std::tuple, + std::shared_ptr, // Q, + std::shared_ptr, // K, + std::shared_ptr, // V, + std::shared_ptr, // O + std::shared_ptr, // dO + std::shared_ptr, // Stats + std::shared_ptr, // Attn_scale, + std::shared_ptr, // dQ, + std::shared_ptr, // dK, + std::shared_ptr>; // dV + +// Need a cache because graph->build_operation_graph() is slow but everything else seems fast +using cache_type_fwd = std::unordered_map; +using cache_type_bwd = std::unordered_map; + +// Loosely based on cuDNN frontend samples functions and massively simplified +template +auto lookup_cache_or_build_graph_fwd(Args... args) { + static cache_type_fwd user_maintained_cache_fwd; + auto [B, H, T, HS, is_inference_only] = std::make_tuple(args...); + + auto graph = std::make_shared(); + graph->set_io_data_type(CUDNN_16BIT) + .set_intermediate_data_type(fe::DataType_t::FLOAT) + .set_compute_data_type(fe::DataType_t::FLOAT); + + // QKV is (B, T, 3, NH, HS) which cuDNN can handle directly without an external permute + auto Q = graph->tensor(fe::graph::Tensor_attributes() + .set_name("Q") + .set_dim({B, H, T, HS}) + .set_stride({3 * H * HS * T, HS, 3 * H * HS, 1})); + auto K = graph->tensor(fe::graph::Tensor_attributes() + .set_name("K") + .set_dim({B, H, T, HS}) + .set_stride({3 * H * HS * T, HS, 3 * H * HS, 1})); + auto V = graph->tensor(fe::graph::Tensor_attributes() + .set_name("V") + .set_dim({B, H, T, HS}) + .set_stride({3 * H * HS * T, HS, 3 * H * HS, 1})); + auto attn_scale = graph->tensor(fe::graph::Tensor_attributes() + .set_name("attn_scale") + .set_dim({1, 1, 1, 1}) + .set_stride({1, 1, 1, 1}) + .set_is_pass_by_value(true) + .set_data_type(fe::DataType_t::FLOAT)); + + auto sdpa_options = fe::graph::SDPA_attributes().set_name("flash_attention"); + sdpa_options.set_is_inference(is_inference_only); + sdpa_options.set_attn_scale(attn_scale); + sdpa_options.set_causal_mask(true); + + // Create the graph operation and get the output tensors back + auto [O, stats] = graph->sdpa(Q, K, V, sdpa_options); + + // Output is (B, T, NH, HS) BF16/FP16 and stats for backward pass is (B, NH, T) FP32 + O->set_output(true).set_dim({B, H, T, HS}).set_stride({H * HS * T, HS, H * HS, 1}); + + assert(stats == nullptr || is_inference_only == false); + if (is_inference_only == false) { + stats->set_output(true).set_data_type(fe::DataType_t::FLOAT) + .set_dim({B, H, T, 1}) + .set_stride({H * T, T, 1, 1}); + } + + assert(graph->validate().is_good()); + auto key = graph->key(); + auto it = user_maintained_cache_fwd.find(key); + if (it != user_maintained_cache_fwd.end()) { + return it->second; + } + + // Build the operation graph and execution part (this is the VERY SLOW PART) + assert(graph->build_operation_graph(cudnn_handle).is_good()); + auto plans = graph->create_execution_plans({fe::HeurMode_t::A}); + assert(graph->check_support(cudnn_handle).is_good()); + assert(graph->build_plans(cudnn_handle).is_good()); + + auto tuple = std::make_tuple(graph, Q, K, V, attn_scale, O, stats); + user_maintained_cache_fwd.insert({key, tuple}); + return tuple; +} + +template +auto lookup_cache_or_build_graph_bwd(Args... args) { + static cache_type_bwd user_maintained_cache_bwd; + auto [B, NH, T, HS] = std::make_tuple(args...); + + auto graph = std::make_shared(); + graph->set_io_data_type(CUDNN_16BIT) + .set_intermediate_data_type(fe::DataType_t::FLOAT) + .set_compute_data_type(fe::DataType_t::FLOAT); + + // (B, N, 3, NH, HS) + // must come from inp (which means we also need to convert THAT to FP16) + auto Q = graph->tensor(fe::graph::Tensor_attributes() + .set_name("Q") + .set_dim({B, NH, T, HS}) + .set_stride({3 * NH * HS * T, HS, 3 * NH * HS, 1})); + auto K = graph->tensor(fe::graph::Tensor_attributes() + .set_name("K") + .set_dim({B, NH, T, HS}) + .set_stride({3 * NH * HS * T, HS, 3 * NH * HS, 1})); + auto V = graph->tensor(fe::graph::Tensor_attributes() + .set_name("V") + .set_dim({B, NH, T, HS}) + .set_stride({3 * NH * HS * T, HS, 3 * NH * HS, 1})); + auto O = graph->tensor(fe::graph::Tensor_attributes() + .set_name("O") + .set_dim({B, NH, T, HS}) + .set_stride({NH * HS * T, HS, NH * HS, 1})); + auto dO = graph->tensor(fe::graph::Tensor_attributes() + .set_name("dO") + .set_dim({B, NH, T, HS}) + .set_stride({NH * HS * T, HS, NH * HS, 1})); + + auto stats = graph->tensor(fe::graph::Tensor_attributes() + .set_name("stats") + .set_dim({B, NH, T, 1}) + .set_stride({NH * T, T, 1, 1}) + .set_data_type(fe::DataType_t::FLOAT)); + auto attn_scale = graph->tensor(fe::graph::Tensor_attributes() + .set_name("attn_scale") + .set_dim({1, 1, 1, 1}) + .set_stride({1, 1, 1, 1}) + .set_is_pass_by_value(true) + .set_data_type(fe::DataType_t::FLOAT)); + auto sdpa_backward_options = fe::graph::SDPA_backward_attributes() + .set_name("flash_attention_backward") + .set_causal_mask(true) + .set_attn_scale(attn_scale); + + // Create the graph operation and get the output tensors back + auto [dQ, dK, dV] = graph->sdpa_backward(Q, K, V, O, dO, stats, sdpa_backward_options); + + dQ->set_output(true).set_dim({B, NH, T, HS}).set_stride({3 * NH * HS * T, HS, 3 * NH * HS, 1}); + dK->set_output(true).set_dim({B, NH, T, HS}).set_stride({3 * NH * HS * T, HS, 3 * NH * HS, 1}); + dV->set_output(true).set_dim({B, NH, T, HS}).set_stride({3 * NH * HS * T, HS, 3 * NH * HS, 1}); + + assert(graph->validate().is_good()); + auto key = graph->key(); + auto it = user_maintained_cache_bwd.find(key); + if (it != user_maintained_cache_bwd.end()) { + return it->second; + } + + // Build the operation graph and execution part (this is the VERY SLOW PART) + assert(graph->build_operation_graph(cudnn_handle).is_good()); + auto plans = graph->create_execution_plans({fe::HeurMode_t::A}); + assert(graph->check_support(cudnn_handle).is_good()); + assert(graph->build_plans(cudnn_handle).is_good()); + + auto tuple = std::make_tuple(graph, Q, K, V, O, dO, stats, attn_scale, dQ, dK, dV); + user_maintained_cache_bwd.insert({key, tuple}); + return tuple; +} + +void attention_forward_cudnn(floatX* out, // output: (B, T, NH, HS) + float* stats, // output for backward pass: (B, NH, T) + floatX* inp, // input: (B, T, 3, NH, HS) QKV + int B, int T, int NH, int C) { + int HS = C / NH; // number of features per head + bool is_inference_only = (stats == nullptr); + + // Get graph and tensors from cache (or generate it on first use) + auto [graph, Q, K, V, attn_scale, O, softmax_stats] = + lookup_cache_or_build_graph_fwd(B, NH, T, HS, is_inference_only); + + // Prepare all the tensor pointers for executing the graph + void* devPtrQ = inp; + void* devPtrK = (inp + C); + void* devPtrV = (inp + 2 * C); + float attn_scale_cpu = 1.0 / sqrtf(HS); + void* devPtrO = out; + + // Build variant pack + std::unordered_map, void*> variant_pack = { + {Q, devPtrQ}, {K, devPtrK}, {V, devPtrV}, {attn_scale, &attn_scale_cpu}, {O, devPtrO}}; + + // Add the stats tensor unless we are only doing inference (only needed for backward pass) + if (is_inference_only == false) { + variant_pack[softmax_stats] = stats; + } + + // Reallocate the workspace if the required size is greater than the current workspace + // By default, cuDNN uses up to 256MiB of workspace, so we don't want to just allocate the maximum + if (graph->get_workspace_size() > cudnn_workspace_size) { + if (cudnn_workspace_size > 0) { + cudaCheck(cudaFree(cudnn_workspace)); + } + cudnn_workspace_size = graph->get_workspace_size(); + cudaCheck(cudaMalloc(&cudnn_workspace, cudnn_workspace_size)); + } + + // Execute graph + assert(graph->execute(cudnn_handle, variant_pack, cudnn_workspace).is_good()); + cudaCheck(cudaGetLastError()); +} + +void attention_backward_cudnn(floatX* dqkvr, // output + floatX* dout, floatX* qkvr, floatX* o, float* stats, // inputs + int B, int T, int NH, int C) { + int HS = C / NH; // number of features per head + + // Get graph and tensors from cache (or generate it on first use) + auto [graph, Q, K, V, O, dO, Stats, attn_scale, dQ, dK, dV] = + lookup_cache_or_build_graph_bwd(B, NH, T, HS); + + // Prepare all the tensor pointers for executing the graph + void* devPtrQ = qkvr; + void* devPtrK = (qkvr + NH * HS); + void* devPtrV = (qkvr + 2 * NH * HS); + void* devPtrO = o; + void* devPtrdO = dout; + void* devPtrStats = stats; + float attn_scale_cpu = 1.0 / sqrtf(HS); + + void* devPtrdQ = dqkvr; + void* devPtrdK = (dqkvr + NH * HS); + void* devPtrdV = (dqkvr + 2 * NH * HS); + + // Build variant pack that links each tensor to its data pointer + std::unordered_map, void*> variant_pack = { + {Q, devPtrQ}, {K, devPtrK}, {V, devPtrV}, {O, devPtrO}, {dO, devPtrdO}, {Stats, devPtrStats}, + {dQ, devPtrdQ}, {dK, devPtrdK}, {dV, devPtrdV}, + {attn_scale, &attn_scale_cpu}}; + + // Reallocate the workspace if the required size is greater than the current workspace + // By default, cuDNN uses up to 256MiB of workspace, so we don't want to just allocate the maximum + if (graph->get_workspace_size() > cudnn_workspace_size) { + if (cudnn_workspace_size > 0) { + cudaCheck(cudaFree(cudnn_workspace)); + } + cudnn_workspace_size = graph->get_workspace_size(); + cudaCheck(cudaMalloc(&cudnn_workspace, cudnn_workspace_size)); + } + + // Execute graph + assert(graph->execute(cudnn_handle, variant_pack, cudnn_workspace).is_good()); + cudaCheck(cudaGetLastError()); +} +#endif // ENABLE_CUDNN + // ---------------------------------------------------------------------------- // all the kernels @@ -1460,19 +1733,29 @@ void fill_in_activation_sizes(size_t* act_sizes, size_t B, size_t T, GPT2Config // Backward pass is conceptually quite different from forward, because we can discard // the activations of a layer as soon as we're done with it. This lets us aggressively // reuse memory, so that we need far fewer tensors for backward state. +#ifdef ENABLE_CUDNN +#define NUM_BACKWARD_TENSORS 2 +#else #define NUM_BACKWARD_TENSORS 3 +#endif + typedef struct { floatX* bt4c; // (B, T, 4*C) - floatX* preatt; // (B, NH, T, T) floatX* residual3; // (B, T, C) + #ifndef ENABLE_CUDNN + floatX* preatt; // (B, NH, T, T) + #endif } GradActTensors; void fill_in_grad_act_sizes(size_t* act_sizes, size_t B, size_t T, GPT2Config config) { - size_t NH = config.num_heads; size_t C = config.channels; act_sizes[0] = B * T * 4 * C; // bt4c - act_sizes[1] = B * NH * T * T; // preatt - act_sizes[2] = B * T * C; // residual3 + act_sizes[1] = B * T * C; // residual3 + + #ifndef ENABLE_CUDNN + size_t NH = config.num_heads; + act_sizes[2] = B * NH * T * T; // preatt + #endif } void* malloc_and_point(floatX** targets[], const size_t* act_sizes, size_t n) { @@ -1502,7 +1785,10 @@ void* malloc_and_point_activations(ActivationTensors* acts, const size_t* act_si void* malloc_and_point_backward(GradActTensors* acts, const size_t* act_sizes) { floatX** ptrs[] = { - &acts->bt4c, &acts->preatt, &acts->residual3 + &acts->bt4c, &acts->residual3, + #ifndef ENABLE_CUDNN + &acts->preatt, + #endif }; return malloc_and_point(ptrs, act_sizes, NUM_BACKWARD_TENSORS); } @@ -1708,14 +1994,21 @@ void gpt2_forward(GPT2 *model, int* inputs, int* targets, size_t B, size_t T) { floatX* l_fch_gelu = acts.fch_gelu + l * B * T * 4*C; floatX* l_fcproj = acts.fcproj + l * B * T * C; floatX* 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 - floatX* scratch = (floatX*)acts.output; // now do the forward pass layernorm_forward(l_ln1, l_ln1_mean, l_ln1_rstd, residual, l_ln1w, l_ln1b, B, T, C); + + #ifdef ENABLE_CUDNN + matmul_forward_cublaslt(l_qkvr, l_ln1, l_qkvw, l_qkvb, B, T, C, 3*C); + attention_forward_cudnn(l_atty, (float*)l_att, l_qkvr, B, T, NH, C); + #else + // these are only needed as scratchpads for the forward pass, but + // need not be stored for backward + floatX* scratch = (floatX*)acts.output; matmul_forward_cublaslt(scratch, l_ln1, l_qkvw, l_qkvb, B, T, C, 3*C); attention_forward(l_atty, l_qkvr, l_att, scratch, B, T, C, NH); + #endif + matmul_forward_cublaslt(l_attproj, l_atty, l_attprojw, l_attprojb, B, T, C, C); residual_forward(l_residual2, residual, l_attproj, B*T*C); layernorm_forward(l_ln2, l_ln2_mean, l_ln2_rstd, l_residual2, l_ln2w, l_ln2b, B, T, C); @@ -1794,9 +2087,8 @@ void gpt2_backward(GPT2 *model) { ActivationTensors acts = model->acts; GradActTensors grads_acts = model->grads_acts; - // re-use scratch buffer of the forward pass + // re-use the output buffer of the forward pass as a scratchpad during backward pass float* scratchF = (float*)acts.output; - floatX* scratchX = (floatX*)acts.output; // we kick off the chain rule by filling in dlosses with 1.0f/(B*T) // this was done in the fused classifier kernel as last step of forward pass @@ -1854,7 +2146,6 @@ void gpt2_backward(GPT2 *model) { // so we can co-opt it here. floatX* dl_btc = (floatX*)acts.lnf; floatX* dl_bt4c = (floatX*)grads_acts.bt4c; - floatX* dl_preatt = (floatX*)grads_acts.preatt; // backprop this layer matmul_backward(dl_bt4c, dl_fcprojw, dl_fcprojb, dresidual, l_fch_gelu, l_fcprojw, B, T, 4*C, C); @@ -1863,11 +2154,19 @@ 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, scratchF, 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); - // 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 + + #ifdef ENABLE_CUDNN + attention_backward_cudnn(dl_bt4c, dl_btc, l_qkvr, l_atty, (float*)l_att, B, T, NH, C); + #else + // we need B x T x (4)C buffers. l_atty and l_fch aren't needed anymore at this point, so reuse their memory floatX* buffer_a = l_atty; floatX* buffer_b = l_fch; // this is B x T x 4C, so even larger than what we need - + floatX* dl_preatt = (floatX*)grads_acts.preatt; // dedicated scratchpad allocation + floatX* scratchX = (floatX*)acts.output; attention_backward(dl_bt4c, buffer_b, dl_preatt, scratchX, buffer_a, dl_btc, l_qkvr, l_att, B, T, C, NH); + #endif + + // QKV parameter gradients matmul_backward(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, scratchF, dl_btc, residual, l_ln1w, l_ln1_mean, l_ln1_rstd, B, T, C); @@ -2155,6 +2454,10 @@ int main(int argc, char *argv[]) { cublasCheck(cublasSetMathMode(cublas_handle, cublas_math_mode)); if(cublas_compute_type); // unused in BF16 mode, avoid warning + #ifdef ENABLE_CUDNN + checkCudnnErr(cudnnCreate(&cudnn_handle)); + #endif + printf0("| device | %-50s |\n", deviceProp.name); printf0("| TF32 | %-50s |\n", enable_tf32 ? "enabled" : "disabled"); printf0("| precision | %-50s |\n", precision_mode_str); @@ -2294,14 +2597,17 @@ int main(int argc, char *argv[]) { cudaCheck(cudaDeviceSynchronize()); // finish all CUDA work to get correct precise timings clock_gettime(CLOCK_MONOTONIC, &end); double time_elapsed_s = (end.tv_sec - start.tv_sec) + (end.tv_nsec - start.tv_nsec) / 1e9; - total_sum_iteration_time_s += time_elapsed_s; + + if (step > 0) { // consider the first batch to be a warmup (e.g. cuBLAS/cuDNN initialisation) + total_sum_iteration_time_s += time_elapsed_s; + } int tokens_per_second = multi_gpu_config.num_processes * (B * T) / time_elapsed_s; float accumulated_loss = multi_gpu_config.num_processes == 1 ? model.mean_loss : model.accumulated_mean_loss; printf0("step %4d/%d: train loss %f (acc %f) (%f ms, %d tok/s)\n", step + 1, train_num_batches, model.mean_loss, accumulated_loss, time_elapsed_s * 1000, tokens_per_second); logger_log_train(&logger, step, model.mean_loss); } - // add a total average, for optimizations that are only mild improvements - printf0("total average iteration time: %f ms\n", total_sum_iteration_time_s / train_num_batches * 1000); + // add a total average, for optimizations that are only mild improvements (excluding 1st batch as warmup) + printf0("total average iteration time: %f ms\n", total_sum_iteration_time_s / (train_num_batches-1) * 1000); // free dataloader_free(&train_loader); @@ -2311,6 +2617,13 @@ int main(int argc, char *argv[]) { free(cpu_logits_raw); free(cpu_logits); free(gen_tokens); + + #ifdef ENABLE_CUDNN + if (cudnn_workspace_size > 0) { + cudaCheck(cudaFree(cudnn_workspace)); + } + #endif + cudaCheck(cudaFree(cublaslt_workspace)); cublasCheck(cublasDestroy(cublas_handle)); cublasCheck(cublasLtDestroy(cublaslt_handle));