From 8cff886ee535eb5ba94f8545bbd48f1b72928d4d Mon Sep 17 00:00:00 2001 From: Francesco Gatti Date: Thu, 23 Apr 2020 01:18:27 +0200 Subject: [PATCH] Flatten batch to be checked --- include/tkDNN/pluginsRT/FlattenConcatRT.h | 14 +++++++++----- tests/test_rtinference/rtinference.cpp | 3 ++- 2 files changed, 11 insertions(+), 6 deletions(-) diff --git a/include/tkDNN/pluginsRT/FlattenConcatRT.h b/include/tkDNN/pluginsRT/FlattenConcatRT.h index e6e4eb2..4ebce69 100644 --- a/include/tkDNN/pluginsRT/FlattenConcatRT.h +++ b/include/tkDNN/pluginsRT/FlattenConcatRT.h @@ -47,11 +47,15 @@ public: virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override { dnnType *srcData = (dnnType*)reinterpret_cast(inputs[0]); dnnType *dstData = reinterpret_cast(outputs[0]); - checkCuda( cudaMemcpy(dstData, srcData, rows*cols*sizeof(dnnType), cudaMemcpyDeviceToDevice)); - - float const alpha(1.0); - float const beta(0.0); - checkERROR( cublasSgeam( handle, CUBLAS_OP_T, CUBLAS_OP_N, rows, cols, &alpha, srcData, cols, &beta, srcData, rows, dstData, rows )); + checkCuda( cudaMemcpyAsync(dstData, srcData, batchSize*rows*cols*sizeof(dnnType), cudaMemcpyDeviceToDevice, stream)); + + checkERROR( cublasSetStream(handle, stream) ); + for(int i=0; i