TRT8 works with almost every nerual network now!!!!(including demo3d)

This commit is contained in:
perseusdg
2021-10-28 23:35:37 +05:30
parent 8c36dd0431
commit c5e66c6bf6
37 changed files with 1042 additions and 700 deletions
+137 -97
View File
@@ -1,14 +1,19 @@
#include <tkDNN/pluginsRT/DeformableConvRT.h>
#include <utility>
using namespace nvinfer1;
using namespace tk::dnn;
std::vector<PluginField> DeformableConvRTPluginCreator::mPluginAttributes;
PluginFieldCollection DeformableConvRTPluginCreator::mFC{};
static const char* DEFORMABLECONVRT_PLUGIN_VERSION{"1"};
static const char* DEFORMABLECONVRT_PLUGIN_NAME{"DeformableConvRT_tkDNN"};
DeformableConvRT::DeformableConvRT(int chunk_dim, int kh, int kw, int sh, int sw, int ph, int pw, int deformableGroup,
int i_n, int i_c, int i_h, int i_w, int o_n, int o_c, int o_h, int o_w,
tk::dnn::DeformConv2d *deformable) {
int i_n, int i_c, int i_h, int i_w, int o_n, int o_c, int o_h, int o_w,std::vector<dnnType> data_H,std::vector<dnnType> bias2_H,
std::vector<dnnType> ones_d1_h,std::vector<dnnType> ones_d2_h,std::vector<dnnType> offsetH,std::vector<dnnType> maskH,int height_ones,int width_ones,int dim_ones) {
this->chunk_dim = chunk_dim;
this->kh = kh;
this->kw = kw;
@@ -25,11 +30,15 @@ DeformableConvRT::DeformableConvRT(int chunk_dim, int kh, int kw, int sh, int sw
this->o_c = o_c;
this->o_h = o_h;
this->o_w = o_w;
this->defRT = deformable;
height_ones = (i_h + 2 * ph - (1 * (kh - 1) + 1)) / sh + 1;
width_ones = (i_w + 2 * pw - (1 * (kw - 1) + 1)) / sw + 1;
dim_ones = i_c * kh * kw * 1 * height_ones * width_ones;
this->mask_v = std::move(maskH);
this->offset_v = std::move(offsetH);
this->ones_d2_v = std::move(ones_d2_h);
this->ones_d1_v = std::move(ones_d1_h);
this->data_d_v = std::move(data_H);
this->bias2_d_v = std::move(bias2_H);
this->height_ones = height_ones;
this->width_ones = width_ones;
this->dim_ones = dim_ones;
checkCuda( cudaMalloc(&data_d, i_c * o_c * kh * kw * 1 * sizeof(dnnType)));
checkCuda( cudaMalloc(&bias2_d, o_c*sizeof(dnnType)));
@@ -37,17 +46,15 @@ DeformableConvRT::DeformableConvRT(int chunk_dim, int kh, int kw, int sh, int sw
checkCuda( cudaMalloc(&offset, 2*chunk_dim*sizeof(dnnType)));
checkCuda( cudaMalloc(&mask, chunk_dim*sizeof(dnnType)));
checkCuda( cudaMalloc(&ones_d2, dim_ones*sizeof(dnnType)));
if(deformable != nullptr) {
checkCuda( cudaMemcpy(data_d, deformable->data_d, sizeof(dnnType)*i_c * o_c * kh * kw * 1, cudaMemcpyDeviceToDevice) );
checkCuda( cudaMemcpy(bias2_d, deformable->bias2_d, sizeof(dnnType)*o_c, cudaMemcpyDeviceToDevice) );
checkCuda( cudaMemcpy(ones_d1, deformable->ones_d1, sizeof(dnnType)*height_ones*width_ones, cudaMemcpyDeviceToDevice) );
checkCuda( cudaMemcpy(offset, deformable->offset, sizeof(dnnType)*2*chunk_dim, cudaMemcpyDeviceToDevice) );
checkCuda( cudaMemcpy(mask, deformable->mask, sizeof(dnnType)*chunk_dim, cudaMemcpyDeviceToDevice) );
checkCuda( cudaMemcpy(ones_d2, deformable->ones_d2, sizeof(dnnType)*dim_ones, cudaMemcpyDeviceToDevice) );
if(!data_d_v.empty() && !bias2_d_v.empty() && !ones_d1_v.empty() && !ones_d2_v.empty() && !mask_v.empty() && !offset_v.empty()) {
checkCuda(cudaMemcpy(data_d, data_d_v.data(), sizeof(dnnType) * data_d_v.size(), cudaMemcpyHostToDevice));
checkCuda(cudaMemcpy(bias2_d, bias2_d_v.data(), sizeof(dnnType) * bias2_d_v.size(), cudaMemcpyHostToDevice));
checkCuda(cudaMemcpy(ones_d1, ones_d1_v.data(), sizeof(dnnType) * ones_d1_v.size(), cudaMemcpyHostToDevice));
checkCuda(cudaMemcpy(offset, offset_v.data(), sizeof(dnnType) * offset_v.size(), cudaMemcpyHostToDevice));
checkCuda(cudaMemcpy(mask, mask_v.data(), sizeof(dnnType) * mask_v.size(), cudaMemcpyHostToDevice));
checkCuda(cudaMemcpy(ones_d2, ones_d2_v.data(), sizeof(dnnType) * ones_d2_v.size(), cudaMemcpyHostToDevice));
}
stat = cublasCreate(&handle);
if (stat != CUBLAS_STATUS_SUCCESS)
FatalError("CUBLAS initialization failed\n");
}
@@ -79,42 +86,27 @@ DeformableConvRT::DeformableConvRT(const void *data, size_t length) {
o_c = readBUF<int>(buf);
o_h = readBUF<int>(buf);
o_w = readBUF<int>(buf);
dnnType *aus = new dnnType[chunk_dim*2];
height_ones = readBUF<int>(buf);
width_ones = readBUF<int>(buf);
dim_ones = readBUF<int>(buf);
offset_v.resize(chunk_dim*2);
for(int i=0;i<chunk_dim*2;i++)
aus[i] = readBUF<dnnType>(buf);
checkCuda(cudaMemcpy(offset,aus,sizeof(dnnType)*2*chunk_dim,cudaMemcpyHostToDevice));
free(aus);
aus = new dnnType[chunk_dim];
offset_v[i] = readBUF<dnnType>(buf);
mask_v.resize(chunk_dim);
for(int i=0;i<chunk_dim;i++)
aus[i] = readBUF<dnnType>(buf);
checkCuda(cudaMemcpy(mask,aus,sizeof(dnnType)*chunk_dim,cudaMemcpyHostToDevice));
free(aus);
aus = new dnnType[i_c*o_c*kh*kw*1];
mask_v[i] = readBUF<dnnType>(buf);
data_d_v.resize(i_c*o_c*kh*kw*1);
for(int i=0;i<(i_c*o_c*kh*kw*1);i++)
aus[i] = readBUF<dnnType>(buf);
checkCuda(cudaMemcpy(data_d,aus,sizeof(dnnType)*(i_c*o_c*kh*kw*1),cudaMemcpyHostToDevice));
free(aus);
aus = new dnnType[o_c];
data_d_v[i] = readBUF<dnnType>(buf);
bias2_d_v.resize(o_c);
for(int i=0; i < o_c; i++)
aus[i] = readBUF<dnnType>(buf);
checkCuda( cudaMemcpy(bias2_d, aus, sizeof(dnnType)*o_c, cudaMemcpyHostToDevice) );
free(aus);
aus = new dnnType[height_ones * width_ones];
bias2_d_v[i] = readBUF<dnnType>(buf);
ones_d1_v.resize(height_ones*width_ones);
for(int i=0; i<height_ones * width_ones; i++)
aus[i] = readBUF<dnnType>(buf);
checkCuda( cudaMemcpy(ones_d1, aus, sizeof(dnnType)*height_ones * width_ones, cudaMemcpyHostToDevice) );
free(aus);
aus = new dnnType[dim_ones];
ones_d1_v[i] = readBUF<dnnType>(buf);
ones_d2_v.resize(dim_ones);
for(int i=0; i<dim_ones; i++)
aus[i] = readBUF<dnnType>(buf);
checkCuda( cudaMemcpy(ones_d2, aus, sizeof(dnnType)*dim_ones, cudaMemcpyHostToDevice) );
free(aus);
ones_d2_v[i] = readBUF<dnnType>(buf);
assert(buf == bufCheck + length);
}
@@ -124,11 +116,9 @@ int DeformableConvRT::getNbOutputs() const NOEXCEPT {
}
Dims DeformableConvRT::getOutputDimensions(int index, const Dims *inputs, int nbInputDims) NOEXCEPT {
return Dims3{defRT->output_dim.c, defRT->output_dim.h, defRT->output_dim.w};
return Dims3{o_c, o_h, o_w};
}
void DeformableConvRT::configureWithFormat(const Dims *inputDims, int nbInputs, const Dims *outputDims, int nbOutputs,DataType type, PluginFormat format, int maxBatchSize) NOEXCEPT {}
int DeformableConvRT::initialize() NOEXCEPT {
return 0;
}
@@ -166,7 +156,7 @@ int DeformableConvRT::enqueue(int batchSize, const void *const *inputs, void *co
}
return 0;
}
#elif NV_TENSORRT_MAJOR == 7
#elif NV_TENSORRT_MAJOR <= 7
int32_t DeformableConvRT::enqueue(int32_t batchSize, const void *const *inputs, void **outputs, void *workspace,
cudaStream_t stream) {
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
@@ -198,7 +188,7 @@ int32_t DeformableConvRT::enqueue(int32_t batchSize, const void *const *inputs,
#endif
size_t DeformableConvRT::getSerializationSize() const NOEXCEPT {
return 16 * sizeof(int) + chunk_dim * 3 * sizeof(dnnType) + (i_c * o_c * kh * kw * 1 ) * sizeof(dnnType) +
return 19 * sizeof(int) + chunk_dim * 3 * sizeof(dnnType) + (i_c * o_c * kh * kw * 1 ) * sizeof(dnnType) +
o_c * sizeof(dnnType) + height_ones * width_ones * sizeof(dnnType) + dim_ones * sizeof(dnnType);
}
@@ -220,45 +210,27 @@ void DeformableConvRT::serialize(void *buffer) const NOEXCEPT {
writeBUF(buf, o_c);
writeBUF(buf, o_h);
writeBUF(buf, o_w);
dnnType *aus = new dnnType[chunk_dim*2];
checkCuda( cudaMemcpy(aus, offset, sizeof(dnnType)*2*chunk_dim, cudaMemcpyDeviceToHost) );
for(int i=0; i<chunk_dim*2; i++)
writeBUF(buf, aus[i]);
free(aus);
aus = new dnnType[chunk_dim];
checkCuda( cudaMemcpy(aus, mask, sizeof(dnnType)*chunk_dim, cudaMemcpyDeviceToHost) );
for(int i=0; i<chunk_dim; i++)
writeBUF(buf, aus[i]);
free(aus);
aus = new dnnType[(i_c * o_c * kh * kw * 1 )];
checkCuda( cudaMemcpy(aus, data_d, sizeof(dnnType)*(i_c * o_c * kh * kw * 1 ), cudaMemcpyDeviceToHost) );
for(int i=0; i<(i_c * o_c * kh * kw * 1 ); i++)
writeBUF(buf, aus[i]);
free(aus);
aus = new dnnType[o_c];
checkCuda( cudaMemcpy(aus, bias2_d, sizeof(dnnType)*o_c, cudaMemcpyDeviceToHost) );
for(int i=0; i < o_c; i++)
writeBUF(buf, aus[i]);
free(aus);
aus = new dnnType[height_ones * width_ones];
checkCuda( cudaMemcpy(aus, ones_d1, sizeof(dnnType)*height_ones * width_ones, cudaMemcpyDeviceToHost) );
for(int i=0; i<height_ones * width_ones; i++)
writeBUF(buf, aus[i]);
free(aus);
aus = new dnnType[dim_ones];
checkCuda( cudaMemcpy(aus, ones_d2, sizeof(dnnType)*dim_ones, cudaMemcpyDeviceToHost) );
for(int i=0; i<dim_ones; i++)
writeBUF(buf, aus[i]);
free(aus);
writeBUF(buf,height_ones);
writeBUF(buf,width_ones);
writeBUF(buf,dim_ones);
for(int i=0; i<offset_v.size(); i++)
writeBUF(buf, offset_v[i]);
for(int i=0; i<mask_v.size(); i++)
writeBUF(buf, mask_v[i]);
for(int i=0; i<data_d_v.size(); i++)
writeBUF(buf, data_d_v[i]);
for(int i=0; i < bias2_d_v.size(); i++)
writeBUF(buf, bias2_d_v[i]);
for(int i=0; i<ones_d1_v.size(); i++)
writeBUF(buf, ones_d1_v[i]);
for(int i=0; i<ones_d2_v.size(); i++)
writeBUF(buf, ones_d2_v[i]);
assert(buf == a + getSerializationSize());
}
void DeformableConvRT::destroy() NOEXCEPT { delete this; }
bool DeformableConvRT::supportsFormat(DataType type, PluginFormat format) const NOEXCEPT {
return true;
//todo assert
}
const char *DeformableConvRT::getPluginNamespace() const NOEXCEPT {
return mPluginNamespace.c_str();
@@ -269,21 +241,81 @@ void DeformableConvRT::setPluginNamespace(const char *pluginNamespace) NOEXCEPT
}
const char *DeformableConvRT::getPluginType() const NOEXCEPT {
return "DeformableConvRT_tkDNN";
return DEFORMABLECONVRT_PLUGIN_NAME;
}
const char *DeformableConvRT::getPluginVersion() const NOEXCEPT {
return "1";
return DEFORMABLECONVRT_PLUGIN_VERSION;
}
IPluginV2 *DeformableConvRT::clone() const NOEXCEPT {
auto *p = new DeformableConvRT(chunk_dim,kh,kw,sh,sw,ph,pw,deformableGroup,i_n,i_c,i_h,i_w,o_n,o_c,o_h,o_w,defRT);
IPluginV2Ext *DeformableConvRT::clone() const NOEXCEPT {
auto *p = new DeformableConvRT(chunk_dim,kh,kw,sh,sw,ph,pw,deformableGroup,i_n,i_c,i_h,i_w,o_n,o_c,o_h,o_w,data_d_v,bias2_d_v,ones_d1_v,ones_d2_v,offset_v,mask_v,height_ones,width_ones,dim_ones);
p->setPluginNamespace(mPluginNamespace.c_str());
return p;
}
DataType
DeformableConvRT::getOutputDataType(int index, const nvinfer1::DataType *inputTypes, int nbInputs) const NOEXCEPT {
return DataType::kFLOAT;
}
void DeformableConvRT::attachToContext(cudnnContext *cudnnContext, cublasContext *cublasContext,
IGpuAllocator *gpuAllocator) NOEXCEPT {
handle = cublasContext;
}
bool DeformableConvRT::isOutputBroadcastAcrossBatch(int outputIndex, const bool *inputIsBroadcasted,
int nbInputs) const NOEXCEPT {
return false;
}
bool DeformableConvRT::canBroadcastInputAcrossBatch(int inputIndex) const NOEXCEPT {
return false;
}
void DeformableConvRT::configurePlugin(const Dims *inputDims, int32_t nbInputs, const Dims *outputDims, int32_t nbOutputs,
const DataType *inputTypes, const DataType *outputTypes, const bool *inputIsBroadcast,
const bool *outputIsBroadcast, PluginFormat floatFormat,
int32_t maxBatchSize) NOEXCEPT {
}
void DeformableConvRT::detachFromContext() NOEXCEPT {
}
bool DeformableConvRT::supportsFormat(DataType type, PluginFormat format) const NOEXCEPT {
return true;
}
DeformableConvRTPluginCreator::DeformableConvRTPluginCreator() {
mPluginAttributes.clear();
mPluginAttributes.emplace_back(PluginField("chunk_dim", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("kh", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("kw", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("sh", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("sw", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("ph", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("pw", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("deformable_group", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("i_n", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("i_c", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("i_h", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("i_w", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("o_n", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("o_c", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("o_h", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("o_w", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("mask_v", nullptr,PluginFieldType::kFLOAT32,1));
mPluginAttributes.emplace_back(PluginField("offset_v", nullptr,PluginFieldType::kFLOAT32,1));
mPluginAttributes.emplace_back(PluginField("ones_d2_v", nullptr,PluginFieldType::kFLOAT32,1));
mPluginAttributes.emplace_back(PluginField("ones_d1_v", nullptr,PluginFieldType::kFLOAT32,1));
mPluginAttributes.emplace_back(PluginField("data_d_v", nullptr,PluginFieldType::kFLOAT32,1));
mPluginAttributes.emplace_back(PluginField("bias2_d_v", nullptr,PluginFieldType::kFLOAT32,1));
mPluginAttributes.emplace_back(PluginField("height_ones", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("width_ones", nullptr,PluginFieldType::kINT32,1));
mPluginAttributes.emplace_back(PluginField("dim_ones", nullptr,PluginFieldType::kINT32,1));
mFC.nbFields = mPluginAttributes.size();
mFC.fields = mPluginAttributes.data();
}
@@ -296,14 +328,14 @@ const char *DeformableConvRTPluginCreator::getPluginNamespace() const NOEXCEPT {
return mPluginNamespace.c_str();
}
IPluginV2 *DeformableConvRTPluginCreator::deserializePlugin(const char *name, const void *serialData,
IPluginV2Ext *DeformableConvRTPluginCreator::deserializePlugin(const char *name, const void *serialData,
size_t serialLength) NOEXCEPT {
auto *pluginObj = new DeformableConvRT(serialData,serialLength);
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
return pluginObj;
}
IPluginV2 *DeformableConvRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
IPluginV2Ext *DeformableConvRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
const PluginField *fields = fc->fields;
int chunk_dim = *(static_cast<const int *>(fields[0].data));
int kh = *(static_cast<const int *>(fields[1].data));
@@ -320,19 +352,27 @@ IPluginV2 *DeformableConvRTPluginCreator::createPlugin(const char *name, const P
int o_n = *(static_cast<const int *>(fields[12].data));
int o_c = *(static_cast<const int *>(fields[13].data));
int o_h = *(static_cast<const int *>(fields[14].data));
int o_w = *(static_cast<const int *>(fields[14].data));
auto *defRT = const_cast<DeformConv2d *>(static_cast<const DeformConv2d *>(fields[15].data));
auto *pluginObj = new DeformableConvRT(chunk_dim,kh,kw,sh,sw,ph,pw,deformableGroup,i_n,i_c,i_h,i_w,o_n,o_c,o_h,o_w,defRT);
int o_w = *(static_cast<const int *>(fields[15].data));
std::vector<dnnType> mask_v(static_cast<const dnnType*>(fields[16].data),static_cast<const dnnType*>(fields[16].data)+fields[16].length);
std::vector<dnnType> offset_v(static_cast<const dnnType*>(fields[17].data),static_cast<const dnnType*>(fields[17].data)+fields[17].length);
std::vector<dnnType> ones_d2_v(static_cast<const dnnType*>(fields[18].data),static_cast<const dnnType*>(fields[18].data)+fields[18].length);
std::vector<dnnType> ones_d1_v(static_cast<const dnnType*>(fields[19].data),static_cast<const dnnType*>(fields[19].data)+fields[19].length);
std::vector<dnnType> data_d_v(static_cast<const dnnType*>(fields[20].data),static_cast<const dnnType*>(fields[20].data)+fields[20].length);
std::vector<dnnType> bias2_d_v(static_cast<const dnnType*>(fields[21].data),static_cast<const dnnType*>(fields[21].data)+fields[21].length);
int height_ones = *(static_cast<const int *>(fields[22].data));
int width_ones = *(static_cast<const int *>(fields[23].data));
int dim_ones = *(static_cast<const int *>(fields[24].data));
auto *pluginObj = new DeformableConvRT(chunk_dim,kh,kw,sh,sw,ph,pw,deformableGroup,i_n,i_c,i_h,i_w,o_n,o_c,o_h,o_w,data_d_v,bias2_d_v,ones_d1_v,ones_d2_v,offset_v,mask_v,height_ones,width_ones,dim_ones);
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
return pluginObj;
}
const char *DeformableConvRTPluginCreator::getPluginName() const NOEXCEPT {
return "DeformableConvRT_tkDNN";
return DEFORMABLECONVRT_PLUGIN_NAME;
}
const char *DeformableConvRTPluginCreator::getPluginVersion() const NOEXCEPT {
return "1";
return DEFORMABLECONVRT_PLUGIN_VERSION;
}
const PluginFieldCollection *DeformableConvRTPluginCreator::getFieldNames() NOEXCEPT {