!31722 [MSLITE] fix compile support for cuda 10.1

Merge pull request !31722 from Liu_Xuu/trt_0322_101
This commit is contained in:
i-robot 2022-03-23 03:13:02 +00:00 committed by Gitee
commit 556c882ad9
No known key found for this signature in database
GPG Key ID: 173E9B9CA92EEF8F
6 changed files with 20 additions and 23 deletions

View File

@ -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

View File

@ -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_

View File

@ -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<int>(desc_a.type)
<< ", b: " << static_cast<int>(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 {

View File

@ -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 {

View File

@ -116,15 +116,15 @@ nvinfer1::DataType ConvertDataType(DataType type_id) {
return data_type;
}
cudaDataType_t ConvertDataType(nvinfer1::DataType type_id) {
std::map<nvinfer1::DataType, cudaDataType_t> data_type_map = {
cudaDataType ConvertDataType(nvinfer1::DataType type_id) {
std::map<nvinfer1::DataType, cudaDataType> 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 {

View File

@ -65,7 +65,7 @@ std::vector<int64_t> NHWC2NCHW(std::vector<int64_t> 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);