| 442 | |
| 443 | #ifdef USE_CUDNN |
| 444 | CudnnConvHandle::CudnnConvHandle( |
| 445 | const Tensor &input, const std::vector<size_t> &kernel_size, |
| 446 | const std::vector<size_t> &stride, const std::vector<size_t> &padding, |
| 447 | const size_t in_channels, const size_t out_channels, const bool bias, |
| 448 | const size_t groups, const size_t workspace_byte_limit, |
| 449 | const std::string &prefer_) |
| 450 | : ConvHandle(input, kernel_size, stride, padding, in_channels, out_channels, |
| 451 | bias, groups) { |
| 452 | std::string prefer = prefer_; |
| 453 | if (const char *env_p = std::getenv("CUDNN_CONV_ALG")) { |
| 454 | prefer = std::string(env_p); |
| 455 | std::transform(prefer.begin(), prefer.end(), prefer.begin(), tolower); |
| 456 | LOG(INFO) << "CUDNN_CONV_ALG: " << prefer; |
| 457 | } |
| 458 | DataType dtype = input.data_type(); |
| 459 | auto dev = input.device(); |
| 460 | Context *ctx = dev->context(0); |
| 461 | channels_per_filter = channels / groups; |
| 462 | |
| 463 | CUDNN_CHECK(cudnnCreateTensorDescriptor(&x_desc)); |
| 464 | CUDNN_CHECK(cudnnCreateTensorDescriptor(&y_desc)); |
| 465 | if (bias_term) CUDNN_CHECK(cudnnCreateTensorDescriptor(&bias_desc)); |
| 466 | CUDNN_CHECK(cudnnCreateFilterDescriptor(&filter_desc)); |
| 467 | CUDNN_CHECK(cudnnCreateConvolutionDescriptor(&conv_desc)); |
| 468 | |
| 469 | CUDNN_CHECK(cudnnSetTensor4dDescriptor(x_desc, CUDNN_TENSOR_NCHW, |
| 470 | GetCudnnDataType(dtype), batchsize, |
| 471 | channels, height, width)); |
| 472 | CUDNN_CHECK(cudnnSetTensor4dDescriptor(y_desc, CUDNN_TENSOR_NCHW, |
| 473 | GetCudnnDataType(dtype), batchsize, |
| 474 | num_filters, conv_height, conv_width)); |
| 475 | if (bias_term) |
| 476 | CUDNN_CHECK(cudnnSetTensor4dDescriptor(bias_desc, CUDNN_TENSOR_NCHW, |
| 477 | GetCudnnDataType(dtype), 1, |
| 478 | num_filters, 1, 1)); |
| 479 | CUDNN_CHECK(cudnnSetConvolution2dDescriptor( |
| 480 | conv_desc, pad_h, pad_w, stride_h, stride_w, 1, 1, CUDNN_CROSS_CORRELATION |
| 481 | #if CUDNN_MAJOR >= 7 |
| 482 | , |
| 483 | GetCudnnDataType(dtype) |
| 484 | #endif |
| 485 | )); |
| 486 | if (CUDNN_MAJOR >= 7 && groups > 1) { |
| 487 | CUDNN_CHECK(cudnnSetConvolutionGroupCount(conv_desc, groups)); |
| 488 | } else if (groups > 1) { |
| 489 | LOG(FATAL) |
| 490 | << "The current version of cuDNN not support grouped convolution."; |
| 491 | }; |
| 492 | |
| 493 | CUDNN_CHECK(cudnnSetFilter4dDescriptor( |
| 494 | filter_desc, GetCudnnDataType(dtype), CUDNN_TENSOR_NCHW, num_filters, |
| 495 | channels / groups, kernel_h, kernel_w)); |
| 496 | |
| 497 | if (prefer == "tensor_ops") { |
| 498 | // std::cout<<"using tensor op\n"; |
| 499 | CUDNN_CHECK(cudnnSetConvolutionMathType(conv_desc, CUDNN_TENSOR_OP_MATH)); |
| 500 | fp_alg = CUDNN_CONVOLUTION_FWD_ALGO_IMPLICIT_PRECOMP_GEMM; |
| 501 | bp_filter_alg = CUDNN_CONVOLUTION_BWD_FILTER_ALGO_1; |