| 53 | } |
| 54 | |
| 55 | void CudnnConvolution::InitCudnn(const Tensor &input) { |
| 56 | DataType dtype = input.data_type(); |
| 57 | auto dev = input.device(); |
| 58 | Context *ctx = dev->context(0); |
| 59 | size_t batchsize = input.shape(0); |
| 60 | if (!has_init_cudnn_) { |
| 61 | CUDNN_CHECK(cudnnCreateTensorDescriptor(&x_desc_)); |
| 62 | CUDNN_CHECK(cudnnCreateTensorDescriptor(&y_desc_)); |
| 63 | if (bias_term_) |
| 64 | CUDNN_CHECK(cudnnCreateTensorDescriptor(&bias_desc_)); |
| 65 | CUDNN_CHECK(cudnnCreateFilterDescriptor(&filter_desc_)); |
| 66 | CUDNN_CHECK(cudnnCreateConvolutionDescriptor(&conv_desc_)); |
| 67 | } |
| 68 | |
| 69 | CUDNN_CHECK(cudnnSetTensor4dDescriptor(x_desc_, CUDNN_TENSOR_NCHW, |
| 70 | GetCudnnDataType(dtype), batchsize, |
| 71 | channels_, height_, width_)); |
| 72 | CUDNN_CHECK(cudnnSetTensor4dDescriptor( |
| 73 | y_desc_, CUDNN_TENSOR_NCHW, GetCudnnDataType(dtype), batchsize, |
| 74 | num_filters_, conv_height_, conv_width_)); |
| 75 | if (bias_term_) |
| 76 | CUDNN_CHECK(cudnnSetTensor4dDescriptor(bias_desc_, CUDNN_TENSOR_NCHW, |
| 77 | GetCudnnDataType(dtype), 1, |
| 78 | num_filters_, 1, 1)); |
| 79 | CUDNN_CHECK(cudnnSetConvolution2dDescriptor(conv_desc_, pad_h_, pad_w_, |
| 80 | stride_h_, stride_w_, 1, 1, // dilation x and y |
| 81 | CUDNN_CROSS_CORRELATION |
| 82 | #if CUDNN_MAJOR >= 7 |
| 83 | , GetCudnnDataType(dtype) |
| 84 | #endif // CUDNN_MAJOR |
| 85 | )); |
| 86 | CUDNN_CHECK(cudnnSetFilter4dDescriptor(filter_desc_, GetCudnnDataType(dtype), |
| 87 | CUDNN_TENSOR_NCHW, num_filters_, |
| 88 | channels_, kernel_h_, kernel_w_)); |
| 89 | if (prefer_ == "fastest" || prefer_ == "limited_workspace" || |
| 90 | prefer_ == "no_workspace") { |
| 91 | cudnnConvolutionFwdPreference_t fwd_pref; |
| 92 | cudnnConvolutionBwdFilterPreference_t bwd_filt_pref; |
| 93 | cudnnConvolutionBwdDataPreference_t bwd_data_pref; |
| 94 | if (prefer_ == "fastest") { |
| 95 | fwd_pref = CUDNN_CONVOLUTION_FWD_PREFER_FASTEST; |
| 96 | bwd_filt_pref = CUDNN_CONVOLUTION_BWD_FILTER_PREFER_FASTEST; |
| 97 | bwd_data_pref = CUDNN_CONVOLUTION_BWD_DATA_PREFER_FASTEST; |
| 98 | } else if (prefer_ == "limited_workspace") { |
| 99 | fwd_pref = CUDNN_CONVOLUTION_FWD_SPECIFY_WORKSPACE_LIMIT; |
| 100 | bwd_filt_pref = CUDNN_CONVOLUTION_BWD_FILTER_SPECIFY_WORKSPACE_LIMIT; |
| 101 | bwd_data_pref = CUDNN_CONVOLUTION_BWD_DATA_SPECIFY_WORKSPACE_LIMIT; |
| 102 | } else { |
| 103 | fwd_pref = CUDNN_CONVOLUTION_FWD_NO_WORKSPACE; |
| 104 | bwd_filt_pref = CUDNN_CONVOLUTION_BWD_FILTER_NO_WORKSPACE; |
| 105 | bwd_data_pref = CUDNN_CONVOLUTION_BWD_DATA_SPECIFY_WORKSPACE_LIMIT; |
| 106 | } |
| 107 | CUDNN_CHECK(cudnnGetConvolutionForwardAlgorithm( |
| 108 | ctx->cudnn_handle, x_desc_, filter_desc_, conv_desc_, y_desc_, fwd_pref, |
| 109 | workspace_byte_limit_, &fp_alg_)); |
| 110 | CUDNN_CHECK(cudnnGetConvolutionBackwardFilterAlgorithm( |
| 111 | ctx->cudnn_handle, x_desc_, y_desc_, conv_desc_, filter_desc_, |
| 112 | bwd_filt_pref, workspace_byte_limit_, &bp_filter_alg_)); |