From cee5231eff0e406de95c52311596bbcfb98edff0 Mon Sep 17 00:00:00 2001 From: Dmitri Smirnov Date: Mon, 24 Mar 2025 12:13:08 -0700 Subject: [PATCH 1/2] Fix debug CUDA build issue on Windows --- .../contrib_ops/cuda/bert/attention_qk.cu | 12 ++++-------- .../contrib_ops/cuda/bert/attention_qk.h | 18 ++++++++++++++++++ 2 files changed, 22 insertions(+), 8 deletions(-) diff --git a/onnxruntime/contrib_ops/cuda/bert/attention_qk.cu b/onnxruntime/contrib_ops/cuda/bert/attention_qk.cu index b81783377936f..bb69059150fb7 100644 --- a/onnxruntime/contrib_ops/cuda/bert/attention_qk.cu +++ b/onnxruntime/contrib_ops/cuda/bert/attention_qk.cu @@ -37,15 +37,11 @@ Status CopyQK(cudaStream_t stream, const int qk_size, const T* input, QK* output) { - if constexpr (std::is_same::value) { - CUDA_RETURN_IF_ERROR(cudaMemcpyAsync(output, input, qk_size * sizeof(QK), cudaMemcpyDeviceToDevice, stream)); - return Status::OK(); - } - const bool half2float = std::is_same::value && std::is_same::value; - const bool float2half = std::is_same::value && std::is_same::value; - ORT_ENFORCE(half2float || float2half); + constexpr const bool half2float = std::is_same::value && std::is_same::value; + constexpr const bool float2half = std::is_same::value && std::is_same::value; + static_assert(half2float || float2half, "This function supports either or "); - int block_size = 256; + constexpr const int block_size = 256; int num_blocks = (qk_size + block_size - 1) / block_size; ConvertAndCopyQK<<>>(qk_size, input, output); diff --git a/onnxruntime/contrib_ops/cuda/bert/attention_qk.h b/onnxruntime/contrib_ops/cuda/bert/attention_qk.h index 3dead308e7d17..69d5816eeff70 100644 --- a/onnxruntime/contrib_ops/cuda/bert/attention_qk.h +++ b/onnxruntime/contrib_ops/cuda/bert/attention_qk.h @@ -18,6 +18,24 @@ Status CopyQK(cudaStream_t stream, const T* input, QK* output); +template <> +Status CopyQK(cudaStream_t stream, + const int qk_size, + const float* input, + float* output) { + CUDA_RETURN_IF_ERROR(cudaMemcpyAsync(output, input, qk_size * sizeof(float), cudaMemcpyDeviceToDevice, stream)); + return Status::OK(); +} + +template <> +Status CopyQK(cudaStream_t stream, + const int qk_size, + const half* input, + half* output) { + CUDA_RETURN_IF_ERROR(cudaMemcpyAsync(output, input, qk_size * sizeof(half), cudaMemcpyDeviceToDevice, stream)); + return Status::OK(); +} + } // namespace cuda } // namespace contrib } // namespace onnxruntime From 63c1eed8726f15a172e7c36d14db7120d06a591b Mon Sep 17 00:00:00 2001 From: Dmitri Smirnov Date: Mon, 24 Mar 2025 12:25:19 -0700 Subject: [PATCH 2/2] Adjust CopyQK() specialiazation --- .../contrib_ops/cuda/bert/attention_qk.cu | 26 ++++++++++++------- .../contrib_ops/cuda/bert/attention_qk.h | 18 ------------- 2 files changed, 17 insertions(+), 27 deletions(-) diff --git a/onnxruntime/contrib_ops/cuda/bert/attention_qk.cu b/onnxruntime/contrib_ops/cuda/bert/attention_qk.cu index bb69059150fb7..3f02a441da73e 100644 --- a/onnxruntime/contrib_ops/cuda/bert/attention_qk.cu +++ b/onnxruntime/contrib_ops/cuda/bert/attention_qk.cu @@ -48,11 +48,6 @@ Status CopyQK(cudaStream_t stream, return CUDA_CALL(cudaGetLastError()); } -template Status CopyQK(cudaStream_t stream, - const int qk_size, - const float* input, - float* output); - template Status CopyQK(cudaStream_t stream, const int qk_size, const float* input, @@ -63,10 +58,23 @@ template Status CopyQK(cudaStream_t stream, const half* input, float* output); -template Status CopyQK(cudaStream_t stream, - const int qk_size, - const half* input, - half* output); +template <> +Status CopyQK(cudaStream_t stream, + const int qk_size, + const float* input, + float* output) { + CUDA_RETURN_IF_ERROR(cudaMemcpyAsync(output, input, qk_size * sizeof(float), cudaMemcpyDeviceToDevice, stream)); + return Status::OK(); +} + +template <> +Status CopyQK(cudaStream_t stream, + const int qk_size, + const half* input, + half* output) { + CUDA_RETURN_IF_ERROR(cudaMemcpyAsync(output, input, qk_size * sizeof(half), cudaMemcpyDeviceToDevice, stream)); + return Status::OK(); +} } // namespace cuda } // namespace contrib diff --git a/onnxruntime/contrib_ops/cuda/bert/attention_qk.h b/onnxruntime/contrib_ops/cuda/bert/attention_qk.h index 69d5816eeff70..3dead308e7d17 100644 --- a/onnxruntime/contrib_ops/cuda/bert/attention_qk.h +++ b/onnxruntime/contrib_ops/cuda/bert/attention_qk.h @@ -18,24 +18,6 @@ Status CopyQK(cudaStream_t stream, const T* input, QK* output); -template <> -Status CopyQK(cudaStream_t stream, - const int qk_size, - const float* input, - float* output) { - CUDA_RETURN_IF_ERROR(cudaMemcpyAsync(output, input, qk_size * sizeof(float), cudaMemcpyDeviceToDevice, stream)); - return Status::OK(); -} - -template <> -Status CopyQK(cudaStream_t stream, - const int qk_size, - const half* input, - half* output) { - CUDA_RETURN_IF_ERROR(cudaMemcpyAsync(output, input, qk_size * sizeof(half), cudaMemcpyDeviceToDevice, stream)); - return Status::OK(); -} - } // namespace cuda } // namespace contrib } // namespace onnxruntime