#include class FlattenConcatRT : public IPlugin { public: FlattenConcatRT() { stat = cublasCreate(&handle); if (stat != CUBLAS_STATUS_SUCCESS) { printf ("CUBLAS initialization failed\n"); return; } } ~FlattenConcatRT(){ } int getNbOutputs() const override { return 1; } Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override { return DimsCHW{ inputs[0].d[0] * inputs[0].d[1] * inputs[0].d[2], 1, 1}; } void configure(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, int maxBatchSize) override { assert(nbOutputs == 1 && nbInputs ==1); rows = inputDims[0].d[0]; cols = inputDims[0].d[1] * inputDims[0].d[2]; c = inputDims[0].d[0] * inputDims[0].d[1] * inputDims[0].d[2]; h = 1; w = 1; } int initialize() override { return 0; } virtual void terminate() override { checkERROR(cublasDestroy(handle)); } virtual size_t getWorkspaceSize(int maxBatchSize) const override { return 0; } 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( cudaMemcpyAsync(dstData, srcData, batchSize*rows*cols*sizeof(dnnType), cudaMemcpyDeviceToDevice, stream)); checkERROR( cublasSetStream(handle, stream) ); for(int i=0; i(buffer),*a = buf; tk::dnn::writeBUF(buf, c); tk::dnn::writeBUF(buf, h); tk::dnn::writeBUF(buf, w); tk::dnn::writeBUF(buf, rows); tk::dnn::writeBUF(buf, cols); assert(buf == a + getSerializationSize()); } int c, h, w; int rows, cols; cublasStatus_t stat; cublasHandle_t handle; };