From 902c5f53aebf272337d377740ba5f0c6abe14f72 Mon Sep 17 00:00:00 2001 From: Yufeng Li Date: Thu, 11 May 2023 13:47:15 -0700 Subject: [PATCH] add cutlass fmha support in PackedAttention (#15838) ### Description Support cutlass fMHA in PackedAttention. Though we have fMHA trt kernel, it doesn't support relative bias position. Cutlass fmha has support for RBP and also support lower end GPUs(5.3, 6.x). ### Motivation and Context --- .../contrib_ops/cuda/bert/attention_impl.cu | 6 +- .../contrib_ops/cuda/bert/packed_attention.cc | 19 +- .../cuda/bert/packed_attention_impl.cu | 170 ++++++++++++------ .../cuda/bert/packed_attention_impl.h | 5 +- .../contrib_ops/packed_attention_op_test.cc | 38 ++-- .../testdata/packed_attention_fp16.rbp.onnx | Bin 0 -> 1299 bytes .../testdata/packed_attention_fp16.rbp.py | 131 ++++++++++++++ 7 files changed, 302 insertions(+), 67 deletions(-) create mode 100644 onnxruntime/test/testdata/packed_attention_fp16.rbp.onnx create mode 100644 onnxruntime/test/testdata/packed_attention_fp16.rbp.py diff --git a/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu b/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu index d6741d68f3..f6b77263e7 100644 --- a/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu +++ b/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu @@ -858,6 +858,10 @@ Status QkvToContext( query = data.query; } + DUMP_TENSOR_D("attention q(BSNH)", q, batch_size * sequence_length, num_heads * qk_head_size); + DUMP_TENSOR_D("attention k(BSNH)", k, batch_size * sequence_length, num_heads * qk_head_size); + DUMP_TENSOR_D("attention v(BSNH)", v, batch_size * sequence_length, num_heads * v_head_size); + MemoryEfficientAttentionParams p; p.sm = device_prop.major * 10 + device_prop.minor; p.is_half = sizeof(T) == 2; @@ -881,7 +885,7 @@ Status QkvToContext( p.workspace = MemoryEfficientAttentionParams::need_workspace(v_head_size, sizeof(T) == sizeof(float)) ? scratch1 : nullptr; p.stream = stream; run_memory_efficient_attention(p); - DUMP_TENSOR("cutlass output", data.output, batch_size * sequence_length, num_heads, v_head_size); + DUMP_TENSOR("attention cutlass output", data.output, batch_size * sequence_length, num_heads, v_head_size); return Status::OK(); } #endif diff --git a/onnxruntime/contrib_ops/cuda/bert/packed_attention.cc b/onnxruntime/contrib_ops/cuda/bert/packed_attention.cc index 54d69fd8f1..8ea74dd178 100644 --- a/onnxruntime/contrib_ops/cuda/bert/packed_attention.cc +++ b/onnxruntime/contrib_ops/cuda/bert/packed_attention.cc @@ -274,6 +274,20 @@ Status PackedAttention::ComputeInternal(OpKernelContext* context) const { MHARunner* fused_runner = TryGettingFusedRunner(parameters); + bool use_memory_efficient_attention = false; + auto& device_prop = GetDeviceProp(); +#if USE_FLASH_ATTENTION + if (nullptr == fused_runner) { + int sm = device_prop.major * 10 + device_prop.minor; + bool is_good_for_rpb = !parameters.has_relative_position_bias || parameters.sequence_length % (4 * sizeof(T)) == 0; + use_memory_efficient_attention = is_good_for_rpb && + sizeof(T) == 2 && // only enable for fp16 + (parameters.head_size & 7) == 0 && + (parameters.v_head_size & 7) == 0 && + has_memory_efficient_attention(sm, sizeof(T) == 2); + } +#endif + typedef typename ToCudaType::MappedType CudaT; CudaT one = ToCudaType::FromFloat(1.0f); CudaT zero = ToCudaType::FromFloat(0.0f); @@ -284,7 +298,6 @@ Status PackedAttention::ComputeInternal(OpKernelContext* context) const { int k = parameters.input_hidden_size; gemm_buffer = GetScratchBuffer(static_cast(m) * n, context->GetComputeStream()); - auto& device_prop = GetDeviceProp(); cublasHandle_t cublas = GetCublasHandle(context); // Gemm, note that CUDA assumes col-major, so result(N, M) = 1 * weights x input + 1 x bias @@ -302,7 +315,8 @@ Status PackedAttention::ComputeInternal(OpKernelContext* context) const { parameters.head_size, parameters.v_head_size, parameters.sequence_length, - fused_runner); + fused_runner, + use_memory_efficient_attention); auto work_space = GetScratchBuffer(workSpaceSize, context->GetComputeStream()); typedef typename ToCudaType::MappedType CudaT; @@ -315,6 +329,7 @@ Status PackedAttention::ComputeInternal(OpKernelContext* context) const { data.cumulative_sequence_length = cumulative_sequence_length->Data(); data.output = reinterpret_cast(output->MutableData()); data.fused_runner = reinterpret_cast(fused_runner); + data.use_memory_efficient_attention = use_memory_efficient_attention; return QkvToContext(device_prop, cublas, Stream(context), parameters, data); } diff --git a/onnxruntime/contrib_ops/cuda/bert/packed_attention_impl.cu b/onnxruntime/contrib_ops/cuda/bert/packed_attention_impl.cu index 0a218fab47..15718ec3c2 100644 --- a/onnxruntime/contrib_ops/cuda/bert/packed_attention_impl.cu +++ b/onnxruntime/contrib_ops/cuda/bert/packed_attention_impl.cu @@ -46,7 +46,8 @@ size_t GetAttentionWorkspaceSize( size_t qk_head_size, size_t v_head_size, size_t sequence_length, - void* fused_runner) { + void* fused_runner, + bool use_memory_efficient_attention) { // Note that q, k and v might need alignment for fused attention kernels. const size_t qkv_bytes = element_size * batch_size * num_heads * sequence_length * (qk_head_size + qk_head_size + v_head_size); @@ -54,6 +55,19 @@ size_t GetAttentionWorkspaceSize( return qkv_bytes; } +#if USE_FLASH_ATTENTION + if (use_memory_efficient_attention) { + size_t fmha_buffer_bytes = 0; + if (MemoryEfficientAttentionParams::need_workspace(v_head_size, element_size == sizeof(float))) { + fmha_buffer_bytes = batch_size * sequence_length * num_heads * v_head_size * sizeof(float); + } + + return qkv_bytes + fmha_buffer_bytes; + } +#else + ORT_UNUSED_PARAMETER(use_memory_efficient_attention); +#endif + return qkv_bytes + 2 * GetAttentionScratchSize(element_size, batch_size, num_heads, sequence_length); } @@ -142,7 +156,7 @@ __global__ void AddBiasTransposeQKVPacked( // Block: 256 // For memory efficient fMHA from CUTLASS. For future use, doesn't support fMHA from CUTLASS yet. // Input: Tx3xNxH -// Output: 3xBxNxSxH +// Output: 3xTxNxH // T is token_count // B is batch_size // S is sequence_length @@ -152,48 +166,27 @@ template __global__ void AddBiasTransposeQKVPackedCutlass( const T* input, const T* biases, - int32_t N, - int32_t H_QK, - int32_t H_V, + int32_t D_QK, + int32_t D_V, T* q, T* k, T* v, - const int32_t* token_offset, int32_t token_count) { - int s = blockIdx.x; - int b = blockIdx.y; + int token_idx = blockIdx.x; - int S = gridDim.x; + input += token_idx * (D_QK + D_QK + D_V); + q += token_idx * D_QK; + k += token_idx * D_QK; + v += token_idx * D_V; - const int packing_token_idx = b * S + s; - const int padding_token_idx = token_offset[packing_token_idx]; - b = padding_token_idx / S; - s = padding_token_idx - b % S; - - input += packing_token_idx * N * (H_QK + H_QK + H_V); - int k_offset = N * H_QK; - int v_offset = N * H_QK + N * H_QK; - q += (b * S * N + s * N) * H_QK; - k += (b * S * N + s * N) * H_QK; - v += (b * S * N + s * N) * H_V; - - if (packing_token_idx < token_count) { - for (int i = threadIdx.x; i < N * H_QK; i += blockDim.x) { + if (token_idx < token_count) { + for (int i = threadIdx.x; i < D_QK; i += blockDim.x) { q[i] = input[i] + biases[i]; - k[i] = input[i + k_offset] + biases[i + k_offset]; + k[i] = input[D_QK + i] + biases[D_QK + i]; } - for (int i = threadIdx.x; i < N * H_V; i += blockDim.x) { - v[i] = input[i + v_offset] + biases[i + v_offset]; - } - } else { - for (int i = threadIdx.x; i < N * H_QK; i += blockDim.x) { - q[i] = biases[i]; - k[i] = biases[i + k_offset]; - } - - for (int i = threadIdx.x; i < N * H_V; i += blockDim.x) { - v[i] = biases[i + v_offset]; + for (int i = threadIdx.x; i < D_V; i += blockDim.x) { + v[i] = input[D_QK + D_QK + i] + biases[D_QK + D_QK + i]; } } } @@ -254,18 +247,16 @@ void InvokeAddBiasTranspose( output + 2 * batch_size * sequence_length * num_heads * qk_head_size, token_offset, token_count); - } else if (format == AttentionQkvFormat::Q_K_V_BSNH) { // TODO: add memory efficient support - const dim3 grid(sequence_length, batch_size); + } else if (format == AttentionQkvFormat::Q_K_V_BSNH) { + const dim3 grid(token_count); AddBiasTransposeQKVPackedCutlass<<>>( input, biases, - num_heads, - qk_head_size, - v_head_size, + num_heads * qk_head_size, + num_heads * v_head_size, output, - output + batch_size * sequence_length * num_heads * qk_head_size, - output + 2 * batch_size * sequence_length * num_heads * qk_head_size, - token_offset, + output + token_count * num_heads * qk_head_size, + output + 2 * token_count * num_heads * qk_head_size, token_count); } else { ORT_ENFORCE(format == AttentionQkvFormat::QKV_BSN3H); @@ -381,15 +372,14 @@ Status LaunchTransposeRemovePadding( const int batch_size, const int seq_len, const int number_heads, const int head_size, cudaStream_t stream); - // input: [batch_size, number_heads, seq_len, head_size] - // output: [token_count, number_heads * head_size] +// input: [batch_size, number_heads, seq_len, head_size] +// output: [token_count, number_heads * head_size] template <> Status LaunchTransposeRemovePadding( half* output, const half* input, const int* token_offset, const int token_count, const int batch_size, const int seq_len, const int number_heads, const int head_size, cudaStream_t stream) { - // Make sure memory is aligned to 128 bit ORT_ENFORCE(!(reinterpret_cast(input) & 0xF) && !(reinterpret_cast(output) & 0xF), "alignment"); @@ -476,6 +466,74 @@ Status FusedScaledDotProductAttention( return Status::OK(); } +#if USE_FLASH_ATTENTION +template +Status FusedScaledDotProductAttentionCutlass( + const cudaDeviceProp& device_prop, + cudaStream_t stream, + PackedAttentionParameters& parameters, + PackedAttentionData& data) { + const int batch_size = parameters.batch_size; + const int sequence_length = parameters.sequence_length; + const int num_heads = parameters.num_heads; + const int qk_head_size = parameters.head_size; + const int v_head_size = parameters.v_head_size; + LaunchAddBiasTranspose(data.gemm_buffer, data.bias, data.workspace, + batch_size, sequence_length, + num_heads, qk_head_size, v_head_size, + AttentionQkvFormat::Q_K_V_BSNH, data.token_offset, + parameters.token_count, stream); + DUMP_TENSOR_INIT(); + + DUMP_TENSOR_D("PackedAttention cutlass data.gemm_buffer", data.gemm_buffer, parameters.token_count, 3, num_heads * qk_head_size); + DUMP_TENSOR_D("PackedAttention cutlass data.bias", data.bias, 1, 3 * num_heads * qk_head_size); + + // Q, K and V pointers + const int model_dimension_qk = num_heads * qk_head_size; + const int model_dimension_v = num_heads * v_head_size; + const size_t elements_qk = static_cast(parameters.token_count) * static_cast(model_dimension_qk); + const size_t elements_v = static_cast(parameters.token_count) * static_cast(model_dimension_v); + T* qkv = data.workspace; + T* query = qkv; + T* key = query + elements_qk; + T* value = key + elements_qk; + T* accum_workspace = value + elements_v; + + DUMP_TENSOR_D("PackedAttention cutlass q(BSNH)", query, parameters.token_count, num_heads * qk_head_size); + DUMP_TENSOR_D("PackedAttention cutlass k(BSNH)", key, parameters.token_count, num_heads * qk_head_size); + DUMP_TENSOR_D("PackedAttention cutlass v(BSNH)", value, parameters.token_count, num_heads * v_head_size); + DUMP_TENSOR_D("PackedAttention cutlass cumulative_sequence_length", data.cumulative_sequence_length, 1, batch_size + 1); + + MemoryEfficientAttentionParams p; + p.sm = device_prop.major * 10 + device_prop.minor; + p.is_half = sizeof(T) == 2; + p.batch_size = parameters.batch_size; + p.num_heads = parameters.num_heads; + p.sequence_length = parameters.sequence_length; + p.kv_sequence_length = parameters.sequence_length; + p.qk_head_size = parameters.head_size; + p.v_head_size = parameters.v_head_size; + p.causal = false; + p.scale = parameters.scale == 0.0f ? 1.f / sqrt(static_cast(qk_head_size)) + : parameters.scale; + p.seqlen_k_ptr = nullptr; + p.seqstart_q_ptr = const_cast(data.cumulative_sequence_length); + p.seqstart_k_ptr = const_cast(data.cumulative_sequence_length); + p.query = query; + p.key = key; + p.value = value; + p.attn_bias = data.relative_position_bias; + p.is_attn_bias_batched = !parameters.broadcast_res_pos_bias; + p.output = data.output; + p.workspace = MemoryEfficientAttentionParams::need_workspace(v_head_size, sizeof(T) == sizeof(float)) ? accum_workspace : nullptr; + p.stream = stream; + run_memory_efficient_attention(p); + + DUMP_TENSOR("PackedAttention cutlass output", data.output, parameters.token_count, num_heads, v_head_size); + return Status::OK(); +} +#endif + template Status UnfusedScaledDotProductAttention( const cudaDeviceProp& device_prop, @@ -515,8 +573,8 @@ Status UnfusedScaledDotProductAttention( // Q, K and V are ready now DUMP_TENSOR_INIT(); - DUMP_TENSOR_D("gemm_buffer", data.gemm_buffer, parameters.token_count, (num_heads * (qk_head_size * 2 + v_head_size))); - DUMP_TENSOR_D("data.workspace", data.workspace, 3 * batch_size, num_heads, sequence_length, qk_head_size); + DUMP_TENSOR_D("PackedAttention unfused gemm_buffer", data.gemm_buffer, parameters.token_count, (num_heads * (qk_head_size * 2 + v_head_size))); + DUMP_TENSOR_D("PackedAttention unfused data.workspace", data.workspace, 3 * batch_size, num_heads, sequence_length, qk_head_size); // Compute Q*K' (as K'*Q), scaled by 1/sqrt(H) and store in scaled_qk: BxNxSxT // Q: BxNxSxH, K: BxNxSxH, Q*K': BxNxSxS @@ -537,7 +595,7 @@ Status UnfusedScaledDotProductAttention( scaled_qk, sequence_length, sequence_length * sequence_length, batches, device_prop)); - DUMP_TENSOR_D("QK", scaled_qk, batch_size * num_heads, sequence_length, sequence_length); + DUMP_TENSOR_D("PackedAttention unfused QK", scaled_qk, batch_size * num_heads, sequence_length, sequence_length); const size_t bytes = GetAttentionScratchSize(element_size, batch_size, num_heads, sequence_length); @@ -554,7 +612,7 @@ Status UnfusedScaledDotProductAttention( num_heads, attention_score, stream)); - DUMP_TENSOR_D("Softmax", attention_score, batch_size * num_heads, sequence_length, sequence_length); + DUMP_TENSOR_D("PackedAttention unfused Softmax", attention_score, batch_size * num_heads, sequence_length, sequence_length); // compute R*V (as V*R), and store in temp_output (space used by Q): BxNxSxH_v T* temp_output = qkv; @@ -572,7 +630,7 @@ Status UnfusedScaledDotProductAttention( batch_size, sequence_length, num_heads, v_head_size, stream); - DUMP_TENSOR("unfused output", data.output, parameters.token_count, num_heads, v_head_size); + DUMP_TENSOR("PackedAttention unfused output", data.output, parameters.token_count, num_heads, v_head_size); return result; } @@ -586,9 +644,15 @@ Status QkvToContext( void* fused_runner = data.fused_runner; if (nullptr != fused_runner) { return FusedScaledDotProductAttention(device_prop, stream, parameters, data); - } else { - return UnfusedScaledDotProductAttention(device_prop, cublas, stream, parameters, data); } + +#if USE_FLASH_ATTENTION + if (data.use_memory_efficient_attention) { + return FusedScaledDotProductAttentionCutlass(device_prop, stream, parameters, data); + } +#endif + + return UnfusedScaledDotProductAttention(device_prop, cublas, stream, parameters, data); } template Status QkvToContext( diff --git a/onnxruntime/contrib_ops/cuda/bert/packed_attention_impl.h b/onnxruntime/contrib_ops/cuda/bert/packed_attention_impl.h index 1eb26575df..e3288649e2 100644 --- a/onnxruntime/contrib_ops/cuda/bert/packed_attention_impl.h +++ b/onnxruntime/contrib_ops/cuda/bert/packed_attention_impl.h @@ -24,7 +24,8 @@ size_t GetAttentionWorkspaceSize( size_t qk_head_size, size_t v_head_size, size_t sequence_length, - void* fused_runner); + void* fused_runner, + bool use_memory_efficient_attention); template struct PackedAttentionData { @@ -38,6 +39,8 @@ struct PackedAttentionData { T* output; void* fused_runner; + + bool use_memory_efficient_attention; }; template diff --git a/onnxruntime/test/contrib_ops/packed_attention_op_test.cc b/onnxruntime/test/contrib_ops/packed_attention_op_test.cc index 5bf785185e..98792e9d20 100644 --- a/onnxruntime/test/contrib_ops/packed_attention_op_test.cc +++ b/onnxruntime/test/contrib_ops/packed_attention_op_test.cc @@ -385,14 +385,8 @@ static void RunModelWithRandomInput( int64_t batch_size, int64_t sequence_length, std::string& onnx_model, - bool is_float16) { - // ORT enables TF32 in GEMM for A100. TF32 will cause precsion loss and fail this test. - // Do not run this test unless TF32 is disabled explicitly. - if (HasCudaEnvironment(800) && ParseEnvironmentVariableWithDefault("NVIDIA_TF32_OVERRIDE", 1) != 0) { - GTEST_SKIP() << "Skipping RunModelWithRandomInput in A100 since TF32 is enabled"; - return; - } - + bool is_float16, + bool has_rbp = false) { RandomValueGenerator random{234}; constexpr int hidden_size = 768; @@ -455,6 +449,16 @@ static void RunModelWithRandomInput( test.AddInput("token_offset", token_offset_dims, token_offset); test.AddInput("cumulative_sequence_length", cum_seq_len_dims, cum_seq_len); + if (has_rbp) { + std::vector rbp_dims{1, num_heads, sequence_length, sequence_length}; + std::vector rbp_data = random.Gaussian(rbp_dims, 0.0f, 0.1f); + if (is_float16) { + test.AddInput("rbp", rbp_dims, ToFloat16(rbp_data)); + } else { + test.AddInput("rbp", rbp_dims, rbp_data); + } + } + std::vector> execution_providers; execution_providers.push_back(DefaultCudaExecutionProvider()); test.AddReferenceOutputs(onnx_model, gpu_threshold, DefaultCudaExecutionProvider()); @@ -462,7 +466,7 @@ static void RunModelWithRandomInput( } } -TEST(PackedAttentionTest, test_on_random_data) { +TEST(PackedAttentionTest, TestWithRandomData) { std::string onnx_model = "testdata/packed_attention_fp32.onnx"; std::string onnx_model_fp16 = "testdata/packed_attention_fp16.onnx"; for (int batch_size : std::vector({1, 2, 3, 4, 5, 6, 7, 8})) { @@ -481,7 +485,21 @@ TEST(PackedAttentionTest, test_on_random_data) { } } -TEST(PackedAttentionTest, test_on_random_data_large_seq) { +TEST(PackedAttentionTest, TestWithRandomDataWithRBP) { + std::string onnx_model_fp16 = "testdata/packed_attention_fp16.rbp.onnx"; // mainly for cutlass + for (int batch_size : std::vector({1, 2, 3, 4, 5, 6, 7, 8})) { + for (int sequence_length : std::vector({32, 48, 64, 95, 128})) { + RunModelWithRandomInput( + batch_size, + sequence_length, + onnx_model_fp16, + true /*is_float16*/, + true /*has_rbp*/); + } + } +} + +TEST(PackedAttentionTest, TestWithRandomDataLargeSeq) { int batch_size = 2; int sequence_length = 1152; // > 1024 std::string onnx_model = "testdata/packed_attention_fp32.onnx"; diff --git a/onnxruntime/test/testdata/packed_attention_fp16.rbp.onnx b/onnxruntime/test/testdata/packed_attention_fp16.rbp.onnx new file mode 100644 index 0000000000000000000000000000000000000000..89257df2ff38d21cbd94590a7282a7ef459b583e GIT binary patch literal 1299 zcma)5O>fgc5UragaVBY-Y*lC-0m4EEMM{W46{3eq1WrguD81#t%CR@eQhPUe*Cpc8 z-@re}pW^p~^;cpC2;I}pc;3w0H_v(%opHt%LC7TaNlykg6B8czBoIc$>RA`P{^hS4 zTp@5c6Vhf9_l?}$)ZeH1jsfLc+Y3qONY|RxkNRGEmfAr9ZyAtbu(boPP4q6YyHk_QCOZP%r2ig_qimC+m~`Oj(EGLMRxOZ7~@=Tj%( zf=?tYHTXv_sC*}MI+JiPrnZ1JFW7@pFrK59EhSMlr~&Yt%FVzYAp2e zjM%O?Yb>?>P8(YBJ15kWB*3u|PgkqE)l$>8DgM;p&siHjsFI9F?SS}PEjg2