| 103 | |
| 104 | #ifdef USE_CUDNN |
| 105 | class CudnnConvHandle : public ConvHandle { |
| 106 | public: |
| 107 | CudnnConvHandle(const Tensor &input, const std::vector<size_t> &kernel_size, |
| 108 | const std::vector<size_t> &stride, |
| 109 | const std::vector<size_t> &padding, const size_t in_channels, |
| 110 | const size_t out_channels, const bool bias, |
| 111 | const size_t groups = 1, |
| 112 | const size_t workspace_byte_limit = 1024 * 1024 * 1024, |
| 113 | const std::string &prefer = "fastest"); |
| 114 | ~CudnnConvHandle(); |
| 115 | // TODO(wangwei) add the destructor |
| 116 | |
| 117 | cudnnTensorDescriptor_t x_desc = nullptr; |
| 118 | cudnnTensorDescriptor_t y_desc = nullptr; |
| 119 | cudnnTensorDescriptor_t bias_desc = nullptr; |
| 120 | cudnnFilterDescriptor_t filter_desc = nullptr; |
| 121 | cudnnConvolutionDescriptor_t conv_desc = nullptr; |
| 122 | cudnnConvolutionFwdAlgo_t fp_alg; |
| 123 | cudnnConvolutionBwdFilterAlgo_t bp_filter_alg; |
| 124 | cudnnConvolutionBwdDataAlgo_t bp_data_alg; |
| 125 | |
| 126 | size_t workspace_count; |
| 127 | Tensor workspace; |
| 128 | size_t channels_per_filter; |
| 129 | }; |
| 130 | |
| 131 | Tensor GpuConvForward(const Tensor &x, const Tensor &W, const Tensor &b, |
| 132 | const CudnnConvHandle &cch); |