tensorrt7 support for ipluginv2

This commit is contained in:
perseusdg
2021-10-18 18:07:57 +05:30
parent 18cc6abbb6
commit 54e7af11ed
30 changed files with 313 additions and 64 deletions
+19 -1
View File
@@ -54,6 +54,7 @@ size_t FlattenConcatRT::getWorkspaceSize(int maxBatchSize) const NOEXCEPT {
return 0;
}
#if NV_TENSORRT_MAJOR > 7
int FlattenConcatRT::enqueue(int batchSize, const void *const *inputs, void *const *outputs, void *workspace,
cudaStream_t stream) NOEXCEPT {
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
@@ -69,6 +70,24 @@ int FlattenConcatRT::enqueue(int batchSize, const void *const *inputs, void *con
}
return 0;
}
#elif NV_TENSORRT_MAJOR == 7
int32_t FlattenConcatRT::enqueue(int32_t batchSize, const void *const *inputs, void **outputs, void *workspace,
cudaStream_t stream) {
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
dnnType *dstData = reinterpret_cast<dnnType*>(outputs[0]);
checkCuda( cudaMemcpyAsync(dstData, srcData, batchSize*rows*cols*sizeof(dnnType), cudaMemcpyDeviceToDevice, stream));
checkERROR( cublasSetStream(handle, stream) );
for(int i=0; i<batchSize; i++) {
float const alpha(1.0);
float const beta(0.0);
int offset = i*rows*cols;
checkERROR( cublasSgeam( handle, CUBLAS_OP_T, CUBLAS_OP_N, rows, cols, &alpha, srcData + offset, cols, &beta, srcData + offset, rows, dstData + offset, rows ));
}
return 0;
}
#endif
size_t FlattenConcatRT::getSerializationSize() const NOEXCEPT {
return 5*sizeof(int);
@@ -114,7 +133,6 @@ IPluginV2 *FlattenConcatRT::clone() const NOEXCEPT {
return p;
}
FlattenConcatRTPluginCreator::FlattenConcatRTPluginCreator() {
mPluginAttributes.clear();
mFC.nbFields = mPluginAttributes.size();