attention_forward.cu ~matches train_gpt2cu & allocate less memory for stats in "att"

This commit is contained in:
ademeure 2024-05-01 16:53:01 +01:00
parent be54842622
commit 2b540b61da
2 changed files with 105 additions and 68 deletions

View file

@ -5,8 +5,11 @@ If you do not have CUDNN, you can remove ENABLE_CUDNN to run the other kernels
You need cuDNN from: https://developer.nvidia.com/cudnn
And the cuDNN front-end from: https://github.com/NVIDIA/cudnn-frontend/tree/main
Compile example:
nvcc -I/path/to/cudnn-frontend/include -DENABLE_CUDNN -O3 --use_fast_math -lcublas -lcudnn attention_forward.cu -o attention_forward
Compile example with cuDNN:
nvcc -I/PATH/TO/cudnn-frontend/include -DENABLE_CUDNN -O3 --use_fast_math -lcublas -lcudnn attention_forward.cu -o attention_forward
Compile example without cuDNN:
nvcc -O3 --use_fast_math -lcublas attention_forward.cu -o attention_forward
version 1 is naive port from CPU code to kernel, parallelize over batch, time, heads only
./attention_forward 1
@ -37,7 +40,7 @@ version 10 is using cuDNN Flash Attention using FP16 or BF16, see:
https://github.com/NVIDIA/cudnn-frontend/blob/main/docs/operations/Attention.md
./attention_forward 10
version 11 is kernel 10 skipping FP16/FP32 conversions (requires fully FP16 network)
version 11 is kernel 10 skipping FP16/FP32 conversions (full FP16/BF16 network)
./attention_forward 11
*/
//#define ENABLE_CUDNN // can be enabled via nvcc "-DENABLE_CUDNN"
@ -74,7 +77,7 @@ namespace fe = cudnn_frontend;
#endif
static cudnnHandle_t cudnn_handle;
static size_t cudnn_workspace_size = 32 * 1024 * 1024; // TODO is this only for backward?
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);
@ -1081,7 +1084,7 @@ void attention_forward5(float* out, floatX* vaccum, floatX* qkvr, floatX* preatt
}
#ifdef ENABLE_CUDNN
using graph_and_tensors = std::tuple<std::shared_ptr<fe::graph::Graph>,
using graph_tensors_fwd = std::tuple<std::shared_ptr<fe::graph::Graph>,
std::shared_ptr<fe::graph::Tensor_attributes>, // Q,
std::shared_ptr<fe::graph::Tensor_attributes>, // K,
std::shared_ptr<fe::graph::Tensor_attributes>, // V,
@ -1090,32 +1093,32 @@ using graph_and_tensors = std::tuple<std::shared_ptr<fe::graph::Graph>,
std::shared_ptr<fe::graph::Tensor_attributes>>; // Stats
// Need a cache because graph->build_operation_graph() is slow but everything else seems fast
using cache_type = std::unordered_map<std::size_t, graph_and_tensors>;
using cache_type_fwd = std::unordered_map<std::size_t, graph_tensors_fwd>;
// Loosely based on cuDNN frontend samples functions and massively simplified
template <typename... Args>
auto lookup_cache_or_build_graph(Args... args) {
static cache_type user_maintained_cache;
auto [b, h, s_qkv, d, is_inference] = std::make_tuple(args...);
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<fe::graph::Graph>();
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, N, 3, NH, d) which cuDNN can handle directly without an external permute
// 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, s_qkv, d})
.set_stride({3 * h * d * s_qkv, d, 3 * h * d, 1}));
.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, s_qkv, d})
.set_stride({3 * h * d * s_qkv, d, 3 * h * d, 1}));
.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, s_qkv, d})
.set_stride({3 * h * d * s_qkv, d, 3 * h * d, 1}));
.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})
@ -1124,90 +1127,114 @@ auto lookup_cache_or_build_graph(Args... args) {
.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);
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, N, NH, d) BF16/FP16 and stats for backward pass is (B, NH, N) FP32
O->set_output(true).set_dim({b, h, s_qkv, d}).set_stride({h * d * s_qkv, d, h * d, 1});
assert(stats == nullptr || is_inference == false);
if (!is_inference) {
stats->set_output(true).set_data_type(fe::DataType_t::FLOAT);
// 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.find(key);
if (it != user_maintained_cache.end()) {
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.insert({key, tuple});
user_maintained_cache_fwd.insert({key, tuple});
return tuple;
}
__global__ void fp32_to_lowp_kernel(lowp_float* out, const float* inp) {
// Used on first run only so we can validate against the CPU results
__global__ void fp32_to_lowp_kernel(floatX* out, const float* inp) {
int idx = blockIdx.x * blockDim.x + threadIdx.x;
out[idx] = (lowp_float)inp[idx];
out[idx] = (floatX)inp[idx];
}
__global__ void lowp_to_fp32_kernel(const lowp_float* inp, float *out) {
__global__ void lowp_to_fp32_kernel(const floatX* inp, float *out) {
int idx = blockIdx.x * blockDim.x + threadIdx.x;
out[idx] = (float)inp[idx];
}
void attention_forward10(float* out, // output: (B, T, NH, HS)
float* stats, // for use in backward pass: (B, NH, T)
lowp_float* vaccum, lowp_float* qkvr,
const float* inp, // input: (B, T, 3, NH, HS) QKV
int B, int T, int C, int NH,
bool skip_conversion=false) {
int64_t HS = C / NH; // number of features per head
bool is_inference = stats != NULL;
float attn_scale_cpu = 1.0 / sqrtf(HS);
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
float* in_fp32, // fp32 input
float* out_fp32, // fp32 output for validation
int B, int T, int C, int NH) {
static bool first_run_validation = true;
int HS = C / NH; // number of features per head
bool is_inference_only = (stats == nullptr);
// Optionally convert from FP32 to FP16/BF16 (always on 1st run to get correct results)
// Convert from FP32 to FP16/BF16 on 1st run to get correct results
const int block_size = 64; // smallest full occupancy block size on modern GPUs
int total_threads = B * T * C * 3;
assert(total_threads % block_size == 0);
int num_blocks = total_threads / block_size;
if (!skip_conversion || first_run_validation) {
fp32_to_lowp_kernel<<<num_blocks, block_size>>>(qkvr, inp);
if (first_run_validation) {
int total_threads = B * T * C * 3;
assert(total_threads % block_size == 0);
int num_blocks = total_threads / block_size;
fp32_to_lowp_kernel<<<num_blocks, block_size>>>(inp, in_fp32);
}
// 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(B, NH, T, HS, is_inference);
lookup_cache_or_build_graph_fwd(B, NH, T, HS, is_inference_only);
//// Build variant pack
void* devPtrQ = qkvr;
void* devPtrK = (qkvr + NH * HS);
void* devPtrV = (qkvr + 2 * NH * HS);
void* devPtrO = (void*)vaccum;
// 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<std::shared_ptr<fe::graph::Tensor_attributes>, void*> variant_pack = {
{Q, devPtrQ}, {K, devPtrK}, {V, devPtrV}, {attn_scale, &attn_scale_cpu}, {O, devPtrO}};
if (is_inference == false) {
variant_pack[softmax_stats] = (void*)stats;
// 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->get_workspace_size() <= cudnn_workspace_size); // TODO - not needed for forward?
assert(graph->execute(cudnn_handle, variant_pack, cudnn_workspace).is_good());
cudaCheck(cudaGetLastError());
// Optionally convert back from FP16/BF16 to FP32
total_threads = B * T * C;
assert(total_threads % block_size == 0);
num_blocks = total_threads / block_size;
if (!skip_conversion || first_run_validation) {
lowp_to_fp32_kernel<<<num_blocks, block_size>>>(vaccum, out);
if (first_run_validation) {
int total_threads = B * T * C;
assert(total_threads % block_size == 0);
int num_blocks = total_threads / block_size;
lowp_to_fp32_kernel<<<num_blocks, block_size>>>(out, out_fp32);
}
cudaCheck(cudaGetLastError());
first_run_validation = false;
}
#endif // ENABLE_CUDNN
@ -1216,7 +1243,7 @@ void attention_forward10(float* out, // output: (B, T, NH, HS)
void attention_forward(int kernel_num,
float* out, float* stats, float* vaccum,
float* qkvr, float* preatt, float* att,
const float* inp,
float* inp,
int B, int T, int C, int NH,
const int block_size) {
switch (kernel_num) {
@ -1244,12 +1271,10 @@ void attention_forward(int kernel_num,
break;
#ifdef ENABLE_CUDNN
case 10:
attention_forward10(out, stats, (lowp_float*)vaccum, (lowp_float*)qkvr,
inp, B, T, C, NH, false);
break;
case 11: // skip permutes for perf passes (to analyse perf as if in/out were truly 16-bit)
attention_forward10(out, stats, (lowp_float*)vaccum, (lowp_float*)qkvr,
inp, B, T, C, NH, true);
// note: validation only cares about out, which is out_fp32 of the function
// inp is hackily converted to FP16 into qkvr only on the first run
// similarly, vaccum is converted to FP32 into out only on the first run
attention_forward_cudnn((floatX*)vaccum, stats, (floatX*)qkvr, inp, out, B, T, C, NH);
break;
#endif
default:
@ -1281,7 +1306,6 @@ int main(int argc, char **argv) {
#ifdef ENABLE_CUDNN
checkCudnnErr(cudnnCreate(&cudnn_handle));
cudaCheck(cudaMalloc(&cudnn_workspace, cudnn_workspace_size));
#endif
// create host memory of random numbers
@ -1369,5 +1393,11 @@ int main(int argc, char **argv) {
cudaCheck(cudaFree(d_inp));
cublasDestroy(cublas_handle);
#ifdef ENABLE_CUDNN
if (cudnn_workspace_size > 0) {
cudaCheck(cudaFree(cudnn_workspace));
}
#endif
return 0;
}

View file

@ -1691,7 +1691,7 @@ typedef struct {
floatX* ln1_mean; // (L, B, T)
floatX* ln1_rstd; // (L, B, T)
floatX* atty; // (L, B, T, C)
floatX* att; // (L, B, NH, T, T)
floatX* att; // (L, B, NH, T, T) (smaller with cuDNN)
floatX* attproj; // (L, B, T, C)
floatX* residual2; // (L, B, T, C)
floatX* ln2; // (L, B, T, C)
@ -1725,7 +1725,12 @@ void fill_in_activation_sizes(size_t* act_sizes, size_t B, size_t T, GPT2Config
act_sizes[2] = L * B * T; // ln1_mean
act_sizes[3] = L * B * T; // ln1_rstd
act_sizes[4] = L * B * T * C; // atty
#ifdef ENABLE_CUDNN
// FP32 stats tensor for cuDNN to be passed to backward pass
act_sizes[5] = L * B * NH * T * (sizeof(float) / sizeof(floatX));
#else
act_sizes[5] = L * B * NH * T * T; // att
#endif
act_sizes[6] = L * B * T * C; // attproj
act_sizes[7] = L * B * T * C; // residual2
act_sizes[8] = L * B * T * C; // ln2
@ -1997,7 +2002,6 @@ void gpt2_forward(GPT2 *model, int* inputs, int* targets, size_t B, size_t T) {
floatX* l_ln1_rstd = acts.ln1_rstd + l * B * T;
floatX* l_qkvr = acts.qkvr + l * B * T * 3*C;
floatX* l_atty = acts.atty + l * B * T * C;
floatX* l_att = acts.att + l * B * NH * T * T;
floatX* l_attproj = acts.attproj + l * B * T * C;
floatX* l_residual2 = acts.residual2 + l * B * T * C;
floatX* l_ln2 = acts.ln2 + l * B * T * C;
@ -2012,9 +2016,11 @@ void gpt2_forward(GPT2 *model, int* inputs, int* targets, size_t B, size_t T) {
layernorm_forward(l_ln1, l_ln1_mean, l_ln1_rstd, residual, l_ln1w, l_ln1b, B, T, C);
#ifdef ENABLE_CUDNN
float* l_att = (float*)acts.att + l * B * NH * T; // cuDNN needs a smaller FP32 tensor
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
floatX* l_att = acts.att + l * B * NH * T * T;
// these are only needed as scratchpads for the forward pass, but
// need not be stored for backward
floatX* scratch = (floatX*)acts.output;
@ -2144,7 +2150,6 @@ void gpt2_backward(GPT2 *model) {
floatX* l_ln1_rstd = acts.ln1_rstd + l * B * T;
floatX* l_qkvr = acts.qkvr + l * B * T * 3*C;
floatX* l_atty = acts.atty + l * B * T * C;
floatX* l_att = acts.att + l * B * NH * T * T;
floatX* l_residual2 = acts.residual2 + l * B * T * C;
floatX* l_ln2 = acts.ln2 + l * B * T * C;
floatX* l_ln2_mean = acts.ln2_mean + l * B * T;
@ -2169,8 +2174,10 @@ void gpt2_backward(GPT2 *model) {
matmul_backward(dl_btc, dl_attprojw, dl_attprojb, dresidual, l_atty, l_attprojw, B, T, C, C);
#ifdef ENABLE_CUDNN
float* l_att = (float*)acts.att + l * B * NH * T; // cuDNN needs a smaller FP32 tensor
attention_backward_cudnn(dl_bt4c, dl_btc, l_qkvr, l_atty, (float*)l_att, B, T, NH, C);
#else
floatX* l_att = acts.att + l * B * NH * T * T;
// 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