1 Commits

Author SHA1 Message Date
Micaela Verucchi fe2b06d607 Fix patch
Signed-off-by: Micaela Verucchi <micaelaverucchi@gmail.com>
2020-07-16 18:37:37 +02:00
+7 -2
View File
@@ -62,9 +62,10 @@ void Conv2d::initCUDNN(bool back) {
// init workspace // init workspace
workSpace = NULL; workSpace = NULL;
ws_sizeInBytes = 0; ws_sizeInBytes = 0;
int algo_count = 0;
if(back) { if(back) {
checkCUDNN( cudnnGetConvolutionBackwardDataAlgorithm_v7(net->cudnnHandle, checkCUDNN( cudnnGetConvolutionBackwardDataAlgorithm_v7(net->cudnnHandle,
filterDesc, dstTensor, convDesc, srcTensor, 1, 0, &bwAlgo) ); filterDesc, dstTensor, convDesc, srcTensor, 1, &algo_count, &bwAlgo) );
checkCUDNN(cudnnGetConvolutionBackwardDataWorkspaceSize(net->cudnnHandle, checkCUDNN(cudnnGetConvolutionBackwardDataWorkspaceSize(net->cudnnHandle,
filterDesc, dstTensor, convDesc, srcTensor, filterDesc, dstTensor, convDesc, srcTensor,
bwAlgo.algo, &ws_sizeInBytes)); bwAlgo.algo, &ws_sizeInBytes));
@@ -74,13 +75,17 @@ void Conv2d::initCUDNN(bool back) {
srcTensorDesc = dstTensor; srcTensorDesc = dstTensor;
dstTensorDesc = srcTensor; dstTensorDesc = srcTensor;
} else { } else {
checkCUDNN( cudnnGetConvolutionForwardAlgorithm_v7(net->cudnnHandle, checkCUDNN( cudnnGetConvolutionForwardAlgorithm_v7(net->cudnnHandle,
srcTensor, filterDesc, convDesc, dstTensor, srcTensor, filterDesc, convDesc, dstTensor,
1, 0, &algo) ); 1, &algo_count, &algo) );
checkCUDNN(cudnnGetConvolutionForwardWorkspaceSize(net->cudnnHandle, checkCUDNN(cudnnGetConvolutionForwardWorkspaceSize(net->cudnnHandle,
srcTensor, filterDesc, convDesc, dstTensor, srcTensor, filterDesc, convDesc, dstTensor,
algo.algo, &ws_sizeInBytes)); algo.algo, &ws_sizeInBytes));
} }
if(algo_count < 1)
FatalError("Cannot retrieve convolutional algo");
} }
void Conv2d::inferCUDNN(dnnType* srcData, bool back) { void Conv2d::inferCUDNN(dnnType* srcData, bool back) {