diff --git a/onnxruntime/contrib_ops/cuda/bert/attention.cc b/onnxruntime/contrib_ops/cuda/bert/attention.cc index f5b344430d..25a23a5111 100644 --- a/onnxruntime/contrib_ops/cuda/bert/attention.cc +++ b/onnxruntime/contrib_ops/cuda/bert/attention.cc @@ -87,6 +87,7 @@ Status Attention::ComputeInternal(OpKernelContext* context) const { size_t workSpaceSize = GetAttentionWorkspaceSize(element_size, batch_size, num_heads_, head_size, sequence_length, past_sequence_length); auto temp_buffer = GetScratchBuffer(workSpaceSize); if (!LaunchAttentionKernel( + device_prop, reinterpret_cast(gemm_buffer.get()), nullptr == mask_index ? nullptr : mask_index->template Data(), nullptr == mask_index ? nullptr : &(mask_index->Shape().GetDims()), diff --git a/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu b/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu index 54206eb5fc..72634fa6d4 100644 --- a/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu +++ b/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu @@ -635,7 +635,7 @@ cublasStatus_t inline CublasGemmStridedBatched( template bool QkvToContext( - cublasHandle_t& cublas, cudaStream_t stream, + const cudaDeviceProp& prop, cublasHandle_t& cublas, cudaStream_t stream, const int batch_size, const int sequence_length, const int num_heads, const int head_size, const size_t element_size, const T* input, T* output, T* workspace, const int* mask_index, const std::vector* mask_index_dims, @@ -661,7 +661,7 @@ bool QkvToContext( const T* v = k + total_size; cublasSetStream(cublas, stream); - CublasMathModeSetter helper(cublas, CUBLAS_TENSOR_OP_MATH); + CublasMathModeSetter helper(prop, cublas, CUBLAS_TENSOR_OP_MATH); // Concat past (2xBxNxS'xH) to present (2xBxNxS*xH): // past_k (BxNxS'xH) + k (BxNxSxH) => present_k (BxNxS*xH) @@ -720,6 +720,7 @@ bool QkvToContext( } bool LaunchAttentionKernel( + const cudaDeviceProp& prop, const void* input, const int* mask_index, const std::vector* mask_index_dims, @@ -739,13 +740,13 @@ bool LaunchAttentionKernel( const cudaStream_t stream = nullptr; if (element_size == 2) { - return QkvToContext(cublas, stream, + return QkvToContext(prop, cublas, stream, batch_size, sequence_length, num_heads, head_size, element_size, reinterpret_cast(input), reinterpret_cast(output), reinterpret_cast(workspace), mask_index, mask_index_dims, is_unidirectional, past_sequence_length, reinterpret_cast(past), reinterpret_cast(present)); } else { - return QkvToContext(cublas, stream, + return QkvToContext(prop, cublas, stream, batch_size, sequence_length, num_heads, head_size, element_size, reinterpret_cast(input), reinterpret_cast(output), reinterpret_cast(workspace), mask_index, mask_index_dims, is_unidirectional, diff --git a/onnxruntime/contrib_ops/cuda/bert/attention_impl.h b/onnxruntime/contrib_ops/cuda/bert/attention_impl.h index 8a4ecffe4b..0ba287ee3c 100644 --- a/onnxruntime/contrib_ops/cuda/bert/attention_impl.h +++ b/onnxruntime/contrib_ops/cuda/bert/attention_impl.h @@ -16,6 +16,7 @@ size_t GetAttentionWorkspaceSize( int past_sequence_length); bool LaunchAttentionKernel( + const cudaDeviceProp& prop, // Device Properties const void* input, // Input tensor const int* mask_index, // Attention mask raw data or index (end position of each sequence, or end positions and start positions). NULL means no mask. const std::vector* mask_index_dims, // Mask index shape diff --git a/onnxruntime/contrib_ops/cuda/quantization/attention_quantization.cc b/onnxruntime/contrib_ops/cuda/quantization/attention_quantization.cc index 3eef3615e1..6e76b343ca 100644 --- a/onnxruntime/contrib_ops/cuda/quantization/attention_quantization.cc +++ b/onnxruntime/contrib_ops/cuda/quantization/attention_quantization.cc @@ -171,6 +171,7 @@ Status QAttention::ComputeInternal(OpKernelContext* context) const { size_t workSpaceSize = GetAttentionWorkspaceSize(element_size, batch_size, num_heads_, head_size, sequence_length, past_sequence_length); auto temp_buffer = GetScratchBuffer(workSpaceSize); if (!LaunchAttentionKernel( + GetDeviceProp(), reinterpret_cast(gemm_buffer.get()), nullptr == mask_index ? nullptr : mask_index->template Data(), nullptr == mask_index ? nullptr : &(mask_index->Shape().GetDims()), diff --git a/onnxruntime/core/providers/cuda/cuda_common.h b/onnxruntime/core/providers/cuda/cuda_common.h index a24e52c9e2..e1a7148288 100644 --- a/onnxruntime/core/providers/cuda/cuda_common.h +++ b/onnxruntime/core/providers/cuda/cuda_common.h @@ -225,15 +225,20 @@ inline bool CalculateFdmStrides(gsl::span p, const std::vector= 7) { + cublasGetMathMode(handle, &mode_); + cublasSetMathMode(handle, mode); + } } ~CublasMathModeSetter() { - cublasSetMathMode(handle_, mode_); + if (prop_.major >= 7) { + cublasSetMathMode(handle_, mode_); + } } private: + const cudaDeviceProp& prop_; cublasHandle_t handle_; cublasMath_t mode_; }; diff --git a/onnxruntime/core/providers/cuda/math/einsum_utils/einsum_auxiliary_ops.cc b/onnxruntime/core/providers/cuda/math/einsum_utils/einsum_auxiliary_ops.cc index 64981689eb..b1da3135f9 100644 --- a/onnxruntime/core/providers/cuda/math/einsum_utils/einsum_auxiliary_ops.cc +++ b/onnxruntime/core/providers/cuda/math/einsum_utils/einsum_auxiliary_ops.cc @@ -60,7 +60,8 @@ Status MatMul(const T* input_1_data, const T* input_2_data, T* output_data, reinterpret_cast(output_data), static_cast(N), static_cast(output_stride), - static_cast(num_batches))); + static_cast(num_batches), + static_cast(einsum_cuda_assets)->cuda_ep_->GetDeviceProp())); return Status::OK(); } diff --git a/onnxruntime/core/providers/cuda/math/matmul.cc b/onnxruntime/core/providers/cuda/math/matmul.cc index 307dbb0e15..8f1fd30c5e 100644 --- a/onnxruntime/core/providers/cuda/math/matmul.cc +++ b/onnxruntime/core/providers/cuda/math/matmul.cc @@ -143,7 +143,8 @@ Status MatMul::ComputeInternal(OpKernelContext* ctx) const { reinterpret_cast(Y->template MutableData()), ldc, stride_C, - static_cast(batch_count))); + static_cast(batch_count), + device_prop)); return Status::OK(); } @@ -175,7 +176,8 @@ Status MatMul::ComputeInternal(OpKernelContext* ctx) const { &zero, output_arrays.GpuPtr(), ldc, - static_cast(helper.OutputOffsets().size()))); + static_cast(helper.OutputOffsets().size()), + device_prop)); return Status::OK(); } diff --git a/onnxruntime/core/providers/cuda/shared_inc/fpgeneric.h b/onnxruntime/core/providers/cuda/shared_inc/fpgeneric.h index 82e4fdc386..52cacf022c 100644 --- a/onnxruntime/core/providers/cuda/shared_inc/fpgeneric.h +++ b/onnxruntime/core/providers/cuda/shared_inc/fpgeneric.h @@ -18,49 +18,175 @@ // Generalize library calls to be use in template functions // gemm -inline cublasStatus_t cublasGemmHelper(cublasHandle_t handle, cublasOperation_t transa, cublasOperation_t transb, - int m, int n, int k, const float* alpha, const float* A, int lda, - const float* B, int ldb, const float* beta, float* C, int ldc, +inline cublasStatus_t cublasGemmHelper(cublasHandle_t handle, + cublasOperation_t transa, + cublasOperation_t transb, + int m, int n, int k, + const float* alpha, + const float* A, int lda, + const float* B, int ldb, + const float* beta, + float* C, int ldc, const cudaDeviceProp& /*prop*/) { - return cublasSgemm(handle, transa, transb, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc); + return cublasSgemm(handle, + transa, + transb, + m, n, k, + alpha, + A, lda, + B, ldb, + beta, + C, ldc); } -inline cublasStatus_t cublasGemmHelper(cublasHandle_t handle, cublasOperation_t transa, cublasOperation_t transb, - int m, int n, int k, const double* alpha, const double* A, int lda, - const double* B, int ldb, const double* beta, double* C, int ldc, +inline cublasStatus_t cublasGemmHelper(cublasHandle_t handle, + cublasOperation_t transa, + cublasOperation_t transb, + int m, int n, int k, + const double* alpha, + const double* A, int lda, + const double* B, int ldb, + const double* beta, + double* C, int ldc, const cudaDeviceProp& /*prop*/) { - return cublasDgemm(handle, transa, transb, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc); + return cublasDgemm(handle, + transa, + transb, + m, n, k, + alpha, + A, lda, + B, ldb, + beta, + C, ldc); } -inline cublasStatus_t cublasGemmHelper(cublasHandle_t handle, cublasOperation_t transa, cublasOperation_t transb, - int m, int n, int k, const half* alpha, const half* A, int lda, - const half* B, int ldb, const half* beta, half* C, int ldc, +inline cublasStatus_t cublasGemmHelper(cublasHandle_t handle, + cublasOperation_t transa, + cublasOperation_t transb, + int m, int n, int k, + const half* alpha, + const half* A, int lda, + const half* B, int ldb, + const half* beta, + half* C, int ldc, const cudaDeviceProp& prop) { -#ifndef ENABLE_TRAINING - // This does true FP16 computation which is slow for non-Volta GPUs - if (prop.major >= 7) { - onnxruntime::cuda::CublasMathModeSetter math_mode_setter(handle, CUBLAS_TENSOR_OP_MATH); - return cublasHgemm(handle, transa, transb, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc); - } -#else - ORT_UNUSED_PARAMETER(prop); -#endif + onnxruntime::cuda::CublasMathModeSetter math_mode_setter(prop, handle, CUBLAS_TENSOR_OP_MATH); - //This does pseudo FP16 computation (input/output in fp16, computation in fp32) +#ifdef ENABLE_TRAINING float h_a = onnxruntime::math::halfToFloat(*reinterpret_cast(alpha)); float h_b = onnxruntime::math::halfToFloat(*reinterpret_cast(beta)); - cublasSetMathMode(handle, CUBLAS_TENSOR_OP_MATH); - return cublasGemmEx(handle, transa, transb, m, n, k, &h_a, A, CUDA_R_16F, lda, B, CUDA_R_16F, ldb, &h_b, C, CUDA_R_16F, ldc, CUDA_R_32F, CUBLAS_GEMM_DFALT); + + // accumulating in FP32 + return cublasGemmEx(handle, + transa, + transb, + m, n, k, + &h_a, + A, CUDA_R_16F, lda, + B, CUDA_R_16F, ldb, + &h_b, + C, CUDA_R_16F, ldc, + CUDA_R_32F, + CUBLAS_GEMM_DEFAULT); +#else + // accumulating in FP16 + return cublasHgemm(handle, + transa, + transb, + m, n, k, + alpha, + A, lda, + B, ldb, + beta, + C, ldc); +#endif } // batched gemm -inline cublasStatus_t cublasGemmBatchedHelper(cublasHandle_t handle, cublasOperation_t transa, cublasOperation_t transb, int m, int n, int k, const float* alpha, const float* Aarray[], int lda, const float* Barray[], int ldb, const float* beta, float* Carray[], int ldc, int batch_count) { - return cublasSgemmBatched(handle, transa, transb, m, n, k, alpha, Aarray, lda, Barray, ldb, beta, Carray, ldc, batch_count); +inline cublasStatus_t cublasGemmBatchedHelper(cublasHandle_t handle, + cublasOperation_t transa, + cublasOperation_t transb, + int m, int n, int k, + const float* alpha, + const float* Aarray[], int lda, + const float* Barray[], int ldb, + const float* beta, + float* Carray[], int ldc, + int batch_count, + const cudaDeviceProp& /*prop*/) { + return cublasSgemmBatched(handle, + transa, + transb, + m, n, k, + alpha, + Aarray, lda, + Barray, ldb, + beta, + Carray, ldc, + batch_count); } -inline cublasStatus_t cublasGemmBatchedHelper(cublasHandle_t handle, cublasOperation_t transa, cublasOperation_t transb, int m, int n, int k, const double* alpha, const double* Aarray[], int lda, const double* Barray[], int ldb, const double* beta, double* Carray[], int ldc, int batch_count) { - return cublasDgemmBatched(handle, transa, transb, m, n, k, alpha, Aarray, lda, Barray, ldb, beta, Carray, ldc, batch_count); +inline cublasStatus_t cublasGemmBatchedHelper(cublasHandle_t handle, + cublasOperation_t transa, + cublasOperation_t transb, + int m, int n, int k, + const double* alpha, + const double* Aarray[], int lda, + const double* Barray[], int ldb, + const double* beta, + double* Carray[], int ldc, + int batch_count, + const cudaDeviceProp& /*prop*/) { + return cublasDgemmBatched(handle, + transa, + transb, + m, n, k, + alpha, + Aarray, lda, + Barray, ldb, + beta, + Carray, ldc, + batch_count); } -inline cublasStatus_t cublasGemmBatchedHelper(cublasHandle_t handle, cublasOperation_t transa, cublasOperation_t transb, int m, int n, int k, const half* alpha, const half* Aarray[], int lda, const half* Barray[], int ldb, const half* beta, half* Carray[], int ldc, int batch_count) { - cublasSetMathMode(handle, CUBLAS_TENSOR_OP_MATH); - return cublasHgemmBatched(handle, transa, transb, m, n, k, alpha, (const __half**)Aarray, lda, (const __half**)Barray, ldb, beta, (__half**)Carray, ldc, batch_count); +inline cublasStatus_t cublasGemmBatchedHelper(cublasHandle_t handle, + cublasOperation_t transa, + cublasOperation_t transb, + int m, int n, int k, + const half* alpha, + const half* Aarray[], int lda, + const half* Barray[], int ldb, + const half* beta, + half* Carray[], int ldc, + int batch_count, + const cudaDeviceProp& prop) { + onnxruntime::cuda::CublasMathModeSetter math_mode_setter(prop, handle, CUBLAS_TENSOR_OP_MATH); +#ifdef ENABLE_TRAINING + float h_a = onnxruntime::math::halfToFloat(*reinterpret_cast(alpha)); + float h_b = onnxruntime::math::halfToFloat(*reinterpret_cast(beta)); + + // accumulating in FP32 + return cublasGemmBatchedEx(handle, + transa, + transb, + m, n, k, + &h_a, + (const void**)Aarray, CUDA_R_16F, lda, + (const void**)Barray, CUDA_R_16F, ldb, + &h_b, + (void**)Carray, CUDA_R_16F, ldc, + batch_count, + CUDA_R_32F, + CUBLAS_GEMM_DEFAULT); +#else + // accumulating in FP16 + return cublasHgemmBatched(handle, + transa, + transb, + m, n, k, + alpha, + (const __half**)Aarray, lda, + (const __half**)Barray, ldb, + beta, + (__half**)Carray, ldc, + batch_count); +#endif } // strided batched gemm @@ -76,8 +202,18 @@ inline cublasStatus_t cublasGemmStridedBatchedHelper(cublasHandle_t handle, const float* beta, float* C, int ldc, long long int strideC, - int batch_count) { - return cublasSgemmStridedBatched(handle, transa, transb, m, n, k, alpha, A, lda, strideA, B, ldb, strideB, beta, C, ldc, strideC, batch_count); + int batch_count, + const cudaDeviceProp& /*prop*/) { + return cublasSgemmStridedBatched(handle, + transa, + transb, + m, n, k, + alpha, + A, lda, strideA, + B, ldb, strideB, + beta, + C, ldc, strideC, + batch_count); } inline cublasStatus_t cublasGemmStridedBatchedHelper(cublasHandle_t handle, @@ -92,8 +228,18 @@ inline cublasStatus_t cublasGemmStridedBatchedHelper(cublasHandle_t handle, const double* beta, double* C, int ldc, long long int strideC, - int batch_count) { - return cublasDgemmStridedBatched(handle, transa, transb, m, n, k, alpha, A, lda, strideA, B, ldb, strideB, beta, C, ldc, strideC, batch_count); + int batch_count, + const cudaDeviceProp& /*prop*/) { + return cublasDgemmStridedBatched(handle, + transa, + transb, + m, n, k, + alpha, + A, lda, strideA, + B, ldb, strideB, + beta, + C, ldc, strideC, + batch_count); } inline cublasStatus_t cublasGemmStridedBatchedHelper(cublasHandle_t handle, @@ -108,9 +254,38 @@ inline cublasStatus_t cublasGemmStridedBatchedHelper(cublasHandle_t handle, const __half* beta, __half* C, int ldc, long long int strideC, - int batch_count) { - cublasSetMathMode(handle, CUBLAS_TENSOR_OP_MATH); - return cublasHgemmStridedBatched(handle, transa, transb, m, n, k, alpha, A, lda, strideA, B, ldb, strideB, beta, C, ldc, strideC, batch_count); + int batch_count, + const cudaDeviceProp& prop) { + onnxruntime::cuda::CublasMathModeSetter math_mode_setter(prop, handle, CUBLAS_TENSOR_OP_MATH); +#ifdef ENABLE_TRAINING + float h_a = onnxruntime::math::halfToFloat(*reinterpret_cast(alpha)); + float h_b = onnxruntime::math::halfToFloat(*reinterpret_cast(beta)); + // accumulating in FP32 + return cublasGemmStridedBatchedEx(handle, + transa, + transb, + m, n, k, + &h_a, + A, CUDA_R_16F, lda, strideA, + B, CUDA_R_16F, ldb, strideB, + &h_b, + C, CUDA_R_16F, ldc, strideC, + batch_count, + CUDA_R_32F, + CUBLAS_GEMM_DEFAULT); +#else + // accumulating in FP16 + return cublasHgemmStridedBatched(handle, + transa, + transb, + m, n, k, + alpha, + A, lda, strideA, + B, ldb, strideB, + beta, + C, ldc, strideC, + batch_count); +#endif } // axpy