From fe2b06d60797c2f7deff33b3e086ee5c262ed666 Mon Sep 17 00:00:00 2001 From: Micaela Verucchi Date: Thu, 16 Jul 2020 18:37:37 +0200 Subject: [PATCH] Fix patch Signed-off-by: Micaela Verucchi --- src/Conv2d.cpp | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/src/Conv2d.cpp b/src/Conv2d.cpp index 595fec7..b57cf58 100644 --- a/src/Conv2d.cpp +++ b/src/Conv2d.cpp @@ -62,9 +62,10 @@ void Conv2d::initCUDNN(bool back) { // init workspace workSpace = NULL; ws_sizeInBytes = 0; + int algo_count = 0; if(back) { checkCUDNN( cudnnGetConvolutionBackwardDataAlgorithm_v7(net->cudnnHandle, - filterDesc, dstTensor, convDesc, srcTensor, 1, 0, &bwAlgo) ); + filterDesc, dstTensor, convDesc, srcTensor, 1, &algo_count, &bwAlgo) ); checkCUDNN(cudnnGetConvolutionBackwardDataWorkspaceSize(net->cudnnHandle, filterDesc, dstTensor, convDesc, srcTensor, bwAlgo.algo, &ws_sizeInBytes)); @@ -74,13 +75,17 @@ void Conv2d::initCUDNN(bool back) { srcTensorDesc = dstTensor; dstTensorDesc = srcTensor; } else { + checkCUDNN( cudnnGetConvolutionForwardAlgorithm_v7(net->cudnnHandle, srcTensor, filterDesc, convDesc, dstTensor, - 1, 0, &algo) ); + 1, &algo_count, &algo) ); checkCUDNN(cudnnGetConvolutionForwardWorkspaceSize(net->cudnnHandle, srcTensor, filterDesc, convDesc, dstTensor, algo.algo, &ws_sizeInBytes)); } + + if(algo_count < 1) + FatalError("Cannot retrieve convolutional algo"); } void Conv2d::inferCUDNN(dnnType* srcData, bool back) {