MCPcopy Create free account
hub / github.com/apache/singa / CudnnConvHandle

Method CudnnConvHandle

src/model/operation/convolution.cc:444–572  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

442
443#ifdef USE_CUDNN
444CudnnConvHandle::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;

Callers 2

re_new_handleFunction · 0.80
initializeMethod · 0.80

Calls 9

GetCudnnDataTypeFunction · 0.85
maxFunction · 0.85
SizeOfFunction · 0.85
data_typeMethod · 0.80
deviceMethod · 0.80
contextMethod · 0.80
TensorClass · 0.50
beginMethod · 0.45
endMethod · 0.45

Tested by

no test coverage detected