diff --git a/include/tkDNN/Layer.h b/include/tkDNN/Layer.h index bd544b2..f2ec56d 100644 --- a/include/tkDNN/Layer.h +++ b/include/tkDNN/Layer.h @@ -273,8 +273,8 @@ public: protected: cudnnFilterDescriptor_t filterDesc; cudnnConvolutionDescriptor_t convDesc; - cudnnConvolutionFwdAlgo_t algo; - cudnnConvolutionBwdDataAlgo_t bwAlgo; + cudnnConvolutionFwdAlgoPerf_t algo; + cudnnConvolutionBwdDataAlgoPerf_t bwAlgo; cudnnTensorDescriptor_t biasTensorDesc; void initCUDNN(bool back = false); diff --git a/src/Conv2d.cpp b/src/Conv2d.cpp index 4704c66..595fec7 100644 --- a/src/Conv2d.cpp +++ b/src/Conv2d.cpp @@ -63,23 +63,23 @@ void Conv2d::initCUDNN(bool back) { workSpace = NULL; ws_sizeInBytes = 0; if(back) { - checkCUDNN( cudnnGetConvolutionBackwardDataAlgorithm(net->cudnnHandle, - filterDesc, dstTensor, convDesc, srcTensor, - CUDNN_CONVOLUTION_BWD_DATA_PREFER_FASTEST, 0, &bwAlgo) ); + checkCUDNN( cudnnGetConvolutionBackwardDataAlgorithm_v7(net->cudnnHandle, + filterDesc, dstTensor, convDesc, srcTensor, 1, 0, &bwAlgo) ); checkCUDNN(cudnnGetConvolutionBackwardDataWorkspaceSize(net->cudnnHandle, - filterDesc, dstTensor, convDesc, srcTensor, - bwAlgo, &ws_sizeInBytes)); + filterDesc, dstTensor, convDesc, srcTensor, + bwAlgo.algo, &ws_sizeInBytes)); + // invert tensors srcTensorDesc = dstTensor; dstTensorDesc = srcTensor; } else { - checkCUDNN( cudnnGetConvolutionForwardAlgorithm(net->cudnnHandle, - srcTensor, filterDesc, convDesc, dstTensor, - CUDNN_CONVOLUTION_FWD_PREFER_FASTEST, 0, &algo) ); - checkCUDNN(cudnnGetConvolutionForwardWorkspaceSize(net->cudnnHandle, - srcTensor, filterDesc, convDesc, dstTensor, - algo, &ws_sizeInBytes)); + checkCUDNN( cudnnGetConvolutionForwardAlgorithm_v7(net->cudnnHandle, + srcTensor, filterDesc, convDesc, dstTensor, + 1, 0, &algo) ); + checkCUDNN(cudnnGetConvolutionForwardWorkspaceSize(net->cudnnHandle, + srcTensor, filterDesc, convDesc, dstTensor, + algo.algo, &ws_sizeInBytes)); } } @@ -91,12 +91,12 @@ void Conv2d::inferCUDNN(dnnType* srcData, bool back) { checkCUDNN(cudnnConvolutionBackwardData(net->cudnnHandle, &alpha, filterDesc, data_d, srcTensorDesc, srcData, - convDesc, bwAlgo, workSpace, ws_sizeInBytes, + convDesc, bwAlgo.algo, workSpace, ws_sizeInBytes, &beta, dstTensorDesc, dstData)); } else { checkCUDNN(cudnnConvolutionForward(net->cudnnHandle, &alpha, srcTensorDesc, srcData, filterDesc, - data_d, convDesc, algo, workSpace, ws_sizeInBytes, + data_d, convDesc, algo.algo, workSpace, ws_sizeInBytes, &beta, dstTensorDesc, dstData)); }