diff --git a/mindspore/lite/src/delegate/tensorrt/cuda_impl/cublas_utils.cc b/mindspore/lite/src/delegate/tensorrt/cuda_impl/cublas_utils.cc index 02d31bbe73a..64a18ac5613 100644 --- a/mindspore/lite/src/delegate/tensorrt/cuda_impl/cublas_utils.cc +++ b/mindspore/lite/src/delegate/tensorrt/cuda_impl/cublas_utils.cc @@ -27,8 +27,7 @@ void Cublas2DTranspose(const float *in_addr, float *out_addr, const int *params, } void CublasMM1Batch(const void *a_addr, const void *b_addr, void *c_addr, const int *params, - const cublasOperation_t *operations, const cudaDataType_t *data_types, - cublasComputeType_t type_compute, cublasHandle_t cublas_handle) { + const cublasOperation_t *operations, const cudaDataType *data_types, cublasHandle_t cublas_handle) { const int m = params[0]; const int n = params[1]; const int k = params[2]; @@ -37,12 +36,13 @@ void CublasMM1Batch(const void *a_addr, const void *b_addr, void *c_addr, const const int lda = (trans_a == CUBLAS_OP_N) ? k : m; const int ldb = (trans_b == CUBLAS_OP_N) ? n : k; const int ldc = n; - cudaDataType_t type_a = data_types[0]; - cudaDataType_t type_b = data_types[1]; - cudaDataType_t type_c = data_types[2]; + cudaDataType type_a = data_types[0]; + cudaDataType type_b = data_types[1]; + cudaDataType type_c = data_types[2]; + cudaDataType compute_type = data_types[3]; const float alpha = 1.0f; const float beta = 0.0f; CUBLAS_CHECK_VOID(cublasGemmEx(cublas_handle, trans_b, trans_a, n, m, k, &alpha, b_addr, type_b, ldb, a_addr, type_a, - lda, &beta, c_addr, type_c, ldc, type_compute, CUBLAS_GEMM_DEFAULT_TENSOR_OP)); + lda, &beta, c_addr, type_c, ldc, compute_type, CUBLAS_GEMM_DEFAULT_TENSOR_OP)); } } // namespace mindspore::lite diff --git a/mindspore/lite/src/delegate/tensorrt/cuda_impl/cublas_utils.h b/mindspore/lite/src/delegate/tensorrt/cuda_impl/cublas_utils.h index 0f67cd52270..27e74d38263 100644 --- a/mindspore/lite/src/delegate/tensorrt/cuda_impl/cublas_utils.h +++ b/mindspore/lite/src/delegate/tensorrt/cuda_impl/cublas_utils.h @@ -47,9 +47,8 @@ void Cublas2DTranspose(const float *in_addr, float *out_addr, const int *params, // a: m * k, b: k * n, c: m * n // params order: m, n, k // operations order: trans_a, trans_b -// data_types: type_a, type_b, type_c +// data_types: type_a, type_b, type_c, compute type void CublasMM1Batch(const void *a_addr, const void *b_addr, void *c_addr, const int *params, - const cublasOperation_t *operations, const cudaDataType_t *data_types, - cublasComputeType_t type_compute, cublasHandle_t cublas_handle); + const cublasOperation_t *operations, const cudaDataType *data_types, cublasHandle_t cublas_handle); } // namespace mindspore::lite #endif // MINDSPORE_LITE_SRC_DELEGATE_TENSORRT_CDUA_IMPL_CUBLAS_UTILS_H_ diff --git a/mindspore/lite/src/delegate/tensorrt/op/matmul_opt_plugin.cc b/mindspore/lite/src/delegate/tensorrt/op/matmul_opt_plugin.cc index ed2d6147897..5a9895ec0f1 100644 --- a/mindspore/lite/src/delegate/tensorrt/op/matmul_opt_plugin.cc +++ b/mindspore/lite/src/delegate/tensorrt/op/matmul_opt_plugin.cc @@ -42,8 +42,7 @@ int MatmulOptPlugin::enqueue(const nvinfer1::PluginTensorDesc *inputDesc, const const int mm_params[]{m, n, k}; const int trans_params[]{n, m}; if (desc_a.type == nvinfer1::DataType::kFLOAT && desc_b.type == nvinfer1::DataType::kFLOAT) { - CublasMM1Batch(inputs[0], inputs[1], outputs[0], mm_params, operations_, data_types_, type_compute_, - cublas_handle_); + CublasMM1Batch(inputs[0], inputs[1], outputs[0], mm_params, operations_, data_types_, cublas_handle_); } else { MS_LOG(ERROR) << layer_name_ << " input datatype needs check a: " << static_cast(desc_a.type) << ", b: " << static_cast(desc_a.type); @@ -76,12 +75,12 @@ void MatmulOptPlugin::configurePlugin(const nvinfer1::DynamicPluginTensorDesc *i bias_index_ = (nbInputs == INPUT_SIZE3) ? kBiasIndex : -1; operations_[0] = a_trans_ ? CUBLAS_OP_T : CUBLAS_OP_N; operations_[1] = b_trans_ ? CUBLAS_OP_T : CUBLAS_OP_N; - data_types_[0] = ConvertDataType(in[0].desc.type); // input a - data_types_[1] = ConvertDataType(in[1].desc.type); // input b - data_types_[kBiasIndex] = ConvertDataType(out[0].desc.type); // output c - type_compute_ = (in[0].desc.type == nvinfer1::DataType::kHALF || in[1].desc.type == nvinfer1::DataType::kHALF) - ? CUBLAS_COMPUTE_32F_FAST_16BF - : CUBLAS_COMPUTE_32F; + data_types_[0] = ConvertDataType(in[0].desc.type); // input a + data_types_[1] = ConvertDataType(in[1].desc.type); // input b + data_types_[THIRD_INPUT] = ConvertDataType(out[0].desc.type); // output c + data_types_[FOURTH_INPUT] = + (in[0].desc.type == nvinfer1::DataType::kHALF || in[1].desc.type == nvinfer1::DataType::kHALF) ? CUDA_R_16F + : CUDA_R_32F; } int MatmulOptPlugin::initialize() noexcept { diff --git a/mindspore/lite/src/delegate/tensorrt/op/matmul_opt_plugin.h b/mindspore/lite/src/delegate/tensorrt/op/matmul_opt_plugin.h index e8434123eb7..2d0f3e9a6bd 100644 --- a/mindspore/lite/src/delegate/tensorrt/op/matmul_opt_plugin.h +++ b/mindspore/lite/src/delegate/tensorrt/op/matmul_opt_plugin.h @@ -51,8 +51,7 @@ class MatmulOptPlugin : public TensorRTPlugin { int bias_index_{-1}; // -1 means no bias, otherwise should be 2 cublasHandle_t cublas_handle_{nullptr}; cublasOperation_t operations_[2]{CUBLAS_OP_N, CUBLAS_OP_N}; - cudaDataType_t data_types_[3]{CUDA_R_32F, CUDA_R_32F, CUDA_R_32F}; - cublasComputeType_t type_compute_; + cudaDataType data_types_[4]{CUDA_R_32F, CUDA_R_32F, CUDA_R_32F, CUDA_R_32F}; }; class MatmulOptPluginCreater : public TensorRTPluginCreater { diff --git a/mindspore/lite/src/delegate/tensorrt/tensorrt_utils.cc b/mindspore/lite/src/delegate/tensorrt/tensorrt_utils.cc index f1bb1173e3e..1ea02ceb974 100644 --- a/mindspore/lite/src/delegate/tensorrt/tensorrt_utils.cc +++ b/mindspore/lite/src/delegate/tensorrt/tensorrt_utils.cc @@ -116,15 +116,15 @@ nvinfer1::DataType ConvertDataType(DataType type_id) { return data_type; } -cudaDataType_t ConvertDataType(nvinfer1::DataType type_id) { - std::map data_type_map = { +cudaDataType ConvertDataType(nvinfer1::DataType type_id) { + std::map data_type_map = { {nvinfer1::DataType::kINT8, CUDA_R_8I}, {nvinfer1::DataType::kINT32, CUDA_R_32I}, {nvinfer1::DataType::kFLOAT, CUDA_R_32F}, {nvinfer1::DataType::kHALF, CUDA_R_16F}, }; auto iter = data_type_map.find(type_id); - cudaDataType_t data_type; + cudaDataType data_type; if (iter != data_type_map.end()) { data_type = iter->second; } else { diff --git a/mindspore/lite/src/delegate/tensorrt/tensorrt_utils.h b/mindspore/lite/src/delegate/tensorrt/tensorrt_utils.h index d81b37f7335..aa09c518d6d 100644 --- a/mindspore/lite/src/delegate/tensorrt/tensorrt_utils.h +++ b/mindspore/lite/src/delegate/tensorrt/tensorrt_utils.h @@ -65,7 +65,7 @@ std::vector NHWC2NCHW(std::vector nhwc_shape); nvinfer1::DataType ConvertDataType(DataType type_id); -cudaDataType_t ConvertDataType(nvinfer1::DataType type_id); +cudaDataType ConvertDataType(nvinfer1::DataType type_id); nvinfer1::IShuffleLayer *NHWC2NCHW(nvinfer1::INetworkDefinition *network, const nvinfer1::ITensor &input);