diff --git a/cmake/external/composable_kernel.cmake b/cmake/external/composable_kernel.cmake index 8ca9c1516c..b2331582f9 100644 --- a/cmake/external/composable_kernel.cmake +++ b/cmake/external/composable_kernel.cmake @@ -1,5 +1,5 @@ set(composable_kernel_URL https://github.com/ROCmSoftwarePlatform/composable_kernel.git) -set(composable_kernel_TAG 0345963eef4f92e9c5eab608bb8557b5463a1dcb) # 2022-12-15 15:07:24 -0600 +set(composable_kernel_TAG bef0cb20dba0d9b315df46899310478a81c21852) # 2023-02-16 11:54:08 -0800 set(PATCH ${PROJECT_SOURCE_DIR}/patches/composable_kernel/Fix_Clang_Build.patch) diff --git a/cmake/onnxruntime_providers.cmake b/cmake/onnxruntime_providers.cmake index dc6e361a17..91abd44707 100644 --- a/cmake/onnxruntime_providers.cmake +++ b/cmake/onnxruntime_providers.cmake @@ -1508,8 +1508,9 @@ if (onnxruntime_USE_ROCM) target_link_libraries(onnxruntime_providers_rocm PRIVATE onnxruntime_composable_kernel_includes # Currently we shall not use composablekernels::device_operations, the target includes all conv dependencies, which - # are extremely slow to compile. Instead, we only link all gemm related objects. See the following link on updating. - # https://github.com/ROCmSoftwarePlatform/composable_kernel/blob/85978e0201/library/src/tensor_operation_instance/gpu/CMakeLists.txt#L33-L54 + # are extremely slow to compile. Instead, we only link all gemm related objects. See the following directory on + # updating. + # https://github.com/ROCmSoftwarePlatform/composable_kernel/tree/develop/library/src/tensor_operation_instance/gpu device_gemm_instance device_gemm_add_fastgelu_instance device_gemm_fastgelu_instance diff --git a/cmake/patches/composable_kernel/Fix_Clang_Build.patch b/cmake/patches/composable_kernel/Fix_Clang_Build.patch index 7ee4c8bfaf..ae2a015c1d 100644 --- a/cmake/patches/composable_kernel/Fix_Clang_Build.patch +++ b/cmake/patches/composable_kernel/Fix_Clang_Build.patch @@ -1,5 +1,5 @@ diff --git a/CMakeLists.txt b/CMakeLists.txt -index f861e302..f0b6bcea 100644 +index f861e3020..f0b6bceae 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -1,7 +1,7 @@ @@ -71,14 +71,18 @@ index f861e302..f0b6bcea 100644 - HEADER_ONLY -) diff --git a/library/src/tensor_operation_instance/gpu/CMakeLists.txt b/library/src/tensor_operation_instance/gpu/CMakeLists.txt -index c206c4dc..e45fac9d 100644 +index c206c4dc0..b283eeb64 100644 --- a/library/src/tensor_operation_instance/gpu/CMakeLists.txt +++ b/library/src/tensor_operation_instance/gpu/CMakeLists.txt -@@ -1,7 +1,9 @@ +@@ -1,7 +1,13 @@ function(add_instance_library INSTANCE_NAME) message("adding instance ${INSTANCE_NAME}") + set_source_files_properties(${ARGN} PROPERTIES LANGUAGE HIP) add_library(${INSTANCE_NAME} OBJECT ${ARGN}) ++ # Always disable debug symbol and C debug assert due to ++ # - Linker error: ... relocation truncated to fit ..., caused by object files to be linked are too huge. ++ # - https://github.com/ROCmSoftwarePlatform/composable_kernel/issues/622 ++ target_compile_options(${INSTANCE_NAME} PRIVATE -g0 -DNDEBUG) target_compile_features(${INSTANCE_NAME} PUBLIC) + target_compile_definitions(${INSTANCE_NAME} PRIVATE "__HIP_PLATFORM_AMD__=1" "__HIP_PLATFORM_HCC__=1") set_target_properties(${INSTANCE_NAME} PROPERTIES POSITION_INDEPENDENT_CODE ON) diff --git a/onnxruntime/contrib_ops/rocm/bert/attention.cu b/onnxruntime/contrib_ops/rocm/bert/attention.cu index da4013a5e9..c064196ec5 100644 --- a/onnxruntime/contrib_ops/rocm/bert/attention.cu +++ b/onnxruntime/contrib_ops/rocm/bert/attention.cu @@ -83,13 +83,17 @@ Status Attention::ComputeInternal(OpKernelContext* context) const { using HipT = typename ToHipType::MappedType; using QkvProjectGeneric = GemmPermuteGenericPipeline; using AttentionGeneric = GemmSoftmaxGemmPermuteGenericPipeline; + using AttentionTunableOp = GemmSoftmaxGemmPermuteTunableOp; size_t qkv_project_output_bytes = QkvProjectGeneric::GetOutputNumBytes(&attn); - size_t attention_workspace_bytes = AttentionGeneric::GetWorkspaceNumBytes(&attn); - ORT_ENFORCE(QkvProjectGeneric::GetWorkspaceNumBytes(&attn) <= attention_workspace_bytes); // workspace reuse + size_t shared_workspace_bytes = std::max(QkvProjectGeneric::GetWorkspaceNumBytes(&attn), + AttentionGeneric::GetWorkspaceNumBytes(&attn)); + if (GetTuningContext()->IsTunableOpEnabled()) { + shared_workspace_bytes = std::max(shared_workspace_bytes, AttentionTunableOp::GetWorkspaceNumBytes(&attn)); + } auto qkv_project_output = GetScratchBuffer(qkv_project_output_bytes, context->GetComputeStream()); - auto workspace = GetScratchBuffer(attention_workspace_bytes, context->GetComputeStream()); + auto workspace = GetScratchBuffer(shared_workspace_bytes, context->GetComputeStream()); GemmPermuteParams gemm_permute_params; { @@ -105,7 +109,7 @@ Status Attention::ComputeInternal(OpKernelContext* context) const { params.bias_buffer = reinterpret_cast(bias->DataRaw()); params.out_buffer = reinterpret_cast(qkv_project_output.get()); params.ones = GetConstOnes(attn.batch_size * attn.sequence_length, stream); - params.workspace_buffer = reinterpret_cast(workspace.get()); // workspace reuse + params.workspace_buffer = reinterpret_cast(workspace.get()); } ORT_RETURN_IF_ERROR(QkvProjectGeneric::Run(&gemm_permute_params)); @@ -118,16 +122,16 @@ Status Attention::ComputeInternal(OpKernelContext* context) const { const int batches = attn.batch_size * attn.num_heads; const int present_size_per_batch = attn.total_sequence_length * attn.head_size; ORT_RETURN_IF_ERROR( - LaunchConcatPastToPresent(Stream(context), - attn.total_sequence_length, - attn.sequence_length, - attn.batch_size, - attn.head_size, - attn.num_heads, - device_prop.maxThreadsPerBlock, - nullptr == past ? nullptr : reinterpret_cast(past->DataRaw()), - k_buffer, - reinterpret_cast(present->MutableDataRaw()))); + LaunchConcatPastToPresent(Stream(context), + attn.total_sequence_length, + attn.sequence_length, + attn.batch_size, + attn.head_size, + attn.num_heads, + device_prop.maxThreadsPerBlock, + nullptr == past ? nullptr : reinterpret_cast(past->DataRaw()), + k_buffer, + reinterpret_cast(present->MutableDataRaw()))); // update pointers to present_k and present_v. k_buffer = reinterpret_cast(present->MutableDataRaw()); @@ -159,13 +163,18 @@ Status Attention::ComputeInternal(OpKernelContext* context) const { if (mask_index != nullptr) { params.mask_index_buffer = mask_index->Data(); - params.mask_index_dims = mask_index->Shape().GetDims(); + params.mask_index_dims = mask_index->Shape().AsShapeVector(); } params.workspace_buffer = reinterpret_cast(workspace.get()); } - return AttentionGeneric::Run(&gemm_softmax_gemm_permute_params, use_persistent_softmax); + if (this->GetTuningContext()->IsTunableOpEnabled() && + !use_persistent_softmax) { + return AttentionTunableOp{}(&gemm_softmax_gemm_permute_params); + } else { + return AttentionGeneric::Run(&gemm_softmax_gemm_permute_params, use_persistent_softmax); + } } } // namespace rocm diff --git a/onnxruntime/contrib_ops/rocm/bert/attention_impl.cu b/onnxruntime/contrib_ops/rocm/bert/attention_impl.cu index 8750334303..32f225451e 100644 --- a/onnxruntime/contrib_ops/rocm/bert/attention_impl.cu +++ b/onnxruntime/contrib_ops/rocm/bert/attention_impl.cu @@ -67,10 +67,10 @@ size_t GetAttentionWorkspaceSize( int num_heads, int head_size, int sequence_length, - int past_sequence_length) { + int total_sequence_length) { size_t qkv_size = element_size * 3 * batch_size * sequence_length * num_heads * head_size; return qkv_size + 2 * GetAttentionScratchSize(element_size, batch_size, num_heads, - sequence_length, past_sequence_length + sequence_length); + sequence_length, total_sequence_length); } inline int3 Get2DMaskStrides(int total_sequence_length) { diff --git a/onnxruntime/contrib_ops/rocm/bert/attention_softmax.h b/onnxruntime/contrib_ops/rocm/bert/attention_softmax.h index 8a34e29e10..dc0cdb4d6a 100644 --- a/onnxruntime/contrib_ops/rocm/bert/attention_softmax.h +++ b/onnxruntime/contrib_ops/rocm/bert/attention_softmax.h @@ -38,7 +38,6 @@ namespace rocm { template __device__ inline void Softmax(const int all_sequence_length, - const int sequence_length, const int valid_end, const int valid_start, const T* add_before_softmax, @@ -276,10 +275,8 @@ __global__ void SoftmaxKernelSmall(const int all_sequence_length, const int sequ } template -__global__ void SoftmaxKernel(const int all_sequence_length, const int sequence_length, - const T* add_before_softmax, const T* input, T* output) { - Softmax(all_sequence_length, sequence_length, all_sequence_length, 0, - add_before_softmax, input, output); +__global__ void SoftmaxKernel(const int all_sequence_length, const T* add_before_softmax, const T* input, T* output) { + Softmax(all_sequence_length, all_sequence_length, 0, add_before_softmax, input, output); } template @@ -315,7 +312,7 @@ Status ComputeSoftmax( } else if (!is_unidirectional) { const int blockSize = 1024; SoftmaxKernel<<>>( - all_sequence_length, sequence_length, add_before_softmax, input, output); + all_sequence_length, add_before_softmax, input, output); } else { return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Attention ROCM operator does not support total sequence length > 1024."); } @@ -349,8 +346,7 @@ __global__ void MaskedSoftmaxKernelSmall(const int all_sequence_length, const in } template -__global__ void MaskedSoftmaxKernel(const int all_sequence_length, const int sequence_length, - const int* mask_end, const int* mask_start, +__global__ void MaskedSoftmaxKernel(const int all_sequence_length, const int* mask_end, const int* mask_start, const T* add_before_softmax, const T* input, T* output) { __shared__ int start_position; __shared__ int end_position; @@ -368,8 +364,7 @@ __global__ void MaskedSoftmaxKernel(const int all_sequence_length, const int seq } __syncthreads(); - Softmax(all_sequence_length, sequence_length, end_position, start_position, - add_before_softmax, input, output); + Softmax(all_sequence_length, end_position, start_position, add_before_softmax, input, output); } template @@ -400,7 +395,7 @@ Status ComputeSoftmaxWithMask1D( } else if (!is_unidirectional) { const int blockSize = 1024; MaskedSoftmaxKernel<<>>( - all_sequence_length, sequence_length, mask_index, mask_start, + all_sequence_length, mask_index, mask_start, add_before_softmax, input, output); } else { return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Attention ROCM operator does not support total sequence length > 1024."); diff --git a/onnxruntime/contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_ck_impl/impl.cuh b/onnxruntime/contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_ck_impl/impl.cuh new file mode 100644 index 0000000000..c127a1089a --- /dev/null +++ b/onnxruntime/contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_ck_impl/impl.cuh @@ -0,0 +1,146 @@ +// SPDX-License-Identifier: MIT +// Modifications Copyright (c) Microsoft Corporation. All rights reserved. +// Copyright (c) 2018-2022, Advanced Micro Devices, Inc. All rights reserved. + +#pragma once + +#include + +#include "ck/ck.hpp" +#include "ck/tensor_operation/gpu/device/tensor_layout.hpp" +#include "ck/tensor_operation/gpu/device/tensor_specialization.hpp" +#include "ck/tensor_operation/gpu/device/device_batched_gemm_softmax_gemm_permute.hpp" +#include "ck/tensor_operation/gpu/device/gemm_specialization.hpp" +#include "ck/tensor_operation/gpu/device/impl/device_batched_gemm_softmax_gemm_permute_xdl_cshuffle.hpp" +#include "ck/tensor_operation/gpu/element/unary_element_wise_operation.hpp" +#include "ck/utility/data_type.hpp" + +namespace onnxruntime { +namespace contrib { +namespace rocm { +namespace internal { + +using F16 = ck::half_t; +using F32 = float; + +template +using S = ck::Sequence; + +using MaskingSpecialization = ck::tensor_operation::device::MaskingSpecialization; + +using PassThrough = ck::tensor_operation::element_wise::PassThrough; + +using ck::tensor_operation::device::DeviceBatchedGemmSoftmaxGemmPermute; // the interface +using ck::tensor_operation::device::DeviceBatchedGemmSoftmaxGemmPermute_Xdl_CShuffle; // the implementation + +static constexpr auto GemmDefault = ck::tensor_operation::device::GemmSpecialization::Default; +static constexpr auto GemmPadded = ck::tensor_operation::device::GemmSpecialization::MNKOPadding; + +static constexpr auto TensorDefault = ck::tensor_operation::device::TensorSpecialization::Default; + +template +using device_batched_gemm_softmax_gemm_permute_instances = + std::tuple< + // clang-format off + // #############################################| NumDimG| NumDimM| NumDimN| NumDimK| NumDimO| AData| B0Data| B1Data| CData| Acc0BiasData| Acc1BiasData| AccData| CShuffle| A| B0| Acc0| B1| C| GEMM| ATensorSpec| B0TensorSpec| B1TensorSpec| CTensorSpec| NumGemmK| Block| Gemm01| Gemm0| Gemm0| Gemm1| Gemm1| AK1| BK1| B1K1| MPer| NPer| Gemm0| Gemm0| Gemm1| ABlockTransfer| ABlockTransfer| ABlockTransfer| ABlockTransfer| ABlockTransfer| ABlockTransfer| ABlockLds| B0BlockTransfer| B0BlockTransfer| B0BlockTransfer| B0BlockTransfer| B0BlockTransfer| B0BlockTransfer| B0BlockLds| B1BlockTransfer| B1BlockTransfer| B1BlockTransfer| B1BlockTransfer| B1BlockTransfer| B1BlockTransfer| B1BlockLds| CShuffle| CShuffle| CBlockTransferClusterLengths| CBlockTransfer| MaskingSpec| + // #############################################| | | | | | Type| Type| Type| Type| Type| Type| Type| DataType| Elementwise| Elementwise| Elementwise| Elementwise| Elementwise| Specialization| | | | | Prefetch| Size| MPer| NPer| KPer| NPer| KPer| | | | XDL| XDL| MXdl| NXdl| NXdl| ThreadCluster| ThreadCluster| SrcAccessOrder| SrcVectorDim| SrcScalar| DstScalar| AddExtraM| ThreadCluster| ThreadCluster| SrcAccessOrder| SrcVectorDim| SrcScalar| DstScalar| AddExtraN| ThreadCluster| ThreadCluster| SrcAccessOrder| SrcVectorDim| SrcScalar| DstScalar| AddExtraN| MXdlPerWave| NXdlPerWave| _MBlock_MWaveMPerXdl| ScalarPerVector| | + // #############################################| | | | | | | | | | | | | | Operation| Operation| Operation| Operation| Operation| | | | | | Stage| | Block| Block| Block| Block| Block| | | | | | Per| Per| Per| Lengths_K0_M_K1| ArrangeOrder| | | PerVector| PerVector_K1| | Lengths_K0_N_K1| ArrangeOrder| | | PerVector| PerVector_K1| | Lengths_K0_N_K1| ArrangeOrder| | | PerVector| PerVector_K1| | PerShuffle| PerShuffle| _NBlock_NWaveNPerXdl| _NWaveNPerXdl| | + // #############################################| | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | Wave| Wave| Wave| | | | | | | | | | | | | | | | | | | | | | | | | | | + DeviceBatchedGemmSoftmaxGemmPermute_Xdl_CShuffle< NumDimG, NumDimM, NumDimN, NumDimK, NumDimO, DT, DT, DT, DT, D0sDT, ck::Tuple<>, AccDT, DT, PassThrough, PassThrough, D0Op, PassThrough, PassThrough, GemmDefault, TensorDefault, TensorDefault, TensorDefault, TensorDefault, 1, 256, 256, 128, 32, 64, 32, 8, 8, 2, 32, 32, 2, 4, 2, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S<16, 16, 1>, S<0, 2, 1>, S<0, 2, 1>, 1, 4, 2, false, 1, 2, S<1, 32, 1, 8>, 8, MaskingSpec>, + DeviceBatchedGemmSoftmaxGemmPermute_Xdl_CShuffle< NumDimG, NumDimM, NumDimN, NumDimK, NumDimO, DT, DT, DT, DT, D0sDT, ck::Tuple<>, AccDT, DT, PassThrough, PassThrough, D0Op, PassThrough, PassThrough, GemmDefault, TensorDefault, TensorDefault, TensorDefault, TensorDefault, 1, 256, 256, 128, 32, 128, 32, 8, 8, 2, 32, 32, 2, 4, 4, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S< 8, 32, 1>, S<0, 2, 1>, S<0, 2, 1>, 1, 4, 2, false, 1, 2, S<1, 32, 1, 8>, 8, MaskingSpec>, +#if ROCM_VERSION >= 50500 + DeviceBatchedGemmSoftmaxGemmPermute_Xdl_CShuffle< NumDimG, NumDimM, NumDimN, NumDimK, NumDimO, DT, DT, DT, DT, D0sDT, ck::Tuple<>, AccDT, DT, PassThrough, PassThrough, D0Op, PassThrough, PassThrough, GemmDefault, TensorDefault, TensorDefault, TensorDefault, TensorDefault, 1, 256, 128, 256, 32, 64, 32, 8, 8, 2, 32, 32, 1, 8, 2, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S<16, 16, 1>, S<0, 2, 1>, S<0, 2, 1>, 1, 4, 2, false, 1, 2, S<1, 32, 1, 8>, 8, MaskingSpec>, +#endif + DeviceBatchedGemmSoftmaxGemmPermute_Xdl_CShuffle< NumDimG, NumDimM, NumDimN, NumDimK, NumDimO, DT, DT, DT, DT, D0sDT, ck::Tuple<>, AccDT, DT, PassThrough, PassThrough, D0Op, PassThrough, PassThrough, GemmDefault, TensorDefault, TensorDefault, TensorDefault, TensorDefault, 1, 256, 128, 256, 32, 128, 32, 8, 8, 2, 32, 32, 1, 8, 4, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S< 8, 32, 1>, S<0, 2, 1>, S<0, 2, 1>, 1, 4, 2, false, 1, 2, S<1, 32, 1, 8>, 8, MaskingSpec>, + DeviceBatchedGemmSoftmaxGemmPermute_Xdl_CShuffle< NumDimG, NumDimM, NumDimN, NumDimK, NumDimO, DT, DT, DT, DT, D0sDT, ck::Tuple<>, AccDT, DT, PassThrough, PassThrough, D0Op, PassThrough, PassThrough, GemmDefault, TensorDefault, TensorDefault, TensorDefault, TensorDefault, 1, 256, 128, 128, 64, 64, 32, 8, 8, 2, 32, 32, 1, 4, 2, S<8, 32, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, false, S<8, 32, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, false, S<16, 16, 1>, S<0, 2, 1>, S<0, 2, 1>, 1, 4, 2, false, 1, 2, S<1, 32, 1, 8>, 8, MaskingSpec>, + DeviceBatchedGemmSoftmaxGemmPermute_Xdl_CShuffle< NumDimG, NumDimM, NumDimN, NumDimK, NumDimO, DT, DT, DT, DT, D0sDT, ck::Tuple<>, AccDT, DT, PassThrough, PassThrough, D0Op, PassThrough, PassThrough, GemmDefault, TensorDefault, TensorDefault, TensorDefault, TensorDefault, 1, 256, 128, 128, 32, 64, 32, 8, 8, 2, 32, 32, 1, 4, 2, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S<16, 16, 1>, S<0, 2, 1>, S<0, 2, 1>, 1, 4, 2, false, 1, 2, S<1, 32, 1, 8>, 8, MaskingSpec>, + DeviceBatchedGemmSoftmaxGemmPermute_Xdl_CShuffle< NumDimG, NumDimM, NumDimN, NumDimK, NumDimO, DT, DT, DT, DT, D0sDT, ck::Tuple<>, AccDT, DT, PassThrough, PassThrough, D0Op, PassThrough, PassThrough, GemmDefault, TensorDefault, TensorDefault, TensorDefault, TensorDefault, 1, 256, 128, 128, 64, 128, 32, 8, 8, 2, 32, 32, 1, 4, 4, S<8, 32, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, false, S<8, 32, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, false, S< 8, 32, 1>, S<0, 2, 1>, S<0, 2, 1>, 1, 4, 2, false, 1, 2, S<1, 32, 1, 8>, 8, MaskingSpec>, + DeviceBatchedGemmSoftmaxGemmPermute_Xdl_CShuffle< NumDimG, NumDimM, NumDimN, NumDimK, NumDimO, DT, DT, DT, DT, D0sDT, ck::Tuple<>, AccDT, DT, PassThrough, PassThrough, D0Op, PassThrough, PassThrough, GemmDefault, TensorDefault, TensorDefault, TensorDefault, TensorDefault, 1, 256, 128, 128, 32, 128, 32, 8, 8, 2, 32, 32, 1, 4, 4, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S< 8, 32, 1>, S<0, 2, 1>, S<0, 2, 1>, 1, 4, 2, false, 1, 2, S<1, 32, 1, 8>, 8, MaskingSpec>, + DeviceBatchedGemmSoftmaxGemmPermute_Xdl_CShuffle< NumDimG, NumDimM, NumDimN, NumDimK, NumDimO, DT, DT, DT, DT, D0sDT, ck::Tuple<>, AccDT, DT, PassThrough, PassThrough, D0Op, PassThrough, PassThrough, GemmDefault, TensorDefault, TensorDefault, TensorDefault, TensorDefault, 1, 256, 64, 256, 32, 128, 32, 8, 8, 2, 16, 16, 1, 16, 8, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S< 8, 32, 1>, S<0, 2, 1>, S<0, 2, 1>, 1, 4, 2, false, 1, 8, S<1, 16, 1,16>, 8, MaskingSpec>, + DeviceBatchedGemmSoftmaxGemmPermute_Xdl_CShuffle< NumDimG, NumDimM, NumDimN, NumDimK, NumDimO, DT, DT, DT, DT, D0sDT, ck::Tuple<>, AccDT, DT, PassThrough, PassThrough, D0Op, PassThrough, PassThrough, GemmDefault, TensorDefault, TensorDefault, TensorDefault, TensorDefault, 1, 256, 64, 256, 32, 64, 32, 8, 8, 2, 16, 16, 1, 16, 4, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S<16, 16, 1>, S<0, 2, 1>, S<0, 2, 1>, 1, 4, 2, false, 1, 4, S<1, 32, 1, 8>, 8, MaskingSpec>, + DeviceBatchedGemmSoftmaxGemmPermute_Xdl_CShuffle< NumDimG, NumDimM, NumDimN, NumDimK, NumDimO, DT, DT, DT, DT, D0sDT, ck::Tuple<>, AccDT, DT, PassThrough, PassThrough, D0Op, PassThrough, PassThrough, GemmDefault, TensorDefault, TensorDefault, TensorDefault, TensorDefault, 1, 256, 64, 256, 64, 128, 32, 8, 8, 2, 16, 16, 1, 16, 8, S<8, 32, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S<8, 32, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S< 8, 32, 1>, S<0, 2, 1>, S<0, 2, 1>, 1, 4, 2, false, 1, 8, S<1, 16, 1,16>, 8, MaskingSpec>, + DeviceBatchedGemmSoftmaxGemmPermute_Xdl_CShuffle< NumDimG, NumDimM, NumDimN, NumDimK, NumDimO, DT, DT, DT, DT, D0sDT, ck::Tuple<>, AccDT, DT, PassThrough, PassThrough, D0Op, PassThrough, PassThrough, GemmDefault, TensorDefault, TensorDefault, TensorDefault, TensorDefault, 1, 256, 64, 256, 64, 64, 32, 8, 8, 2, 16, 16, 1, 16, 4, S<8, 32, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S<8, 32, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S<16, 16, 1>, S<0, 2, 1>, S<0, 2, 1>, 1, 4, 2, false, 1, 4, S<1, 32, 1, 8>, 8, MaskingSpec>, + // Padded fallback kernel + DeviceBatchedGemmSoftmaxGemmPermute_Xdl_CShuffle< NumDimG, NumDimM, NumDimN, NumDimK, NumDimO, DT, DT, DT, DT, D0sDT, ck::Tuple<>, AccDT, DT, PassThrough, PassThrough, D0Op, PassThrough, PassThrough, GemmPadded, TensorDefault, TensorDefault, TensorDefault, TensorDefault, 1, 256, 128, 128, 64, 128, 32, 8, 8, 2, 32, 32, 1, 4, 4, S<8, 32, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, false, S<8, 32, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, false, S< 8, 32, 1>, S<0, 2, 1>, S<0, 2, 1>, 1, 4, 2, false, 1, 2, S<1, 32, 1, 8>, 8, MaskingSpec>, + DeviceBatchedGemmSoftmaxGemmPermute_Xdl_CShuffle< NumDimG, NumDimM, NumDimN, NumDimK, NumDimO, DT, DT, DT, DT, D0sDT, ck::Tuple<>, AccDT, DT, PassThrough, PassThrough, D0Op, PassThrough, PassThrough, GemmPadded, TensorDefault, TensorDefault, TensorDefault, TensorDefault, 1, 256, 128, 64, 32, 128, 32, 8, 8, 2, 32, 32, 1, 2, 4, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S<4, 64, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 8, 8, true, S< 8, 32, 1>, S<0, 2, 1>, S<0, 2, 1>, 1, 4, 2, false, 1, 2, S<1, 32, 1, 8>, 8, MaskingSpec> + // clang-format on + >; + + +struct PreSoftmaxAttentionScoreOp { + PreSoftmaxAttentionScoreOp(float scale) : scale_(scale) {} + + // non-biased, non-masked + __host__ __device__ void operator()(float& y, const float& x) const { + y = scale_ * x; + } + + // biased or converted masked + __host__ __device__ void operator()(float& y, const float& x, const F16& bias) const { + y = scale_ * x + ck::type_convert(bias); + } + + // biased and converted masked + __host__ __device__ void operator()(float& y, const float& x, const F16& bias, const F16& converted_mask) const { + y = scale_ * x + ck::type_convert(bias) + ck::type_convert(converted_mask); + } + + float scale_; +}; + + +// Use this function to gat implementation +template +std::vector, + PassThrough, PassThrough, D0Op, PassThrough, PassThrough, + MaskingSpec>>> +GetDeviceBatchedGemmSoftmaxGemmPermuteInstances() { + return {}; +} + +// implemented in impl_{fp16,bf16}[_biased][_masked].cu +// fp16, non-biased, non-masked +template <> +std::vector, ck::Tuple<>, + PassThrough, PassThrough, PreSoftmaxAttentionScoreOp, PassThrough, PassThrough, + MaskingSpecialization::MaskDisabled>>> +GetDeviceBatchedGemmSoftmaxGemmPermuteInstances< + F16, ck::Tuple<>, F32, PreSoftmaxAttentionScoreOp, MaskingSpecialization::MaskDisabled>(); + +// fp16, biased, non-masked +template <> +std::vector, ck::Tuple<>, + PassThrough, PassThrough, PreSoftmaxAttentionScoreOp, PassThrough, PassThrough, + MaskingSpecialization::MaskDisabled>>> +GetDeviceBatchedGemmSoftmaxGemmPermuteInstances< + F16, ck::Tuple, F32, PreSoftmaxAttentionScoreOp, MaskingSpecialization::MaskDisabled>(); + +// fp16, biased, fp16 masked, basically, two bias +template <> +std::vector, ck::Tuple<>, + PassThrough, PassThrough, PreSoftmaxAttentionScoreOp, PassThrough, PassThrough, + MaskingSpecialization::MaskDisabled>>> +GetDeviceBatchedGemmSoftmaxGemmPermuteInstances< + F16, ck::Tuple, F32, PreSoftmaxAttentionScoreOp, MaskingSpecialization::MaskDisabled>(); + +} // namespace ck +} // namespace rocm +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_ck_impl/impl_fp16.cu b/onnxruntime/contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_ck_impl/impl_fp16.cu new file mode 100644 index 0000000000..9499eb84c8 --- /dev/null +++ b/onnxruntime/contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_ck_impl/impl_fp16.cu @@ -0,0 +1,37 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_ck_impl/impl.cuh" +#include "ck/library/tensor_operation_instance/add_device_operation_instance.hpp" +#include "ck/tensor_operation/gpu/device/impl/device_batched_gemm_softmax_gemm_permute_xdl_cshuffle.hpp" + +namespace onnxruntime { +namespace contrib { +namespace rocm { +namespace internal { + +using NonBiasedNonmasked = DeviceBatchedGemmSoftmaxGemmPermute< + 2, 1, 1, 1, 1, + F16, F16, F16, F16, ck::Tuple<>, ck::Tuple<>, + PassThrough, PassThrough, PreSoftmaxAttentionScoreOp, PassThrough, PassThrough, + MaskingSpecialization::MaskDisabled>; + +template <> +std::vector> +GetDeviceBatchedGemmSoftmaxGemmPermuteInstances< + F16, ck::Tuple<>, F32, PreSoftmaxAttentionScoreOp, MaskingSpecialization::MaskDisabled>() { + std::vector> instances; + ck::tensor_operation::device::instance::add_device_operation_instances( + instances, + device_batched_gemm_softmax_gemm_permute_instances< + 2, 1, 1, 1, 1, + F16, ck::Tuple<>, F32, PreSoftmaxAttentionScoreOp, + MaskingSpecialization::MaskDisabled>{}); + + return instances; +} + +} // namespace internal +} // namespace rocm +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_ck_impl/impl_fp16_biased.cu b/onnxruntime/contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_ck_impl/impl_fp16_biased.cu new file mode 100644 index 0000000000..fe1256b920 --- /dev/null +++ b/onnxruntime/contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_ck_impl/impl_fp16_biased.cu @@ -0,0 +1,37 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_ck_impl/impl.cuh" +#include "ck/library/tensor_operation_instance/add_device_operation_instance.hpp" +#include "ck/tensor_operation/gpu/device/impl/device_batched_gemm_softmax_gemm_permute_xdl_cshuffle.hpp" + +namespace onnxruntime { +namespace contrib { +namespace rocm { +namespace internal { + +using BiasedNonmasked = DeviceBatchedGemmSoftmaxGemmPermute< + 2, 1, 1, 1, 1, + F16, F16, F16, F16, ck::Tuple, ck::Tuple<>, + PassThrough, PassThrough, PreSoftmaxAttentionScoreOp, PassThrough, PassThrough, + MaskingSpecialization::MaskDisabled>; + +template <> +std::vector> +GetDeviceBatchedGemmSoftmaxGemmPermuteInstances< + F16, ck::Tuple, F32, PreSoftmaxAttentionScoreOp, MaskingSpecialization::MaskDisabled>() { + std::vector> instances; + ck::tensor_operation::device::instance::add_device_operation_instances( + instances, + device_batched_gemm_softmax_gemm_permute_instances< + 2, 1, 1, 1, 1, + F16, ck::Tuple, F32, PreSoftmaxAttentionScoreOp, + MaskingSpecialization::MaskDisabled>{}); + + return instances; +} + +} // namespace internal +} // namespace rocm +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_ck_impl/impl_fp16_biased_biased.cu b/onnxruntime/contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_ck_impl/impl_fp16_biased_biased.cu new file mode 100644 index 0000000000..a0d99d2e84 --- /dev/null +++ b/onnxruntime/contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_ck_impl/impl_fp16_biased_biased.cu @@ -0,0 +1,37 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_ck_impl/impl.cuh" +#include "ck/library/tensor_operation_instance/add_device_operation_instance.hpp" +#include "ck/tensor_operation/gpu/device/impl/device_batched_gemm_softmax_gemm_permute_xdl_cshuffle.hpp" + +namespace onnxruntime { +namespace contrib { +namespace rocm { +namespace internal { + +using BiasedNonmasked = DeviceBatchedGemmSoftmaxGemmPermute< + 2, 1, 1, 1, 1, + F16, F16, F16, F16, ck::Tuple, ck::Tuple<>, + PassThrough, PassThrough, PreSoftmaxAttentionScoreOp, PassThrough, PassThrough, + MaskingSpecialization::MaskDisabled>; + +template <> +std::vector> +GetDeviceBatchedGemmSoftmaxGemmPermuteInstances< + F16, ck::Tuple, F32, PreSoftmaxAttentionScoreOp, MaskingSpecialization::MaskDisabled>() { + std::vector> instances; + ck::tensor_operation::device::instance::add_device_operation_instances( + instances, + device_batched_gemm_softmax_gemm_permute_instances< + 2, 1, 1, 1, 1, + F16, ck::Tuple, F32, PreSoftmaxAttentionScoreOp, + MaskingSpecialization::MaskDisabled>{}); + + return instances; +} + +} // namespace internal +} // namespace rocm +} // namespace contrib +} // namespace onnxruntime diff --git a/onnxruntime/contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_pipelines.cuh b/onnxruntime/contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_pipelines.cuh index ece4cd8912..f501f20692 100644 --- a/onnxruntime/contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_pipelines.cuh +++ b/onnxruntime/contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_pipelines.cuh @@ -18,10 +18,10 @@ In QKV projection (prior to this pipeline): X --o--> K [B,T,N*H] ->Reshape-> [B,T,N,H] ->Permute0213-> [B,N,T,H] \-> V [B,T,N*H] ->Reshape-> [B,T,N,H] ->Permute0213-> [B,N,T,H] -pre_softmax_attn_scores = Q*K' = [B,N,S,H] * [BxNxTxH]' = [B,N,S,T] Batched GEMM1 -pre_softmax_attn_scores_masked = pre_softmax_attn_scores +? bias +? mask Add Bias, +? is optional -attn_scores = softmax(pre_softmax_attn_scores_masked * scale) = [B,N,S,T] Scale then Softmax -scaled_multi_head_attn = attn_scores * V = [B,N,S,T] * [B,N,T,H] = [B,N,S,H] Batched GEMM2 +pre_softmax_attn_scores = Q*K' = [B,N,S,H] * [BxNxTxH]' = [B,N,S,T] Batched GEMM1 +pre_softmax_attn_scores_masked = pre_softmax_attn_scores * scale +? bias +? mask Scale Add Bias, +? is optional +attn_scores = softmax(pre_softmax_attn_scores_masked) = [B,N,S,T] Softmax +scaled_multi_head_attn = attn_scores * V = [B,N,S,T] * [B,N,T,H] = [B,N,S,H] Batched GEMM2 Op outputs scaled_multi_head_attn: [B,N,S,H] ->Permute0213-> [B,S,N,H] ->Reshape-> [B,S,N*H] @@ -31,13 +31,46 @@ For the computing of pre_softmax_attn_scores +? mask +? bias: GemmSoftmaxGemmPermuteGenericPipeline handles it in specialized softmax. TODO: remove it! +CK in GemmSoftmaxGemmPermuteTunablePipeline + + Q*K' ---> scale ---> [B,N,S,T] -------+?--> masked + bias --------------> [B,N,S,T] --+?--/ +mask_2d ---> [B,T] ---> [B,1,1,T] -/ + + Q*K' ---> scale ---> [B,N,S,T] -------+?--> masked + bias --------------> [B,N,S,T] --+?--/ +mask_3d --> [B,S,T] --> [B,1,S,T] -/ + + Q*K' ---> scale ---> [B,N,S,T] -------+?--> masked + bias --------------> [B,N,S,T] --+?--/ +mask_4d -> [B,1,M,M] -> [B,1,S,T] -/ M is max_sequence_length from megatron, we will create a + **sub-view** from original mask buffer + +For CK implementation, there will be four cases combined: +non-biased, non-masked, no special processing. + biased, non-masked, no special processing, add the mask directly. +non-biased, masked, convert the mask to [B,1,1_or_S,T] and perform broadcast add with scaled Q*K'. + biased, masked, convert the mask to [B,1,1_or_S,T] and perform broadcast add with bias and scaled Q*K'. + +Broadcast add is not actually perform the broadcasting, just broadcast the load operation from memory. The impl details +are in composable kernels. The scale and add logic is performed via Acc0ElementOp + */ +#include "core/framework/tensor_shape.h" #include "core/providers/rocm/tunable/gemm.h" #include "core/providers/rocm/tunable/rocm_tunable.h" #include "contrib_ops/cpu/bert/attention_base.h" #include "contrib_ops/rocm/bert/attention_impl.h" #include "contrib_ops/rocm/bert/attention_softmax.h" +#include "contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_ck_impl/impl.cuh" + +#include "ck/ck.hpp" +#include "ck/tensor_operation/gpu/device/tensor_layout.hpp" +#include "ck/tensor_operation/gpu/element/element_wise_operation.hpp" + +#include +#include namespace blas = onnxruntime::rocm::tunable::blas; @@ -53,7 +86,7 @@ inline int3 Get2DMaskStrides(int total_sequence_length) { inline std::tuple GetRawMaskBufferAddrSizesAndStrides( const int* buffer, const AttentionParameters* attn) { const int* offseted_buffer{buffer}; // how to view the mask buffer - int3 sizes{-1, -1, -1}; // the logical shape of the view + int3 sizes{0, 0, 0}; // the logical shape of the view int3 strides{-1, -1, -1}; // the physical memory layout switch (attn->mask_type) { case MASK_NONE: @@ -74,6 +107,7 @@ inline std::tuple GetRawMaskBufferAddrSizesAndStrides( strides = {attn->max_sequence_length * attn->max_sequence_length, attn->max_sequence_length, 1}; break; default: + LOGS_DEFAULT(FATAL) << "unsupported mask type: " << attn->mask_type; throw std::runtime_error("unsupported mask type"); } return {offseted_buffer, sizes, strides}; @@ -82,8 +116,14 @@ inline std::tuple GetRawMaskBufferAddrSizesAndStrides( template struct GemmSoftmaxGemmPermuteParams : onnxruntime::rocm::tunable::OpParams { std::string Signature() const override { - auto [m, n, k, o, batch] = GetGemmsMNKOBatch(); - return MakeString("M", m, "_N", n, "_K", k, "_O", o, "_B", batch); + return MakeString( + "B", attention->batch_size, + "_S", attention->sequence_length, + "_T", attention->total_sequence_length, + "_N", attention->num_heads, + "_H", attention->head_size, + bias_buffer != nullptr ? "_B" : "_NB", + "_M", mask_index_dims.size()); } std::tuple GetGemmsMNKOBatch() const { @@ -111,10 +151,10 @@ struct GemmSoftmaxGemmPermuteParams : onnxruntime::rocm::tunable::OpParams { // optional, mask value const int* mask_index_buffer{nullptr}; - gsl::span mask_index_dims{}; + TensorShapeVector mask_index_dims{}; // optional, internal - T* workspace_buffer{nullptr}; + void* workspace_buffer{nullptr}; }; template @@ -130,7 +170,7 @@ struct GemmSoftmaxGemmPermuteGenericPipeline { params->attention->num_heads, params->attention->sequence_length, params->attention->total_sequence_length); - auto gemm1_out = params->workspace_buffer; + auto gemm1_out = reinterpret_cast(params->workspace_buffer); auto softmax_out = gemm1_out + (bytes / sizeof(T)); auto gemm2_out = softmax_out + (bytes / sizeof(T)); return {gemm1_out, softmax_out, gemm2_out}; @@ -143,7 +183,7 @@ struct GemmSoftmaxGemmPermuteGenericPipeline { attn->num_heads, attn->head_size, attn->sequence_length, - attn->past_sequence_length); + attn->total_sequence_length); } inline static Status Gemm1(const GemmSoftmaxGemmPermuteParams* params) { @@ -244,6 +284,247 @@ struct GemmSoftmaxGemmPermuteGenericPipeline { } }; +namespace { +template +struct DataTypeAdaptor { + using type = T; +}; + +template <> +struct DataTypeAdaptor { + using type = ck::half_t; +}; + +template <> +struct DataTypeAdaptor { + using type = ck::bhalf16_t; +}; + +} // namespace + +template +class GemmSoftmaxGemmPermuteTunableOp : public tunable::TunableOp> { + public: + GemmSoftmaxGemmPermuteTunableOp(); + + inline static bool IsSupportedMaskType(const AttentionParameters* attn) { + switch (attn->mask_type) { + case MASK_NONE: + case MASK_2D_DUMMY: + case MASK_2D_KEY_PADDING: + case MASK_3D_ATTENTION: + case MASK_4D_MEGATRON: + return true; + default: + return false; + } + } + + inline static size_t GetWorkspaceNumBytes(const AttentionParameters* attn) { + if (!IsSupportedMaskType(attn)) { + return 0; + } + auto [buffer, sizes, strides] = GetRawMaskBufferAddrSizesAndStrides(nullptr, attn); + return sizeof(T) * sizes.x * sizes.y * sizes.z; + } + + template + __global__ static void ConvertToFilledMaskValue( + T* __restrict__ out, + const int3 out_strides, + const int* __restrict__ mask_buffer, + const int3 mask_lengths, // [B,S,T] + const int3 mask_strides, + Converter cvt) { + const int64_t global_idx = blockDim.x * blockIdx.x + threadIdx.x; + if (global_idx >= mask_lengths.x * mask_lengths.y * CeilDiv(mask_lengths.z, VecSize)) { + return; + } + + const int tidx = (global_idx % CeilDiv(mask_lengths.z, VecSize)) * VecSize; + const int bs_idx = global_idx / CeilDiv(mask_lengths.z, VecSize); + const int sidx = bs_idx % mask_lengths.y; + const int bidx = bs_idx / mask_lengths.y; + + int64_t in_offset = mask_strides.x * bidx + mask_strides.y * sidx + mask_strides.z * tidx; + int64_t out_offset = out_strides.x * bidx + out_strides.y * sidx + out_strides.z * tidx; + + if (tidx + VecSize <= mask_lengths.z) { + using LoadT = const aligned_vector; + using StoreT = aligned_vector; + LoadT load = *reinterpret_cast(mask_buffer + in_offset); + StoreT store; + +#pragma unroll + for (int i = 0; i < VecSize; i++) { + store.val[i] = cvt(load.val[i]); + } + *reinterpret_cast(out + out_offset) = store; + } else { +#pragma unroll + for (int i = tidx; i < mask_lengths.z; i++) { + out[out_offset + i] = cvt(mask_buffer[in_offset + i]); + } + } + } + + static Status LaunchConvertToFilledMaskValue(const GemmSoftmaxGemmPermuteParams* params) { + constexpr const int kThreadPerBlock = 256; + constexpr const int kVecSize = 4; + + auto attn = params->attention; + auto [buffer, lengths, strides] = GetRawMaskBufferAddrSizesAndStrides(params->mask_index_buffer, attn); + int64_t total_threads = lengths.x * lengths.y * CeilDiv(lengths.z, kVecSize); + auto num_blocks = CeilDiv(total_threads, kThreadPerBlock); + + auto mask_filter_value = attn->mask_filter_value; + auto cvt = [=] __device__(int v) -> T { + return v == 1 ? 0 : mask_filter_value; + }; + + ConvertToFilledMaskValue<<Stream()>>>( + reinterpret_cast(params->workspace_buffer), {lengths.y * lengths.z, lengths.z, 1}, // out desc + buffer, lengths, strides, // mask desc + cvt); + + return HIP_CALL(hipGetLastError()); + } +}; + +template +auto GetCKGemmSoftmaxGemmPermuteTypeStringAndOps() { + constexpr const int kNumBiasBuffer = static_cast(USE_BIAS) + static_cast(USE_MASK); + + using Nop = ck::tensor_operation::element_wise::PassThrough; + using Acc0ElementOp = internal::PreSoftmaxAttentionScoreOp; + + using CKDataType = typename DataTypeAdaptor::type; + using D0DataType = typename ck::detail::tuple_concat< + std::conditional_t, ck::Tuple<>>, + std::conditional_t, ck::Tuple<>>>::type; + + constexpr static auto MaskingSpec = + ck::tensor_operation::device::MaskingSpecialization::MaskDisabled; + + std::vector>>> ret; + for (auto&& impl : internal::GetDeviceBatchedGemmSoftmaxGemmPermuteInstances< + CKDataType, D0DataType, internal::F32, internal::PreSoftmaxAttentionScoreOp, MaskingSpec>()) { + auto type_string = impl->GetTypeString(); + + auto invoker = impl->MakeInvokerPointer(); + auto op = [impl = std::move(impl), invoker = std::move(invoker)]( + const GemmSoftmaxGemmPermuteParams* params) -> Status { + if constexpr (USE_BIAS) { + TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF(params->bias_buffer == nullptr); + } else { + TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF(params->bias_buffer != nullptr); + } + if constexpr (USE_MASK) { + TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF( + !GemmSoftmaxGemmPermuteTunableOp::IsSupportedMaskType(params->attention)); + TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF(params->mask_index_buffer == nullptr); + } else { + TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF(params->mask_index_buffer != nullptr); + } + + auto attn = params->attention; + int G0 = attn->batch_size; + int G1 = attn->num_heads; + int M = attn->sequence_length; + int N = attn->total_sequence_length; + int K = attn->head_size; + int O = attn->v_head_size; + { + auto [m, n, k, o, batch] = params->GetGemmsMNKOBatch(); + ORT_ENFORCE(M == m && N == n && K == k && O == o && G0 * G1 == batch, "semantic mismatch"); + ORT_ENFORCE(K == O, "inner product dimension mismatch"); + } + + std::vector q_buffer_lengths = {G0, G1, M, K}; + std::vector q_buffer_strides = {G1 * M * K, M * K, K, 1}; + std::vector k_buffer_lengths = {G0, G1, N, K}; + std::vector k_buffer_strides = {G1 * N * K, N * K, K, 1}; // matrices are transposed + std::vector v_buffer_lengths = {G0, G1, O, N}; + std::vector v_buffer_strides = {G1 * N * O, N * O, 1, O}; + std::vector out_buffer_lengths = {G0, G1, M, O}; + std::vector out_buffer_strides = {M * G1 * O, O, G1 * O, 1}; // permute 0213 + + std::array bias_buffers{}; + std::array, kNumBiasBuffer> bias_lengths{}; + std::array, kNumBiasBuffer> bias_strides{}; + if constexpr (USE_BIAS) { + bias_buffers[0] = const_cast(params->bias_buffer); + bias_lengths[0] = {G0, G1, M, N}; // BN(G0*G1), S(M), T(N) + bias_strides[0] = {G1 * M * N, M * N, N, 1}; + } + if constexpr (USE_MASK) { + bias_buffers[kNumBiasBuffer - 1] = params->workspace_buffer; + bias_lengths[kNumBiasBuffer - 1] = {G0, G1, M, N}; // BN(G0*G1), S(M), T(N) + if (params->mask_index_dims.size() == 2) { // [B,T] + bias_strides[kNumBiasBuffer - 1] = {N, 0, 0, 1}; + } else if (params->mask_index_dims.size() == 3) { // [B,S,T] + bias_strides[kNumBiasBuffer - 1] = {M * N, 0, N, 1}; + } else if (params->mask_index_dims.size() == 4) { // [B,1,max_seq_len,max_seq_len] -->convert--> [B,S,T] + bias_strides[kNumBiasBuffer - 1] = {M * N, 0, N, 1}; + } else { + ORT_ENFORCE(false, "Unreachable"); + } + } + + auto arg = impl->MakeArgumentPointer( + params->q_buffer, params->k_buffer, params->v_buffer, params->out_buffer, + bias_buffers, // Gemm1 bias, as attention mask + {}, // Gemm2 bias + q_buffer_lengths, q_buffer_strides, + k_buffer_lengths, k_buffer_strides, + v_buffer_lengths, v_buffer_strides, + out_buffer_lengths, out_buffer_strides, + bias_lengths, bias_strides, + {}, + {}, + Nop{}, + Nop{}, + Acc0ElementOp{params->scale}, + Nop{}, + Nop{}); + + TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF(!impl->IsSupportedArgument(arg.get()), + impl->GetTypeString(), " does not support ", params->Signature()); + + if constexpr (USE_MASK) { + ORT_RETURN_IF_ERROR(GemmSoftmaxGemmPermuteTunableOp::LaunchConvertToFilledMaskValue(params)); + } + invoker->Run(arg.get(), StreamConfig{params->Stream()}); + return Status::OK(); + }; + ret.emplace_back(std::make_pair(std::move(type_string), std::move(op))); + } + return ret; +} + +template +GemmSoftmaxGemmPermuteTunableOp::GemmSoftmaxGemmPermuteTunableOp() { + this->RegisterOp([](const GemmSoftmaxGemmPermuteParams* params) { + return GemmSoftmaxGemmPermuteGenericPipeline::Run(params, false); + }); + + for (auto&& [_, op] : GetCKGemmSoftmaxGemmPermuteTypeStringAndOps()) { + this->RegisterOp(std::move(op)); + } + + for (auto&& [_, op] : GetCKGemmSoftmaxGemmPermuteTypeStringAndOps()) { + this->RegisterOp(std::move(op)); + } + + for (auto&& [_, op] : GetCKGemmSoftmaxGemmPermuteTypeStringAndOps()) { + this->RegisterOp(std::move(op)); + } + + for (auto&& [_, op] : GetCKGemmSoftmaxGemmPermuteTypeStringAndOps()) { + this->RegisterOp(std::move(op)); + } +} + } // namespace rocm } // namespace contrib } // namespace onnxruntime diff --git a/onnxruntime/core/providers/rocm/math/softmax_ck.cuh b/onnxruntime/core/providers/rocm/math/softmax_ck.cuh index 060415990a..710f6e540b 100644 --- a/onnxruntime/core/providers/rocm/math/softmax_ck.cuh +++ b/onnxruntime/core/providers/rocm/math/softmax_ck.cuh @@ -56,8 +56,8 @@ auto GetCKSoftmaxTypeStringAndOps() { auto invoker = impl->MakeInvokerPointer(); auto ck_softmax_op = [impl = std::move(impl), invoker = std::move(invoker)](const SoftmaxParams* params) -> Status { - AccDataType alpha{1.0f}; - AccDataType beta{0.0f}; + double alpha{1.0f}; + double beta{0.0f}; TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF( params->is_log_softmax, @@ -71,7 +71,7 @@ auto GetCKSoftmaxTypeStringAndOps() { std::vector reduce_dims{3}; auto nop = Nop{}; - auto arg = impl->MakeArgumentPointer(in_lengths, in_strides, reduce_dims, &alpha, &beta, + auto arg = impl->MakeArgumentPointer(in_lengths, in_strides, reduce_dims, alpha, beta, params->input, params->output, nop, nop); TUNABLE_OP_RETURN_UNSUPPORTED_ARGUMENT_IF(!impl->IsSupportedArgument(arg.get()), impl->GetTypeString(), " does not support ", params->Signature()); diff --git a/onnxruntime/python/tools/kernel_explorer/device_array.h b/onnxruntime/python/tools/kernel_explorer/device_array.h index e9e345bb3b..89784543e1 100644 --- a/onnxruntime/python/tools/kernel_explorer/device_array.h +++ b/onnxruntime/python/tools/kernel_explorer/device_array.h @@ -45,7 +45,7 @@ class DeviceArray { CALL_THROW(MEMCPY(device_.get(), host_, size_ * itemsize_, MEMCPY_HOST_TO_DEVICE)); } DeviceArray(const DeviceArray&) = default; - DeviceArray& operator=(DeviceArray&) = default; + DeviceArray& operator=(const DeviceArray&) = default; void UpdateHostNumpyArray() { CALL_THROW(MEMCPY(host_, device_.get(), size_ * itemsize_, MEMCPY_DEVICE_TO_HOST)); diff --git a/onnxruntime/python/tools/kernel_explorer/kernel_explorer.cc b/onnxruntime/python/tools/kernel_explorer/kernel_explorer.cc index 93508fcb62..1a2b0c702b 100644 --- a/onnxruntime/python/tools/kernel_explorer/kernel_explorer.cc +++ b/onnxruntime/python/tools/kernel_explorer/kernel_explorer.cc @@ -8,6 +8,7 @@ #include "python/tools/kernel_explorer/kernels/rocm/fast_gelu.h" #include "python/tools/kernel_explorer/kernels/rocm/gemm.h" #include "python/tools/kernel_explorer/kernels/rocm/gemm_fast_gelu.h" +#include "python/tools/kernel_explorer/kernels/rocm/gemm_softmax_gemm_permute.h" #include "python/tools/kernel_explorer/kernels/rocm/skip_layer_norm.h" #include "python/tools/kernel_explorer/kernels/rocm/softmax.h" @@ -26,6 +27,7 @@ PYBIND11_MODULE(_kernel_explorer, m) { InitSkipLayerNorm(m); InitGemmFastGelu(m); InitSoftmax(m); + InitGemmSoftmaxGemmPermute(m); #endif m.def("is_composable_kernel_available", []() { diff --git a/onnxruntime/python/tools/kernel_explorer/kernel_explorer_interface.h b/onnxruntime/python/tools/kernel_explorer/kernel_explorer_interface.h index 7233ef7d46..49d932101f 100644 --- a/onnxruntime/python/tools/kernel_explorer/kernel_explorer_interface.h +++ b/onnxruntime/python/tools/kernel_explorer/kernel_explorer_interface.h @@ -58,18 +58,22 @@ class IKernelExplorer { virtual ~IKernelExplorer() = default; protected: - TuningContextT* TuningContext() { - if (ep_ == nullptr) { + ExecutionProvider* GetEp() { + std::call_once(ep_create_once_, [this]() { ExecutionProviderInfo info{}; - ep_ = std::make_unique(info); - } + this->ep_ = std::make_unique(info); + }); + return ep_.get(); + } - return static_cast(ep_->GetTuningContext()); + TuningContextT* TuningContext() { + return static_cast(GetEp()->GetTuningContext()); } StreamT Stream() { return stream_; } private: + std::once_flag ep_create_once_; std::unique_ptr ep_{}; StreamT stream_{0}; int repeats_{100}; diff --git a/onnxruntime/python/tools/kernel_explorer/kernels/gemm_softmax_gemm_permute_test.py b/onnxruntime/python/tools/kernel_explorer/kernels/gemm_softmax_gemm_permute_test.py new file mode 100644 index 0000000000..2f7df713c8 --- /dev/null +++ b/onnxruntime/python/tools/kernel_explorer/kernels/gemm_softmax_gemm_permute_test.py @@ -0,0 +1,375 @@ +# ------------------------------------------------------------------------- +# Copyright (c) Microsoft Corporation. All rights reserved. +# Licensed under the MIT License. +# -------------------------------------------------------------------------- + + +import os +import sys +from dataclasses import dataclass +from itertools import product + +import kernel_explorer as ke +import numpy as np +import pytest +from utils import dtype_to_suffix, softmax + + +def multinormal_distribution(num_distribution, num_element_per_dist): + arrays = [] + for _ in range(num_distribution): + mean = np.random.rand() - 0.5 + std = np.random.rand() + 0.5 + arrays.append(np.random.normal(mean, std, (num_element_per_dist,))) + return np.array(arrays) + + +def get_ck_binding_name(dtype, biased: bool, masked: bool): + dtype_suffix = "_" + dtype_to_suffix(dtype) + ck_suffix = "" + if biased: + ck_suffix += "Biased" + if masked: + ck_suffix += "Masked" + ck_suffix += dtype_suffix + return "GemmSoftmaxGemmPermuteCK" + ck_suffix + + +dtypes = ["float16"] +batches = [1, 64] +seqlens = [128, 512] +total_seqlens = [128, 512] +num_heads = [8, 12] +head_sizes = [64] +biaseds = [False, True] +mask_dims = [0, 2, 3, 4] + + +def get_biased_id(biased): + return "biased" if biased else "nobias" + + +def get_mask_dim_id(dim): + if dim == 0: + return "nomask" + return f"mask_{dim}d" + + +def _test_gemm_softmax_gemm_permute( + f, dtype, batch, seqlen, total_seqlen, num_heads, head_size, biased, mask_dim, scale +): + v_head_size = head_size + q_shape = [batch, num_heads, seqlen, head_size] + k_shape = [batch, num_heads, total_seqlen, head_size] + v_shape = [batch, num_heads, total_seqlen, v_head_size] + out_shape = [batch, seqlen, num_heads, head_size] + + attn_bias = None + bias_shape = [batch, num_heads, seqlen, total_seqlen] if biased else None + + attn_mask = None + mask_shape = None + mask_shape_broadcasted = None + max_seqlen = None + if mask_dim != 0: + if mask_dim == 2: + mask_shape = [batch, total_seqlen] + mask_shape_broadcasted = [batch, 1, 1, total_seqlen] + elif mask_dim == 3: + mask_shape = [batch, seqlen, total_seqlen] + mask_shape_broadcasted = [batch, 1, seqlen, total_seqlen] + elif mask_dim == 4: + max_seqlen = ((seqlen - 1) // 1024 + 1) * 1024 # round up to multiple of 1024 + mask_shape = [batch, 1, max_seqlen, max_seqlen] + else: + raise ValueError + + np.random.seed(42) + Q = multinormal_distribution(np.prod(q_shape[:-1]), q_shape[-1]).reshape(q_shape).astype(np.float64) + K = multinormal_distribution(np.prod(k_shape[:-1]), k_shape[-1]).reshape(k_shape).astype(np.float64) + V = multinormal_distribution(np.prod(v_shape[:-1]), v_shape[-1]).reshape(v_shape).astype(np.float64) + if bias_shape is not None: + attn_bias = np.random.uniform(-0.5, 0.5, size=bias_shape) + if mask_shape is not None: + attn_mask = (np.random.randint(0, 100, size=mask_shape) < 95).astype(np.int32) + + pre_softmax_attn_scores = Q @ np.swapaxes(K, 2, 3) + pre_softmax_attn_scores = pre_softmax_attn_scores * scale + if attn_bias is not None: + pre_softmax_attn_scores = pre_softmax_attn_scores + attn_bias + if attn_mask is not None: + filter_value = -10000.0 + if mask_dim == 4: + # equivalent to past_sequence_length = max_sequence_length - seqlen + converted_mask = (1 - attn_mask[:, :, -seqlen:, :total_seqlen]) * filter_value + else: + converted_mask = (1 - attn_mask.reshape(mask_shape_broadcasted)) * filter_value + pre_softmax_attn_scores = pre_softmax_attn_scores + converted_mask + attn_scores = softmax(pre_softmax_attn_scores, axis=-1) + attn = attn_scores @ V + ref = np.swapaxes(attn, 2, 1) # permute 0213 + + out = np.empty(out_shape, dtype=dtype) + host_Q = Q.astype(dtype) + host_K = K.astype(dtype) + host_V = V.astype(dtype) + host_attn_bias = attn_bias.astype(dtype) if attn_bias is not None else None + dev_Q = ke.DeviceArray(host_Q) + dev_K = ke.DeviceArray(host_K) + dev_V = ke.DeviceArray(host_V) + dev_out = ke.DeviceArray(out) + dev_attn_bias = ke.DeviceArray(host_attn_bias) if host_attn_bias is not None else None + dev_attn_mask = ke.DeviceArray(attn_mask) if attn_mask is not None else None + + my_gemm_softmax_gemm_permute = f( + batch, + seqlen, + total_seqlen, + max_seqlen, + num_heads, + head_size, + mask_dim, + scale, + dev_Q, + dev_K, + dev_V, + dev_attn_bias, + dev_attn_mask, + dev_out, + ) + + print() # write an empty line in case pytest ... -s -v + failures = {} + for impl in my_gemm_softmax_gemm_permute.ListOps(): + if not my_gemm_softmax_gemm_permute.SelectOp(impl): + print("Unsupport", impl) + continue + print(" Support", impl) + + my_gemm_softmax_gemm_permute.Run() + dev_out.UpdateHostNumpyArray() + + try: + is_strict = int(os.environ.get("KERNEL_EXPLORER_STRICT_TEST", "0")) + if is_strict: + # NOTE: this will always fail, just for manual checking with: + # KERNEL_EXPLORER_STRICT_TEST=1 pytest ... -s -v + np.testing.assert_allclose(out, ref) + else: + is_zero_tol, atol, rtol = 1e-3, 1e-2, 1e-2 + not_close_to_zeros = np.abs(ref) > is_zero_tol + np.testing.assert_allclose(out[not_close_to_zeros], ref[not_close_to_zeros], atol=atol, rtol=rtol) + except Exception as err: + header = "*" * 30 + impl + "*" * 30 + print(header) + print(err) + print("*" * len(header)) + failures[impl] = str(err) + + if failures: + raise Exception(failures) + + +@pytest.mark.parametrize("mask_dim", mask_dims, ids=get_mask_dim_id) +@pytest.mark.parametrize("biased", biaseds, ids=get_biased_id) +@pytest.mark.parametrize("head_size", head_sizes) +@pytest.mark.parametrize("nhead", num_heads) +@pytest.mark.parametrize("total_seqlen", total_seqlens) +@pytest.mark.parametrize("seqlen", seqlens) +@pytest.mark.parametrize("batch", [16]) +@pytest.mark.parametrize("dtype", dtypes) +def test_gemm_softmax_gemm_permute_generic(dtype, batch, seqlen, total_seqlen, nhead, head_size, biased, mask_dim): + f = getattr(ke, "GemmSoftmaxGemmPermuteGeneric_" + dtype_to_suffix(dtype)) + scale = 1.0 / np.sqrt(head_size) + _test_gemm_softmax_gemm_permute(f, dtype, batch, seqlen, total_seqlen, nhead, head_size, biased, mask_dim, scale) + + +@pytest.mark.parametrize("mask_dim", mask_dims, ids=get_mask_dim_id) +@pytest.mark.parametrize("biased", biaseds, ids=get_biased_id) +@pytest.mark.parametrize("head_size", head_sizes) +@pytest.mark.parametrize("nhead", num_heads) +@pytest.mark.parametrize("total_seqlen", total_seqlens) +@pytest.mark.parametrize("seqlen", seqlens) +@pytest.mark.parametrize("batch", batches) +@pytest.mark.parametrize("dtype", dtypes) +def test_gemm_softmax_gemm_permute_ck(dtype, batch, seqlen, total_seqlen, nhead, head_size, biased, mask_dim): + f = getattr(ke, get_ck_binding_name(dtype, biased, mask_dim != 0)) + scale = 1.0 / np.sqrt(head_size) + _test_gemm_softmax_gemm_permute(f, dtype, batch, seqlen, total_seqlen, nhead, head_size, biased, mask_dim, scale) + + +@pytest.mark.parametrize("mask_dim", [2], ids=get_mask_dim_id) +@pytest.mark.parametrize("biased", [False], ids=get_biased_id) +@pytest.mark.parametrize("head_size", [64]) +@pytest.mark.parametrize("nhead", [8]) +@pytest.mark.parametrize("total_seqlen", [128]) +@pytest.mark.parametrize("seqlen", [64]) +@pytest.mark.parametrize("batch", [16]) +@pytest.mark.parametrize("dtype", ["float16"]) +def test_gemm_softmax_gemm_permute_tunable(dtype, batch, seqlen, total_seqlen, nhead, head_size, biased, mask_dim): + f = getattr(ke, "GemmSoftmaxGemmPermuteTunable_" + dtype_to_suffix(dtype)) + scale = 1.0 / np.sqrt(head_size) + _test_gemm_softmax_gemm_permute(f, dtype, batch, seqlen, total_seqlen, nhead, head_size, biased, mask_dim, scale) + + +@dataclass +class GemmSoftmaxGemmPermuteMetric(ke.ComputeMetric): + batch: int + seqlen: int + total_seqlen: int + num_heads: int + head_size: int + biased: bool + mask_dim: int + + def report(self): + bias_str = " biased" if self.biased else "" + mask_str = f" mask_{self.mask_dim}d" if self.mask_dim != 0 else "" + common = ( + f"{self.dtype} B={self.batch} S={self.seqlen} T={self.total_seqlen} " + + f"N={self.num_heads} H={self.head_size}{bias_str}{mask_str}, " + + f"{self.name}" + ) + if self.duration <= 0: + return "not supported " + common + + return f"{self.duration:>6.2f} us {self.tflops:>5.2f} tflops " + common + + +def profile_gemm_softmax_gemm_permute_func( + f, dtype, batch, seqlen, total_seqlen, num_heads, head_size, biased, mask_dim, scale +): + v_head_size = head_size + q_shape = [batch, num_heads, seqlen, head_size] + k_shape = [batch, num_heads, total_seqlen, head_size] + v_shape = [batch, num_heads, total_seqlen, v_head_size] + out_shape = [batch, seqlen, num_heads, head_size] + + attn_bias = None + bias_shape = [batch, num_heads, seqlen, total_seqlen] if biased else None + + attn_mask = None + mask_shape = None + max_seqlen = None + if mask_dim != 0: + if mask_dim == 2: + mask_shape = [batch, total_seqlen] + elif mask_dim == 3: + mask_shape = [batch, seqlen, total_seqlen] + elif mask_dim == 4: + max_seqlen = ((seqlen - 1) // 1024 + 1) * 1024 # round up to multiple of 1024 + mask_shape = [batch, 1, max_seqlen, max_seqlen] + else: + raise ValueError + + np.random.seed(42) + Q = multinormal_distribution(np.prod(q_shape[:-1]), q_shape[-1]).reshape(q_shape).astype(np.float64) + K = multinormal_distribution(np.prod(k_shape[:-1]), k_shape[-1]).reshape(k_shape).astype(np.float64) + V = multinormal_distribution(np.prod(v_shape[:-1]), v_shape[-1]).reshape(v_shape).astype(np.float64) + if bias_shape is not None: + attn_bias = np.random.uniform(-2, 2, size=bias_shape) + if mask_shape is not None: + attn_mask = (np.random.randint(0, 100, size=mask_shape) < 95).astype(np.int32) + + out = np.empty(out_shape, dtype=dtype) + host_Q = Q.astype(dtype) + host_K = K.astype(dtype) + host_V = V.astype(dtype) + host_attn_bias = attn_bias.astype(dtype) if attn_bias is not None else None + dev_Q = ke.DeviceArray(host_Q) + dev_K = ke.DeviceArray(host_K) + dev_V = ke.DeviceArray(host_V) + dev_out = ke.DeviceArray(out) + dev_attn_bias = ke.DeviceArray(host_attn_bias) if host_attn_bias is not None else None + dev_attn_mask = ke.DeviceArray(attn_mask) if attn_mask is not None else None + + my_gemm_softmax_gemm_permute = f( + batch, + seqlen, + total_seqlen, + max_seqlen, + num_heads, + head_size, + mask_dim, + scale, + dev_Q, + dev_K, + dev_V, + dev_attn_bias, + dev_attn_mask, + dev_out, + ) + + for impl in my_gemm_softmax_gemm_permute.ListOps(): + duration_ms = -1 + if my_gemm_softmax_gemm_permute.SelectOp(impl): + duration_ms = my_gemm_softmax_gemm_permute.Profile() + + m, n, k, o, gemm_batch = seqlen, total_seqlen, head_size, head_size, batch * num_heads + flops_per_batch = m * n * k * 2 + m * n * o * 2 + flops_count_bias_and_softmax = True # set to false to be aligned with ck + if flops_count_bias_and_softmax: + flops_per_batch += 2 * n + 1 + if flops_count_bias_and_softmax and attn_bias is not None: + flops_per_batch += m * n + if flops_count_bias_and_softmax and attn_mask is not None: + flops_per_batch += m * n + flops = flops_per_batch * gemm_batch + + ke.report( + GemmSoftmaxGemmPermuteMetric( + impl, dtype, duration_ms, flops, batch, seqlen, total_seqlen, num_heads, head_size, biased, mask_dim + ) + ) + + +def profile_with_args(dtype, batch, seqlen, total_seqlen, num_heads, head_size, biased, mask_dim, scale, *, sort=False): + with ke.benchmark(sort): + args = (dtype, batch, seqlen, total_seqlen, num_heads, head_size, biased, mask_dim, scale) + profile_gemm_softmax_gemm_permute_func( + getattr(ke, "GemmSoftmaxGemmPermuteGeneric_" + dtype_to_suffix(dtype)), *args + ) + profile_gemm_softmax_gemm_permute_func(getattr(ke, get_ck_binding_name(dtype, biased, mask_dim != 0)), *args) + profile_gemm_softmax_gemm_permute_func( + getattr(ke, "GemmSoftmaxGemmPermuteTunable_" + dtype_to_suffix(dtype)), *args + ) + + +def profile(): + for args in product(dtypes, batches, seqlens, total_seqlens, num_heads, head_sizes, biaseds, mask_dims): + profile_with_args(*args, scale=0.125, sort=True) + print() + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser() + group = parser.add_argument_group("profile with args") + group.add_argument("--sort", action="store_true") + group.add_argument("dtype", choices=dtypes) + group.add_argument("batch", type=int) + group.add_argument("seqlen", type=int) + group.add_argument("total_seqlen", type=int) + group.add_argument("num_heads", type=int) + group.add_argument("head_size", type=int) + group.add_argument("biased", type=int, choices=[0, 1], default=0) + group.add_argument("mask_dim", type=int, choices=[0, 2, 3, 4], default=2, help="0 for mask disabled") + group.add_argument("scale", type=float, default=None, help="default to 1.0/sqrt(head_size)") + + if len(sys.argv) == 1: + profile() + else: + args = parser.parse_args() + profile_with_args( + args.dtype, + args.batch, + args.seqlen, + args.total_seqlen, + args.num_heads, + args.head_size, + args.biased, + args.mask_dim, + args.scale, + sort=args.sort, + ) diff --git a/onnxruntime/python/tools/kernel_explorer/kernels/rocm/gemm_softmax_gemm_permute.cu b/onnxruntime/python/tools/kernel_explorer/kernels/rocm/gemm_softmax_gemm_permute.cu new file mode 100644 index 0000000000..85c3ce1ad7 --- /dev/null +++ b/onnxruntime/python/tools/kernel_explorer/kernels/rocm/gemm_softmax_gemm_permute.cu @@ -0,0 +1,298 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "python/tools/kernel_explorer/kernels/rocm/gemm_softmax_gemm_permute.h" + +#include "pybind11/stl.h" + +#include "contrib_ops/rocm/bert/batched_gemm_softmax_gemm_permute_pipelines.cuh" +#include "core/providers/rocm/tunable/rocm_tunable.h" +#include "python/tools/kernel_explorer/device_array.h" +#include "python/tools/kernel_explorer/kernel_explorer_interface.h" + +#include + +namespace py = pybind11; + +using namespace onnxruntime::contrib::rocm; + +namespace onnxruntime { + +template +class IGemmSoftmaxGemmPermuteKernelExplorer : public IKernelExplorer { + public: + IGemmSoftmaxGemmPermuteKernelExplorer( + int64_t batch, + int64_t seqlen, + int64_t total_seqlen, + std::optional max_seqlen, + int64_t num_heads, + int64_t head_size, + int64_t mask_dim, + double scale, + DeviceArray& Q, + DeviceArray& K, + DeviceArray& V, + std::optional& attn_bias, + std::optional& attn_mask, + DeviceArray& out) { + ROCBLAS_CALL_THROW(rocblas_create_handle(&rocblas_handle_)); + + attn_.batch_size = batch; + attn_.sequence_length = seqlen; + attn_.kv_sequence_length = seqlen; // NOTE: not used + attn_.past_sequence_length = 0; + attn_.original_past_sequence_length = 0; // NOTE: not used + attn_.total_sequence_length = total_seqlen; + attn_.max_sequence_length = 0; + attn_.hidden_size = num_heads * head_size; + attn_.head_size = head_size; + attn_.v_hidden_size = attn_.hidden_size; // Q,K,V hidden size must agree now + attn_.v_head_size = attn_.head_size; // Q,K,V hidden size must agree now + attn_.num_heads = num_heads; + attn_.is_unidirectional = false; + attn_.past_present_share_buffer = false; + attn_.do_rotary = false; + attn_.mask_filter_value = -10000.0f; + attn_.scale = scale; + if (mask_dim == 0) { + attn_.mask_type = contrib::MASK_NONE; + } else if (mask_dim == 2) { + attn_.mask_type = contrib::MASK_2D_KEY_PADDING; + } else if (mask_dim == 3) { + attn_.mask_type = contrib::MASK_3D_ATTENTION; + } else if (mask_dim == 4) { + attn_.mask_type = contrib::MASK_4D_MEGATRON; + } else { + ORT_ENFORCE(false, "mask type not supported"); + } + + device_prop = GetEp()->GetDeviceProp(); + + params_.tuning_ctx = TuningContext(); + params_.stream = Stream(); + params_.handle = rocblas_handle_; + params_.attention = &attn_; + params_.device_prop = &device_prop; + params_.scale = scale; + + params_.q_buffer = reinterpret_cast(Q.ptr()); + params_.k_buffer = reinterpret_cast(K.ptr()); + params_.v_buffer = reinterpret_cast(V.ptr()); + if (attn_bias.has_value()) { + params_.bias_buffer = reinterpret_cast(attn_bias->ptr()); + } + if (attn_mask.has_value()) { + params_.mask_index_buffer = reinterpret_cast(attn_mask->ptr()); + if (mask_dim == 2) { + params_.mask_index_dims = {batch, total_seqlen}; + } else if (mask_dim == 3) { + params_.mask_index_dims = {batch, seqlen, total_seqlen}; + } else if (mask_dim == 4) { + ORT_ENFORCE(max_seqlen.has_value()); + attn_.max_sequence_length = max_seqlen.value(); + ORT_ENFORCE(attn_.max_sequence_length >= seqlen); + attn_.past_sequence_length = attn_.max_sequence_length - seqlen; + params_.mask_index_dims = {batch, 1, attn_.max_sequence_length, attn_.max_sequence_length}; + } + } + params_.out_buffer = reinterpret_cast(out.ptr()); + } + + ~IGemmSoftmaxGemmPermuteKernelExplorer() { + ROCBLAS_CALL_THROW(rocblas_destroy_handle(rocblas_handle_)); + } + + void SetWorkspace(size_t num_bytes) { + void* ptr; + HIP_CALL_THROW(hipMalloc(&ptr, num_bytes)); + workspace_.reset(ptr, [](void* ptr) { HIP_CALL_THROW(hipFree(ptr)); }); + params_.workspace_buffer = reinterpret_cast(workspace_.get()); + } + + protected: + using ParamsT = contrib::rocm::GemmSoftmaxGemmPermuteParams; + rocblas_handle rocblas_handle_; + hipDeviceProp_t device_prop; + contrib::AttentionParameters attn_; + ParamsT params_; + std::shared_ptr workspace_; +}; + +// The pipeline composed from rocblas api calls and kernel launches. +template +class GemmSoftmaxGemmPermuteGeneric : public IGemmSoftmaxGemmPermuteKernelExplorer { + public: + GemmSoftmaxGemmPermuteGeneric( + int64_t batch, + int64_t seqlen, + int64_t total_seqlen, + std::optional max_seqlen, + int64_t num_heads, + int64_t head_size, + int64_t mask_dim, + double scale, + DeviceArray& Q, + DeviceArray& K, + DeviceArray& V, + std::optional& attn_bias, + std::optional& attn_mask, + DeviceArray& out) + : IGemmSoftmaxGemmPermuteKernelExplorer(batch, seqlen, total_seqlen, max_seqlen, + num_heads, head_size, mask_dim, scale, + Q, K, V, attn_bias, attn_mask, out) { + this->SetWorkspace(GemmSoftmaxGemmPermuteGenericPipeline::GetWorkspaceNumBytes(&this->attn_)); + } + + std::vector ListOps() const { + return {"Generic"}; + } + + bool SelectOp(const std::string&) { + return true; + } + + void Run() override { + ORT_THROW_IF_ERROR(GemmSoftmaxGemmPermuteGenericPipeline::Run( + &this->params_, /*use_persistent_softmax=*/false)); + } +}; + +template +class GemmSoftmaxGemmPermuteCK : public IGemmSoftmaxGemmPermuteKernelExplorer { + public: + GemmSoftmaxGemmPermuteCK( + int64_t batch, + int64_t seqlen, + int64_t total_seqlen, + std::optional max_seqlen, + int64_t num_heads, + int64_t head_size, + int64_t mask_dim, + double scale, + DeviceArray& Q, + DeviceArray& K, + DeviceArray& V, + std::optional& attn_bias, + std::optional& attn_mask, + DeviceArray& out) + : IGemmSoftmaxGemmPermuteKernelExplorer(batch, seqlen, total_seqlen, max_seqlen, + num_heads, head_size, mask_dim, scale, + Q, K, V, attn_bias, attn_mask, out) { + this->SetWorkspace(GemmSoftmaxGemmPermuteTunableOp::GetWorkspaceNumBytes(&this->attn_)); + + for (auto&& [ts, op] : GetCKGemmSoftmaxGemmPermuteTypeStringAndOps()) { + type_strings_.emplace_back(std::move(ts)); + ops_.emplace_back(std::move(op)); + } + } + + std::vector ListOps() const { + return type_strings_; + } + + bool SelectOp(const std::string& name) { + for (size_t i = 0; i < ops_.size(); i++) { + if (type_strings_[i] == name) { + selected_op_ = i; + Status status = ops_[i].IsSupported(&this->params_); + return status.IsOK(); + } + } + + ORT_THROW("Cannot find implementation ", name); + } + + void Run() override { + ORT_THROW_IF_ERROR(ops_[selected_op_](&this->params_)); + } + + private: + using ParamsT = typename IGemmSoftmaxGemmPermuteKernelExplorer::ParamsT; + using OpT = Op; + + std::vector ops_; + std::vector type_strings_; + size_t selected_op_{}; +}; + +// The pipeline composed from rocblas api calls and kernel launches. +template +class GemmSoftmaxGemmPermuteTunable : public IGemmSoftmaxGemmPermuteKernelExplorer { + public: + GemmSoftmaxGemmPermuteTunable( + int64_t batch, + int64_t seqlen, + int64_t total_seqlen, + std::optional max_seqlen, + int64_t num_heads, + int64_t head_size, + int64_t mask_dim, + double scale, + DeviceArray& Q, + DeviceArray& K, + DeviceArray& V, + std::optional& attn_bias, + std::optional& attn_mask, + DeviceArray& out) + : IGemmSoftmaxGemmPermuteKernelExplorer(batch, seqlen, total_seqlen, max_seqlen, + num_heads, head_size, mask_dim, scale, + Q, K, V, attn_bias, attn_mask, out) { + this->SetWorkspace(std::max( + GemmSoftmaxGemmPermuteGenericPipeline::GetWorkspaceNumBytes(&this->attn_), + GemmSoftmaxGemmPermuteTunableOp::GetWorkspaceNumBytes(&this->attn_))); + + this->params_.TuningContext()->EnableTunableOp(); + } + + std::vector ListOps() const { + return {"Tunable"}; + } + + bool SelectOp(const std::string&) { + return true; + } + + void Run() override { + ORT_THROW_IF_ERROR(GemmSoftmaxGemmPermuteTunableOp{}(&this->params_)); + } +}; + +#define REGISTER_COMMON(name, type, ...) \ + py::class_>(m, name) \ + .def(py::init, int64_t, int64_t, int64_t, \ + float, \ + DeviceArray&, \ + DeviceArray&, \ + DeviceArray&, \ + std::optional&, \ + std::optional&, \ + DeviceArray&>()) \ + .def("SetRepeats", &type<__VA_ARGS__>::SetRepeats) \ + .def("Run", &type<__VA_ARGS__>::Run) \ + .def("Profile", &type<__VA_ARGS__>::Profile) \ + .def("ListOps", &type<__VA_ARGS__>::ListOps) \ + .def("SelectOp", &type<__VA_ARGS__>::SelectOp); + +#define REGISTER_GENERIC(dtype) \ + REGISTER_COMMON("GemmSoftmaxGemmPermuteGeneric_" #dtype, GemmSoftmaxGemmPermuteGeneric, dtype) + +#define REGISTER_CK(dtype, biased, masked, mask_bias_suffix) \ + REGISTER_COMMON( \ + "GemmSoftmaxGemmPermuteCK" mask_bias_suffix "_" #dtype, GemmSoftmaxGemmPermuteCK, dtype, biased, masked) + +#define REGISTER_TUNABLE(dtype) \ + REGISTER_COMMON("GemmSoftmaxGemmPermuteTunable_" #dtype, GemmSoftmaxGemmPermuteTunable, dtype) + +void InitGemmSoftmaxGemmPermute(py::module m) { + REGISTER_GENERIC(half); + + REGISTER_CK(half, false, false, ""); + REGISTER_CK(half, true, false, "Biased"); + REGISTER_CK(half, false, true, "Masked"); + REGISTER_CK(half, true, true, "BiasedMasked"); + + REGISTER_TUNABLE(half); +} + +} // namespace onnxruntime diff --git a/onnxruntime/python/tools/kernel_explorer/kernels/rocm/gemm_softmax_gemm_permute.h b/onnxruntime/python/tools/kernel_explorer/kernels/rocm/gemm_softmax_gemm_permute.h new file mode 100644 index 0000000000..67dc16ca31 --- /dev/null +++ b/onnxruntime/python/tools/kernel_explorer/kernels/rocm/gemm_softmax_gemm_permute.h @@ -0,0 +1,14 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include + +namespace py = pybind11; + +namespace onnxruntime { + +void InitGemmSoftmaxGemmPermute(py::module mod); + +} diff --git a/onnxruntime/python/tools/kernel_explorer/kernels/softmax_test.py b/onnxruntime/python/tools/kernel_explorer/kernels/softmax_test.py index cd3998a282..83d9363a8b 100644 --- a/onnxruntime/python/tools/kernel_explorer/kernels/softmax_test.py +++ b/onnxruntime/python/tools/kernel_explorer/kernels/softmax_test.py @@ -11,7 +11,7 @@ from itertools import product import kernel_explorer as ke import numpy as np import pytest -from utils import dtype_to_bytes, dtype_to_suffix +from utils import dtype_to_bytes, dtype_to_suffix, softmax def get_test_sizes(): @@ -29,13 +29,6 @@ def dtype_to_funcs(dtype): return type_map[dtype] -def softmax(x, is_log_softmax): - x = x - np.max(x, axis=-1, keepdims=1) - if is_log_softmax: - return x - np.log(np.sum(np.exp(x), axis=-1, keepdims=1)) - return (np.exp(x)) / np.sum(np.exp(x), axis=-1, keepdims=1) - - def _test_softmax(batch_count, softmax_elements, is_log_softmax, dtype, func): np.random.seed(0) x = np.random.rand(batch_count, softmax_elements).astype(dtype) @@ -43,7 +36,7 @@ def _test_softmax(batch_count, softmax_elements, is_log_softmax, dtype, func): x_d = ke.DeviceArray(x) y_d = ke.DeviceArray(y) - y_ref = softmax(x, is_log_softmax) + y_ref = softmax(x, is_log_softmax=is_log_softmax) softmax_func = getattr(ke, func) softmax_op = softmax_func( diff --git a/onnxruntime/python/tools/kernel_explorer/kernels/utils.py b/onnxruntime/python/tools/kernel_explorer/kernels/utils.py index d4eb7d6bd9..b005dd98d5 100644 --- a/onnxruntime/python/tools/kernel_explorer/kernels/utils.py +++ b/onnxruntime/python/tools/kernel_explorer/kernels/utils.py @@ -93,3 +93,10 @@ def get_gemm_basic_sizes(full=True): # ck has various impls to be tested, use the full basic cases will result too many cases to test. # So we use a reduced combination here. return list(product([1, 4, 127, 133], [3, 16, 128], [3, 129, 1024])) + + +def softmax(x, *, is_log_softmax=False, axis=-1): + x = x - np.max(x, axis=axis, keepdims=1) + if is_log_softmax: + return x - np.log(np.sum(np.exp(x), axis=axis, keepdims=1)) + return (np.exp(x)) / np.sum(np.exp(x), axis=axis, keepdims=1)