TRT8 works with almost every nerual network now!!!!(including demo3d)
This commit is contained in:
+2
-2
@@ -17,8 +17,8 @@ Flatten::Flatten(Network *net) : Layer(net) {
|
||||
|
||||
this->h = 1;
|
||||
this->w = 1;
|
||||
this->rows = input_dim.w;
|
||||
this->cols = input_dim.h * input_dim.c;
|
||||
this->rows = input_dim.c;
|
||||
this->cols = input_dim.h * input_dim.w;
|
||||
this->c = input_dim.w * input_dim.h * input_dim.c;
|
||||
}
|
||||
|
||||
|
||||
+144
-272
@@ -18,9 +18,9 @@ using namespace nvinfer1;
|
||||
// Logger for info/warning/errors
|
||||
class Logger : public ILogger {
|
||||
void log(Severity severity, const char* msg) NOEXCEPT override {
|
||||
//#ifdef DEBUG
|
||||
#ifdef DEBUG
|
||||
std::cout <<"TENSORRT LOG: "<< msg << std::endl;
|
||||
//#endif
|
||||
#endif
|
||||
}
|
||||
} loggerRT;
|
||||
|
||||
@@ -472,18 +472,36 @@ ILayer* NetworkRT::convert_layer(ITensor *input, Route *l) {
|
||||
}
|
||||
|
||||
ILayer* NetworkRT::convert_layer(ITensor *input, Flatten *l) {
|
||||
auto creator = getPluginRegistry()->getPluginCreator("FlattenConcatRT_tkDNN","1");
|
||||
std::vector<PluginField> mPluginAttributes;
|
||||
PluginFieldCollection mFC{};
|
||||
mPluginAttributes.emplace_back(PluginField("c",&l->c,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("h",&l->h,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("w",&l->w,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("rows",&l->rows,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("cols",&l->cols,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
|
||||
IPluginV2IOExt *plugin = new FlattenConcatRT(l->c,l->h,l->w,l->rows,l->cols);
|
||||
IPluginV2Layer *lRT = networkRT->addPluginV2(&input, 1, *plugin);
|
||||
auto *plugin = creator->createPlugin(l->getLayerName().c_str(),&mFC);
|
||||
auto *lRT = networkRT->addPluginV2(&input, 1, *plugin);
|
||||
checkNULL(lRT);
|
||||
return lRT;
|
||||
}
|
||||
|
||||
ILayer* NetworkRT::convert_layer(ITensor *input, Reshape *l) {
|
||||
// std::cout<<"convert Reshape\n";
|
||||
|
||||
IPluginV2 *plugin = new ReshapeRT(l->output_dim);
|
||||
IPluginV2Layer *lRT = networkRT->addPluginV2(&input, 1, *plugin);
|
||||
auto creator = getPluginRegistry()->getPluginCreator("ReshapeRT_tkDNN","1");
|
||||
std::vector<PluginField> mPluginAttributes;
|
||||
PluginFieldCollection mFC{};
|
||||
mPluginAttributes.emplace_back(PluginField("n",&l->n,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("c",&l->c,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("h",&l->h,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("w",&l->w,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
auto *plugin = creator->createPlugin(l->getLayerName().c_str(),&mFC);
|
||||
auto *lRT = networkRT->addPluginV2(&input, 1, *plugin);
|
||||
checkNULL(lRT);
|
||||
return lRT;
|
||||
}
|
||||
@@ -503,8 +521,17 @@ ILayer* NetworkRT::convert_layer(ITensor *input, Reorg *l) {
|
||||
//std::cout<<"convert Reorg\n";
|
||||
|
||||
//std::cout<<"New plugin REORG\n";
|
||||
IPluginV2 *plugin = new ReorgRT(l->stride);
|
||||
IPluginV2Layer *lRT = networkRT->addPluginV2(&input, 1, *plugin);
|
||||
auto creator = getPluginRegistry()->getPluginCreator("ReorgRT_tkDNN","1");
|
||||
std::vector<PluginField> mPluginAttributes;
|
||||
PluginFieldCollection mFC{};
|
||||
mPluginAttributes.emplace_back(PluginField("stride",&l->stride,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("c",&l->input_dim.c,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("h",&l->input_dim.h,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("w",&l->input_dim.w,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
auto *plugin = creator->createPlugin(l->getLayerName().c_str(),&mFC);
|
||||
auto *lRT = networkRT->addPluginV2(&input, 1, *plugin);
|
||||
checkNULL(lRT);
|
||||
return lRT;
|
||||
}
|
||||
@@ -513,8 +540,19 @@ ILayer* NetworkRT::convert_layer(ITensor *input, Region *l) {
|
||||
//std::cout<<"convert Region\n";
|
||||
|
||||
//std::cout<<"New plugin REGION\n";
|
||||
IPluginV2 *plugin = new RegionRT(l->classes, l->coords, l->num);
|
||||
IPluginV2Layer *lRT = networkRT->addPluginV2(&input, 1, *plugin);
|
||||
auto creator = getPluginRegistry()->getPluginCreator("RegionRT_tkDNN","1");
|
||||
std::vector<PluginField> mPluginAttributes;
|
||||
PluginFieldCollection mFC{};
|
||||
mPluginAttributes.emplace_back(PluginField("classes",&l->classes,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("coords",&l->coords,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("nums",&l->num,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("c",&l->input_dim.c,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("h",&l->input_dim.h,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("w",&l->input_dim.w,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
auto *plugin = creator->createPlugin(l->getLayerName().c_str(),&mFC);
|
||||
auto *lRT = networkRT->addPluginV2(&input, 1, *plugin);
|
||||
checkNULL(lRT);
|
||||
return lRT;
|
||||
}
|
||||
@@ -535,22 +573,52 @@ ILayer* NetworkRT::convert_layer(ITensor *input, Shortcut *l) {
|
||||
else
|
||||
{
|
||||
// plugin version
|
||||
IPluginV2 *plugin = new ShortcutRT(l->backLayer->output_dim, l->mul);
|
||||
ITensor **inputs = new ITensor*[2];
|
||||
auto creator = getPluginRegistry()->getPluginCreator("ShortcutRT_tkDNN","1");
|
||||
std::vector<PluginField> mPluginAttributes;
|
||||
PluginFieldCollection mFC{};
|
||||
mPluginAttributes.emplace_back(PluginField("bc",&l->backLayer->output_dim.c,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("bh",&l->backLayer->output_dim.h,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("bw",&l->backLayer->output_dim.w,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("mul",&l->mul,PluginFieldType::kUNKNOWN,1));
|
||||
mPluginAttributes.emplace_back(PluginField("c",&l->c,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("h",&l->h,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("w",&l->w,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
auto *plugin = creator->createPlugin(l->getLayerName().c_str(),&mFC);
|
||||
auto **inputs = new ITensor*[2];
|
||||
inputs[0] = input;
|
||||
inputs[1] = back_tens;
|
||||
IPluginV2Layer *lRT = networkRT->addPluginV2(inputs, 2, *plugin);
|
||||
auto *lRT = networkRT->addPluginV2(inputs, 2, *plugin);
|
||||
checkNULL(lRT);
|
||||
return lRT;
|
||||
}
|
||||
}
|
||||
|
||||
ILayer* NetworkRT::convert_layer(ITensor *input, Yolo *l) {
|
||||
//std::cout<<"convert Yolo\n";
|
||||
|
||||
//std::cout<<"New plugin YOLO\n";
|
||||
IPluginV2 *plugin = new YoloRT(l->classes, l->num, l, l->n_masks, l->scaleXY, l->nms_thresh, l->nsm_kind, l->new_coords);
|
||||
IPluginV2Layer *lRT = networkRT->addPluginV2(&input, 1, *plugin);
|
||||
std::vector<dnnType> mask_h(l->mask_h,l->mask_h+sizeof(dnnType)*l->n_masks);
|
||||
std::vector<dnnType> bias_h(l->bias_h,l->bias_h+sizeof(dnnType)*2*l->n_masks*l->num);
|
||||
auto creator = getPluginRegistry()->getPluginCreator("YoloRT_tkDNN","1");
|
||||
std::vector<PluginField> mPluginAttributes;
|
||||
PluginFieldCollection mFC{};
|
||||
mPluginAttributes.emplace_back(PluginField("classes",&l->classes,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("num",&l->num,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("c",&l->input_dim.c,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("h",&l->input_dim.h,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("w",&l->input_dim.w,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("classNames",&l->classesNames[0],PluginFieldType::kUNKNOWN,l->classesNames.size()));
|
||||
mPluginAttributes.emplace_back(PluginField("mask_v",&mask_h[0],PluginFieldType::kFLOAT32,mask_h.size()));
|
||||
mPluginAttributes.emplace_back(PluginField("bias_v",&bias_h[0],PluginFieldType::kFLOAT32,bias_h.size()));
|
||||
mPluginAttributes.emplace_back(PluginField("n_masks",&l->n_masks,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("scale_xy",&l->scaleXY,PluginFieldType::kFLOAT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("nms_thresh",&l->nms_thresh,PluginFieldType::kFLOAT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("nms_kins",&l->nsm_kind,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("new_coords",&l->new_coords,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
auto *plugin = creator->createPlugin(l->getLayerName().c_str(),&mFC);
|
||||
auto *lRT = networkRT->addPluginV2(&input, 1, *plugin);
|
||||
checkNULL(lRT);
|
||||
return lRT;
|
||||
}
|
||||
@@ -558,9 +626,17 @@ ILayer* NetworkRT::convert_layer(ITensor *input, Yolo *l) {
|
||||
ILayer* NetworkRT::convert_layer(ITensor *input, Upsample *l) {
|
||||
//std::cout<<"convert Upsample\n";
|
||||
|
||||
std::cout<<"New plugin UPSAMPLE\n";
|
||||
IPluginV2 *plugin = new UpsampleRT(l->stride);
|
||||
IPluginV2Layer *lRT = networkRT->addPluginV2(&input, 1, *plugin);
|
||||
auto creator = getPluginRegistry()->getPluginCreator("UpSample_tkDNN","1");
|
||||
std::vector<PluginField> mPluginAttributes;
|
||||
PluginFieldCollection mFC{};
|
||||
mPluginAttributes.emplace_back(PluginField("stride",&l->stride,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("c",&l->c,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("h",&l->h,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("w",&l->w,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
auto *plugin = creator->createPlugin(l->getLayerName().c_str(),&mFC);
|
||||
auto *lRT = networkRT->addPluginV2(&input, 1, *plugin);
|
||||
checkNULL(lRT);
|
||||
return lRT;
|
||||
}
|
||||
@@ -575,10 +651,53 @@ ILayer* NetworkRT::convert_layer(ITensor *input, DeformConv2d *l) {
|
||||
inputs[1] = preconv->getOutput(0);
|
||||
|
||||
//std::cout<<"New plugin DEFORMABLE\n";
|
||||
IPluginV2 *plugin = new DeformableConvRT(l->chunk_dim, l->kernelH, l->kernelW, l->strideH, l->strideW, l->paddingH, l->paddingW,
|
||||
l->deformableGroup, l->input_dim.n, l->input_dim.c, l->input_dim.h, l->input_dim.w,
|
||||
l->output_dim.n, l->output_dim.c, l->output_dim.h, l->output_dim.w, l);
|
||||
IPluginV2Layer *lRT = networkRT->addPluginV2(inputs, 2, *plugin);
|
||||
int height_ones = (l->input_dim.h + 2 * l->paddingH - (1 * (l->kernelH - 1) + 1)) / l->strideH + 1;
|
||||
int width_ones = (l->input_dim.w + 2 * l->paddingW - (1 * (l->kernelW - 1) + 1)) / l->strideW + 1;
|
||||
int dim_ones = l->input_dim.c * l->kernelH * l->kernelW * 1 * height_ones * width_ones;
|
||||
std::vector<dnnType> offsetV(2*l->chunk_dim);
|
||||
std::vector<dnnType> maskV(l->chunk_dim);
|
||||
std::vector<dnnType> dataV(l->input_dim.c*l->output_dim.c*l->kernelW*l->kernelH*1);
|
||||
std::vector<dnnType> bias2DV(l->output_dim.c);
|
||||
std::vector<dnnType> onesD1V(height_ones*width_ones);
|
||||
std::vector<dnnType> onesD2V(dim_ones);
|
||||
checkCuda(cudaMemcpy(offsetV.data(),l->offset,offsetV.size()*sizeof(dnnType),cudaMemcpyDeviceToHost));
|
||||
checkCuda(cudaMemcpy(maskV.data(),l->mask,sizeof(dnnType)*maskV.size(),cudaMemcpyDeviceToHost));
|
||||
checkCuda(cudaMemcpy(dataV.data(),l->data_d,sizeof(dnnType)*dataV.size(),cudaMemcpyDeviceToHost));
|
||||
checkCuda(cudaMemcpy(bias2DV.data(),l->bias2_d,sizeof(dnnType)*bias2DV.size(),cudaMemcpyDeviceToHost));
|
||||
checkCuda(cudaMemcpy(onesD1V.data(),l->ones_d1,sizeof(dnnType)*onesD1V.size(),cudaMemcpyDeviceToHost));
|
||||
checkCuda(cudaMemcpy(onesD2V.data(),l->ones_d2,sizeof(dnnType)*onesD2V.size(),cudaMemcpyDeviceToHost));
|
||||
auto creator = getPluginRegistry()->getPluginCreator("DeformableConvRT_tkDNN","1");
|
||||
std::vector<PluginField> mPluginAttributes;
|
||||
PluginFieldCollection mFC{};
|
||||
mPluginAttributes.emplace_back(PluginField("chunk_dum",&l->chunk_dim,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("kh",&l->kernelH,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("kw",&l->kernelW,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("sh",&l->strideH,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("sw",&l->strideW,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("ph",&l->paddingH,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("pw",&l->paddingW,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("deformable_group",&l->deformableGroup,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("i_n",&l->input_dim.n,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("i_c",&l->input_dim.c,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("i_h",&l->input_dim.h,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("i_w",&l->input_dim.w,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("o_n",&l->output_dim.n,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("o_c",&l->output_dim.c,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("o_h",&l->output_dim.h,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("o_w",&l->output_dim.w,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("mask_v",&maskV[0],PluginFieldType::kFLOAT32,maskV.size()));
|
||||
mPluginAttributes.emplace_back(PluginField("offset_v",&offsetV[0],PluginFieldType::kFLOAT32,offsetV.size()));
|
||||
mPluginAttributes.emplace_back(PluginField("ones_d2_v",&onesD2V[0],PluginFieldType::kFLOAT32,onesD2V.size()));
|
||||
mPluginAttributes.emplace_back(PluginField("ones_d1_v",&onesD1V[0],PluginFieldType::kFLOAT32,onesD1V.size()));
|
||||
mPluginAttributes.emplace_back(PluginField("data_d_v",&dataV[0],PluginFieldType::kFLOAT32,dataV.size()));
|
||||
mPluginAttributes.emplace_back(PluginField("bias2_d_v",&bias2DV[0],PluginFieldType::kFLOAT32,bias2DV.size()));
|
||||
mPluginAttributes.emplace_back(PluginField("height_ones",&height_ones,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("width_ones",&width_ones,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("dim_ones",&dim_ones,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
auto *plugin = creator->createPlugin(l->getLayerName().c_str(),&mFC);
|
||||
auto *lRT = networkRT->addPluginV2(inputs, 2, *plugin);
|
||||
checkNULL(lRT);
|
||||
lRT->setName( ("Deformable" + std::to_string(l->id)).c_str() );
|
||||
delete[](inputs);
|
||||
@@ -658,254 +777,7 @@ bool NetworkRT::deserialize(const char *filename) {
|
||||
void NetworkRT::destroy() {
|
||||
contextRT->destroy();
|
||||
engineRT->destroy();
|
||||
configRT->destroy();
|
||||
builderRT->destroy();
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
/*
|
||||
IPlugin* PluginFactory::createPlugin(const char* layerName, const void* serialData, size_t serialLength) {
|
||||
const char * buf = reinterpret_cast<const char*>(serialData),*bufCheck = buf;
|
||||
|
||||
std::string name(layerName);
|
||||
//std::cout<<name<<std::endl;
|
||||
|
||||
if(name.find("ActivationLeaky") == 0) {
|
||||
ActivationLeakyRT *a = new ActivationLeakyRT(readBUF<float>(buf));
|
||||
a->size = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + serialLength);
|
||||
return a;
|
||||
}
|
||||
if(name.find("ActivationMish") == 0) {
|
||||
ActivationMishRT *a = new ActivationMishRT();
|
||||
a->size = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + serialLength);
|
||||
return a;
|
||||
}
|
||||
if(name.find("ActivationLogistic") == 0) {
|
||||
ActivationLogisticRT *a = new ActivationLogisticRT();
|
||||
a->size = readBUF<int>(buf);
|
||||
return a;
|
||||
}
|
||||
if(name.find("ActivationLogistic") == 0) {
|
||||
ActivationLogisticRT *a = new ActivationLogisticRT();
|
||||
a->size = readBUF<int>(buf);
|
||||
return a;
|
||||
}
|
||||
if(name.find("ActivationCReLU") == 0) {
|
||||
float activationReluTemp = readBUF<float>(buf);
|
||||
ActivationReLUCeiling* a = new ActivationReLUCeiling(activationReluTemp);
|
||||
a->size = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + serialLength);
|
||||
return a;
|
||||
}
|
||||
|
||||
if(name.find("Region") == 0) {
|
||||
int classesTemp = readBUF<int>(buf);
|
||||
int coordsTemp = readBUF<int>(buf);
|
||||
int numTemp = readBUF<int>(buf);
|
||||
RegionRT* r = new RegionRT(classesTemp, coordsTemp, numTemp);
|
||||
|
||||
r->c = readBUF<int>(buf);
|
||||
r->h = readBUF<int>(buf);
|
||||
r->w = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + serialLength);
|
||||
return r;
|
||||
}
|
||||
|
||||
if(name.find("Reorg") == 0) {
|
||||
int strideTemp = readBUF<int>(buf);
|
||||
ReorgRT *r = new ReorgRT(strideTemp);
|
||||
r->c = readBUF<int>(buf);
|
||||
r->h = readBUF<int>(buf);
|
||||
r->w = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + serialLength);
|
||||
return r;
|
||||
}
|
||||
|
||||
if(name.find("Shortcut") == 0) {
|
||||
tk::dnn::dataDim_t bdim;
|
||||
bdim.c = readBUF<int>(buf);
|
||||
bdim.h = readBUF<int>(buf);
|
||||
bdim.w = readBUF<int>(buf);
|
||||
bdim.l = 1;
|
||||
|
||||
ShortcutRT *r = new ShortcutRT(bdim, readBUF<bool>(buf));
|
||||
r->c = readBUF<int>(buf);
|
||||
r->h = readBUF<int>(buf);
|
||||
r->w = readBUF<int>(buf);
|
||||
return r;
|
||||
assert(buf == bufCheck + serialLength);
|
||||
}
|
||||
|
||||
if(name.find("Pooling") == 0) {
|
||||
int cTemp = readBUF<int>(buf);
|
||||
int hTemp = readBUF<int>(buf);
|
||||
int wTemp = readBUF<int>(buf);
|
||||
int nTemp = readBUF<int>(buf);
|
||||
int strideHTemp = readBUF<int>(buf);
|
||||
int strideWTemp = readBUF<int>(buf);
|
||||
int winSizeTemp = readBUF<int>(buf);
|
||||
int paddingTemp = readBUF<int>(buf);
|
||||
|
||||
MaxPoolFixedSizeRT* r = new MaxPoolFixedSizeRT(cTemp, hTemp, wTemp, nTemp, strideHTemp, strideWTemp, winSizeTemp, paddingTemp);
|
||||
assert(buf == bufCheck + serialLength);
|
||||
return r;
|
||||
}
|
||||
|
||||
if(name.find("Resize") == 0) {
|
||||
int o_cTemp = readBUF<int>(buf);
|
||||
int o_hTemp = readBUF<int>(buf);
|
||||
int o_wTemp = readBUF<int>(buf);
|
||||
ResizeLayerRT* r = new ResizeLayerRT(o_cTemp, o_hTemp, o_wTemp);
|
||||
|
||||
r->i_c = readBUF<int>(buf);
|
||||
r->i_h = readBUF<int>(buf);
|
||||
r->i_w = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + serialLength);
|
||||
return r;
|
||||
}
|
||||
|
||||
if(name.find("Flatten") == 0) {
|
||||
FlattenConcatRT *r = new FlattenConcatRT();
|
||||
r->c = readBUF<int>(buf);
|
||||
r->h = readBUF<int>(buf);
|
||||
r->w = readBUF<int>(buf);
|
||||
r->rows = readBUF<int>(buf);
|
||||
r->cols = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + serialLength);
|
||||
return r;
|
||||
}
|
||||
|
||||
if(name.find("Reshape") == 0) {
|
||||
|
||||
dataDim_t new_dim;
|
||||
new_dim.n = readBUF<int>(buf);
|
||||
new_dim.c = readBUF<int>(buf);
|
||||
new_dim.h = readBUF<int>(buf);
|
||||
new_dim.w = readBUF<int>(buf);
|
||||
ReshapeRT *r = new ReshapeRT(new_dim);
|
||||
assert(buf == bufCheck + serialLength);
|
||||
|
||||
return r;
|
||||
}
|
||||
|
||||
if(name.find("Yolo") == 0) {
|
||||
|
||||
int classes_temp = readBUF<int>(buf);
|
||||
int num_temp = readBUF<int>(buf);
|
||||
int n_masks_temp = readBUF<int>(buf);
|
||||
float scale_xy_temp = readBUF<float>(buf);
|
||||
float nms_thresh_temp = readBUF<float>(buf);
|
||||
int nms_kind_temp = readBUF<int>(buf);
|
||||
int new_coords_temp = readBUF<int>(buf);
|
||||
|
||||
YoloRT *r = new YoloRT(classes_temp,num_temp,nullptr,n_masks_temp,scale_xy_temp,nms_thresh_temp,nms_kind_temp,new_coords_temp);
|
||||
|
||||
|
||||
|
||||
r->c = readBUF<int>(buf);
|
||||
r->h = readBUF<int>(buf);
|
||||
r->w = readBUF<int>(buf);
|
||||
for(int i=0; i<r->n_masks; i++)
|
||||
r->mask[i] = readBUF<dnnType>(buf);
|
||||
for(int i=0; i<r->n_masks*2*r->num; i++)
|
||||
r->bias[i] = readBUF<dnnType>(buf);
|
||||
|
||||
// save classes names
|
||||
r->classesNames.resize(r->classes);
|
||||
for(int i=0; i<r->classes; i++) {
|
||||
char tmp[YOLORT_CLASSNAME_W];
|
||||
for(int j=0; j<YOLORT_CLASSNAME_W; j++)
|
||||
tmp[j] = readBUF<char>(buf);
|
||||
r->classesNames[i] = std::string(tmp);
|
||||
}
|
||||
assert(buf == bufCheck + serialLength);
|
||||
|
||||
yolos[n_yolos++] = r;
|
||||
return r;
|
||||
}
|
||||
if(name.find("Upsample") == 0) {
|
||||
int strideTemp = readBUF<int>(buf);
|
||||
UpsampleRT* r = new UpsampleRT(strideTemp);
|
||||
r->c = readBUF<int>(buf);
|
||||
r->h = readBUF<int>(buf);
|
||||
r->w = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + serialLength);
|
||||
return r;
|
||||
}
|
||||
|
||||
if(name.find("Route") == 0) {
|
||||
int groupsTemp = readBUF<int>(buf);
|
||||
int group_idTemp = readBUF<int>(buf);
|
||||
RouteRT* r = new RouteRT(groupsTemp, group_idTemp);
|
||||
r->in = readBUF<int>(buf);
|
||||
for(int i=0; i<RouteRT::MAX_INPUTS; i++)
|
||||
r->c_in[i] = readBUF<int>(buf);
|
||||
r->c = readBUF<int>(buf);
|
||||
r->h = readBUF<int>(buf);
|
||||
r->w = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + serialLength);
|
||||
return r;
|
||||
}
|
||||
|
||||
if(name.find("Deformable") == 0) {
|
||||
int chuck_dimTemp = readBUF<int>(buf);
|
||||
int khTemp = readBUF<int>(buf);
|
||||
int kwTemp = readBUF<int>(buf);
|
||||
int shTemp = readBUF<int>(buf);
|
||||
int swTemp = readBUF<int>(buf);
|
||||
int phTemp = readBUF<int>(buf);
|
||||
int pwTemp = readBUF<int>(buf);
|
||||
int deformableGroupTemp = readBUF<int>(buf);
|
||||
int i_nTemp = readBUF<int>(buf);
|
||||
int i_cTemp = readBUF<int>(buf);
|
||||
int i_hTemp = readBUF<int>(buf);
|
||||
int i_wTemp = readBUF<int>(buf);
|
||||
int o_nTemp = readBUF<int>(buf);
|
||||
int o_cTemp = readBUF<int>(buf);
|
||||
int o_hTemp = readBUF<int>(buf);
|
||||
int o_wTemp = readBUF<int>(buf);
|
||||
|
||||
DeformableConvRT* r = new DeformableConvRT(chuck_dimTemp, khTemp, kwTemp, shTemp, swTemp, phTemp, pwTemp, deformableGroupTemp, i_nTemp, i_cTemp, i_hTemp, i_wTemp, o_nTemp, o_cTemp, o_hTemp, o_wTemp, nullptr);
|
||||
dnnType *aus = new dnnType[r->chunk_dim*2];
|
||||
for(int i=0; i<r->chunk_dim*2; i++)
|
||||
aus[i] = readBUF<dnnType>(buf);
|
||||
checkCuda( cudaMemcpy(r->offset, aus, sizeof(dnnType)*2*r->chunk_dim, cudaMemcpyHostToDevice) );
|
||||
free(aus);
|
||||
aus = new dnnType[r->chunk_dim];
|
||||
for(int i=0; i<r->chunk_dim; i++)
|
||||
aus[i] = readBUF<dnnType>(buf);
|
||||
checkCuda( cudaMemcpy(r->mask, aus, sizeof(dnnType)*r->chunk_dim, cudaMemcpyHostToDevice) );
|
||||
free(aus);
|
||||
aus = new dnnType[(r->i_c * r->o_c * r->kh * r->kw * 1 )];
|
||||
for(int i=0; i<(r->i_c * r->o_c * r->kh * r->kw * 1 ); i++)
|
||||
aus[i] = readBUF<dnnType>(buf);
|
||||
checkCuda( cudaMemcpy(r->data_d, aus, sizeof(dnnType)*(r->i_c * r->o_c * r->kh * r->kw * 1 ), cudaMemcpyHostToDevice) );
|
||||
free(aus);
|
||||
aus = new dnnType[r->o_c];
|
||||
for(int i=0; i < r->o_c; i++)
|
||||
aus[i] = readBUF<dnnType>(buf);
|
||||
checkCuda( cudaMemcpy(r->bias2_d, aus, sizeof(dnnType)*r->o_c, cudaMemcpyHostToDevice) );
|
||||
free(aus);
|
||||
aus = new dnnType[r->height_ones * r->width_ones];
|
||||
for(int i=0; i<r->height_ones * r->width_ones; i++)
|
||||
aus[i] = readBUF<dnnType>(buf);
|
||||
checkCuda( cudaMemcpy(r->ones_d1, aus, sizeof(dnnType)*r->height_ones * r->width_ones, cudaMemcpyHostToDevice) );
|
||||
free(aus);
|
||||
aus = new dnnType[r->dim_ones];
|
||||
for(int i=0; i<r->dim_ones; i++)
|
||||
aus[i] = readBUF<dnnType>(buf);
|
||||
checkCuda( cudaMemcpy(r->ones_d2, aus, sizeof(dnnType)*r->dim_ones, cudaMemcpyHostToDevice) );
|
||||
free(aus);
|
||||
assert(buf == bufCheck + serialLength);
|
||||
return r;
|
||||
}
|
||||
|
||||
FatalError("Cant deserialize Plugin");
|
||||
return NULL;
|
||||
}
|
||||
*/
|
||||
}}
|
||||
|
||||
@@ -16,7 +16,6 @@ Region::Region(Network *net, int classes, int coords, int num) :
|
||||
this->classes = classes;
|
||||
this->coords = coords;
|
||||
this->num = num;
|
||||
|
||||
// same
|
||||
output_dim.n = input_dim.n;
|
||||
output_dim.c = input_dim.c;
|
||||
|
||||
+4
-1
@@ -8,7 +8,10 @@ namespace tk { namespace dnn {
|
||||
Reshape::Reshape(Network *net, dataDim_t new_dim) : Layer(net) {
|
||||
|
||||
checkCuda( cudaMalloc(&dstData, input_dim.tot()*sizeof(dnnType)) );
|
||||
|
||||
this->n = new_dim.n;
|
||||
this->c = new_dim.c;
|
||||
this->h = new_dim.h;
|
||||
this->w = new_dim.w;
|
||||
output_dim.n = new_dim.n;
|
||||
output_dim.c = new_dim.c;
|
||||
output_dim.h = new_dim.h;
|
||||
|
||||
@@ -9,6 +9,9 @@ Shortcut::Shortcut(Network *net, Layer *backLayer, bool mul) : Layer(net) {
|
||||
|
||||
this->backLayer = backLayer;
|
||||
this->mul = mul;
|
||||
this->c = input_dim.c;
|
||||
this->h = input_dim.h;
|
||||
this->w = input_dim.w;
|
||||
checkCuda( cudaMalloc(&dstData, output_dim.tot()*sizeof(dnnType)) );
|
||||
|
||||
if( ( backLayer->output_dim.c != input_dim.c && mul ) ||
|
||||
|
||||
@@ -14,6 +14,9 @@ Upsample::Upsample(Network *net, int stride) : Layer(net) {
|
||||
output_dim.h = input_dim.h*stride;
|
||||
output_dim.w = input_dim.w*stride;
|
||||
output_dim.l = input_dim.l;
|
||||
this->c = input_dim.c;
|
||||
this->h = input_dim.h;
|
||||
this->w = input_dim.w;
|
||||
|
||||
checkCuda( cudaMalloc(&dstData, output_dim.tot()*sizeof(dnnType)) );
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -4,12 +4,10 @@ using namespace nvinfer1;
|
||||
std::vector<PluginField> FlattenConcatRTPluginCreator::mPluginAttributes;
|
||||
PluginFieldCollection FlattenConcatRTPluginCreator::mFC{};
|
||||
|
||||
static const char* FLATTENCONCATRT_PLUGIN_VERSION{"1"};
|
||||
static const char* FLATTENCONCATRT_PLUGIN_NAME{"FlattenConcatRT_tkDNN"};
|
||||
|
||||
FlattenConcatRT::FlattenConcatRT(int c, int h, int w, int rows, int cols) {
|
||||
stat = cublasCreate(&handle);
|
||||
if (stat != CUBLAS_STATUS_SUCCESS) {
|
||||
printf ("CUBLAS initialization failed\n");
|
||||
return;
|
||||
}
|
||||
this->c = c;
|
||||
this->h = h;
|
||||
this->w = w;
|
||||
@@ -42,7 +40,7 @@ int FlattenConcatRT::initialize() NOEXCEPT {
|
||||
}
|
||||
|
||||
void FlattenConcatRT::terminate() NOEXCEPT {
|
||||
checkERROR(cublasDestroy(handle));
|
||||
|
||||
}
|
||||
|
||||
size_t FlattenConcatRT::getWorkspaceSize(int maxBatchSize) const NOEXCEPT {
|
||||
@@ -105,11 +103,11 @@ void FlattenConcatRT::destroy() NOEXCEPT {
|
||||
|
||||
|
||||
const char *FlattenConcatRT::getPluginType() const NOEXCEPT {
|
||||
return "FlattenConcatRT_tkDNN";
|
||||
return FLATTENCONCATRT_PLUGIN_NAME;
|
||||
}
|
||||
|
||||
const char *FlattenConcatRT::getPluginVersion() const NOEXCEPT {
|
||||
return "1";
|
||||
return FLATTENCONCATRT_PLUGIN_VERSION;
|
||||
}
|
||||
|
||||
const char *FlattenConcatRT::getPluginNamespace() const NOEXCEPT {
|
||||
@@ -120,7 +118,7 @@ void FlattenConcatRT::setPluginNamespace(const char *pluginNamespace) NOEXCEPT {
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
IPluginV2IOExt *FlattenConcatRT::clone() const NOEXCEPT {
|
||||
IPluginV2Ext *FlattenConcatRT::clone() const NOEXCEPT {
|
||||
auto* p = new FlattenConcatRT(c, h, w, rows, cols);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
@@ -131,12 +129,10 @@ DataType FlattenConcatRT::getOutputDataType(int index, const nvinfer1::DataType*
|
||||
return DataType::kFLOAT;
|
||||
}
|
||||
|
||||
void FlattenConcatRT::configurePlugin(const PluginTensorDesc* in, int nbInput, const PluginTensorDesc* out, int nbOutput) NOEXCEPT
|
||||
{
|
||||
}
|
||||
|
||||
void FlattenConcatRT::attachToContext(cudnnContext* cudnnContext, cublasContext* cublasContext, IGpuAllocator* gpuAllocator) NOEXCEPT
|
||||
{
|
||||
handle = cublasContext;
|
||||
}
|
||||
|
||||
bool FlattenConcatRT::isOutputBroadcastAcrossBatch(int outputIndex, const bool* inputIsBroadcasted, int nbInputs) const NOEXCEPT
|
||||
@@ -149,17 +145,29 @@ bool FlattenConcatRT::canBroadcastInputAcrossBatch(int inputIndex) const NOEXCEP
|
||||
return false;
|
||||
}
|
||||
|
||||
bool FlattenConcatRT::supportsFormatCombination(int pos, const PluginTensorDesc* inOut, int nbInputs, int nbOutputs) const NOEXCEPT
|
||||
{
|
||||
return true;
|
||||
}
|
||||
|
||||
void FlattenConcatRT::detachFromContext() NOEXCEPT
|
||||
{
|
||||
}
|
||||
|
||||
void
|
||||
FlattenConcatRT::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 {
|
||||
|
||||
}
|
||||
|
||||
bool FlattenConcatRT::supportsFormat(DataType type, PluginFormat format) const NOEXCEPT {
|
||||
return true;
|
||||
}
|
||||
|
||||
FlattenConcatRTPluginCreator::FlattenConcatRTPluginCreator() {
|
||||
mPluginAttributes.clear();
|
||||
mPluginAttributes.emplace_back(PluginField("c", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("h", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("w", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("rows", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("cols", nullptr,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
@@ -172,14 +180,14 @@ const char *FlattenConcatRTPluginCreator::getPluginNamespace() const NOEXCEPT {
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2IOExt *FlattenConcatRTPluginCreator::deserializePlugin(const char *name, const void *serialData,
|
||||
IPluginV2Ext *FlattenConcatRTPluginCreator::deserializePlugin(const char *name, const void *serialData,
|
||||
size_t serialLength) NOEXCEPT {
|
||||
auto *pluginObj = new FlattenConcatRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2IOExt *FlattenConcatRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
IPluginV2Ext *FlattenConcatRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
const PluginField* fields = fc->fields;
|
||||
int c = *(static_cast<const int*>(fields[0].data));
|
||||
int h = *(static_cast<const int*>(fields[1].data));
|
||||
@@ -192,11 +200,11 @@ IPluginV2IOExt *FlattenConcatRTPluginCreator::createPlugin(const char *name, con
|
||||
}
|
||||
|
||||
const char *FlattenConcatRTPluginCreator::getPluginName() const NOEXCEPT {
|
||||
return "FlattenConcatRT_tkDNN";
|
||||
return FLATTENCONCATRT_PLUGIN_NAME;
|
||||
}
|
||||
|
||||
const char *FlattenConcatRTPluginCreator::getPluginVersion() const NOEXCEPT {
|
||||
return "1";
|
||||
return FLATTENCONCATRT_PLUGIN_VERSION;
|
||||
}
|
||||
|
||||
const PluginFieldCollection *FlattenConcatRTPluginCreator::getFieldNames() NOEXCEPT {
|
||||
|
||||
@@ -40,8 +40,6 @@ Dims MaxPoolFixedSizeRT::getOutputDimensions(int index, const Dims *inputs, int
|
||||
return Dims3{this->c, this->h, this->w};
|
||||
}
|
||||
|
||||
void MaxPoolFixedSizeRT::configureWithFormat(const Dims *inputDims, int nbInputs, const Dims *outputDims, int nbOutputs,DataType type, PluginFormat format, int maxBatchSize) NOEXCEPT {}
|
||||
|
||||
int MaxPoolFixedSizeRT::initialize() NOEXCEPT {
|
||||
return 0;
|
||||
}
|
||||
@@ -62,7 +60,7 @@ int MaxPoolFixedSizeRT::enqueue(int batchSize, const void *const *inputs, void *
|
||||
MaxPoolingForward(srcData, dstData, batchSize, this->c, this->h, this->w, this->stride_H, this->stride_W, this->winSize, this->padding, stream);
|
||||
return 0;
|
||||
}
|
||||
#elif NV_TENSORRT_MAJOR == 7
|
||||
#elif NV_TENSORRT_MAJOR <= 7
|
||||
int32_t MaxPoolFixedSizeRT::enqueue(int32_t batchSize, const void *const *inputs, void **outputs, void *workspace,
|
||||
cudaStream_t stream) {
|
||||
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
|
||||
@@ -115,12 +113,43 @@ const char *MaxPoolFixedSizeRT::getPluginVersion() const NOEXCEPT {
|
||||
return "1";
|
||||
}
|
||||
|
||||
IPluginV2 *MaxPoolFixedSizeRT::clone() const NOEXCEPT {
|
||||
IPluginV2Ext *MaxPoolFixedSizeRT::clone() const NOEXCEPT {
|
||||
auto *p = new MaxPoolFixedSizeRT(c,h,w,n,stride_H,stride_W,winSize,padding);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
DataType
|
||||
MaxPoolFixedSizeRT::getOutputDataType(int index, const nvinfer1::DataType *inputTypes, int nbInputs) const NOEXCEPT {
|
||||
return DataType::kFLOAT;
|
||||
}
|
||||
|
||||
void MaxPoolFixedSizeRT::attachToContext(cudnnContext *cudnnContext, cublasContext *cublasContext,
|
||||
IGpuAllocator *gpuAllocator) NOEXCEPT {
|
||||
|
||||
}
|
||||
|
||||
bool MaxPoolFixedSizeRT::isOutputBroadcastAcrossBatch(int outputIndex, const bool *inputIsBroadcasted,
|
||||
int nbInputs) const NOEXCEPT {
|
||||
return false;
|
||||
}
|
||||
|
||||
bool MaxPoolFixedSizeRT::canBroadcastInputAcrossBatch(int inputIndex) const NOEXCEPT {
|
||||
return false;
|
||||
}
|
||||
|
||||
void
|
||||
MaxPoolFixedSizeRT::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 MaxPoolFixedSizeRT::detachFromContext() NOEXCEPT {
|
||||
IPluginV2Ext::detachFromContext();
|
||||
}
|
||||
|
||||
MaxPoolFixedSizeRTPluginCreator::MaxPoolFixedSizeRTPluginCreator() {
|
||||
mPluginAttributes.clear();
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
@@ -135,15 +164,14 @@ const char *MaxPoolFixedSizeRTPluginCreator::getPluginNamespace() const NOEXCEPT
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *MaxPoolFixedSizeRTPluginCreator::deserializePlugin(const char *name, const void *serialData,size_t serialLength) NOEXCEPT {
|
||||
IPluginV2Ext *MaxPoolFixedSizeRTPluginCreator::deserializePlugin(const char *name, const void *serialData,size_t serialLength) NOEXCEPT {
|
||||
auto *pluginObj = new MaxPoolFixedSizeRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *MaxPoolFixedSizeRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
IPluginV2Ext *MaxPoolFixedSizeRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
const PluginField *fields = fc->fields;
|
||||
//todo assert
|
||||
int c = *(static_cast<const int *>(fields[0].data));
|
||||
int h = *(static_cast<const int *>(fields[1].data));
|
||||
int w = *(static_cast<const int *>(fields[2].data));
|
||||
|
||||
+58
-21
@@ -1,12 +1,19 @@
|
||||
#include <tkDNN/pluginsRT/RegionRT.h>
|
||||
using namespace nvinfer1;
|
||||
|
||||
std::vector<PluginField> RegionRTPluginCreator::mPluginAttributes;
|
||||
PluginFieldCollection RegionRTPluginCreator::mFC{};
|
||||
|
||||
RegionRT::RegionRT(int classes, int coords, int num) {
|
||||
static const char* REGIONRT_PLUGIN_VERSION{"1"};
|
||||
static const char* REGIONRT_PLUGIN_NAME{"RegionRT_tkDNN"};
|
||||
|
||||
RegionRT::RegionRT(int classes, int coords, int num,int c,int h,int w) {
|
||||
this->classes = classes;
|
||||
this->coords = coords;
|
||||
this->num = num;
|
||||
this->c = c;
|
||||
this->h = h;
|
||||
this->w = w;
|
||||
}
|
||||
|
||||
RegionRT::~RegionRT() {}
|
||||
@@ -30,12 +37,6 @@ Dims RegionRT::getOutputDimensions(int index, const Dims *inputs, int nbInputDim
|
||||
return inputs[0];
|
||||
}
|
||||
|
||||
void RegionRT::configureWithFormat(const Dims *inputDims, int nbInputs, const Dims *outputDims, int nbOutputs, DataType type,
|
||||
PluginFormat format, int maxBatchSize) NOEXCEPT {
|
||||
c = inputDims[0].d[0];
|
||||
h = inputDims[0].d[1];
|
||||
w = inputDims[0].d[2];
|
||||
}
|
||||
|
||||
int RegionRT::initialize() NOEXCEPT {return 0;}
|
||||
|
||||
@@ -112,11 +113,11 @@ void RegionRT::serialize(void *buffer) const NOEXCEPT {
|
||||
}
|
||||
|
||||
const char *RegionRT::getPluginType() const NOEXCEPT {
|
||||
return "RegionRT_tkDNN";
|
||||
return REGIONRT_PLUGIN_NAME;
|
||||
}
|
||||
|
||||
const char *RegionRT::getPluginVersion() const NOEXCEPT {
|
||||
return "1";
|
||||
return REGIONRT_PLUGIN_VERSION;
|
||||
}
|
||||
|
||||
void RegionRT::destroy() NOEXCEPT { delete this; }
|
||||
@@ -133,14 +134,48 @@ bool RegionRT::supportsFormat(DataType type, PluginFormat format) const NOEXCEPT
|
||||
return true;
|
||||
}
|
||||
|
||||
IPluginV2 *RegionRT::clone() const NOEXCEPT {
|
||||
auto *p = new RegionRT(classes,coords,num);
|
||||
IPluginV2Ext *RegionRT::clone() const NOEXCEPT {
|
||||
auto *p = new RegionRT(classes,coords,num,c,h,w);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
DataType RegionRT::getOutputDataType(int index, const nvinfer1::DataType *inputTypes, int nbInputs) const NOEXCEPT {
|
||||
return DataType::kFLOAT;
|
||||
}
|
||||
|
||||
void RegionRT::attachToContext(cudnnContext *cudnnContext, cublasContext *cublasContext,
|
||||
IGpuAllocator *gpuAllocator) NOEXCEPT {
|
||||
|
||||
}
|
||||
|
||||
bool RegionRT::isOutputBroadcastAcrossBatch(int outputIndex, const bool *inputIsBroadcasted, int nbInputs) const NOEXCEPT {
|
||||
return false;
|
||||
}
|
||||
|
||||
bool RegionRT::canBroadcastInputAcrossBatch(int inputIndex) const NOEXCEPT {
|
||||
return false;
|
||||
}
|
||||
|
||||
void RegionRT::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 RegionRT::detachFromContext() NOEXCEPT {
|
||||
|
||||
}
|
||||
|
||||
|
||||
RegionRTPluginCreator::RegionRTPluginCreator() {
|
||||
mPluginAttributes.clear();
|
||||
mPluginAttributes.emplace_back(PluginField("classes", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("coords", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("num", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("c", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("h", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("w", nullptr,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
@@ -153,32 +188,35 @@ const char *RegionRTPluginCreator::getPluginNamespace() const NOEXCEPT {
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *RegionRTPluginCreator::deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT {
|
||||
IPluginV2Ext *RegionRTPluginCreator::deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT {
|
||||
auto *pluginObj = new RegionRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *RegionRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
IPluginV2Ext *RegionRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
const PluginField *fields = fc->fields;
|
||||
assert(fc->nbFields == 3);
|
||||
assert(fields[0].type == PluginFieldType::kINT32);
|
||||
assert(fields[1].type == PluginFieldType::kINT32);
|
||||
assert(fields[2].type == PluginFieldType::kINT32);
|
||||
assert(fc->nbFields == 6);
|
||||
for(int i=0;i<6;i++){
|
||||
assert(fields[i].type == PluginFieldType::kINT32);
|
||||
}
|
||||
int classes = *(static_cast<const int*>(fields[0].data));
|
||||
int coords = *(static_cast<const int*>(fields[1].data));
|
||||
int num = *(static_cast<const int*>(fields[2].data));
|
||||
RegionRT *pluginObj = new RegionRT(classes,coords,num);
|
||||
int c = *(static_cast<const int*>(fields[3].data));
|
||||
int h = *(static_cast<const int*>(fields[4].data));
|
||||
int w = *(static_cast<const int*>(fields[5].data));
|
||||
auto *pluginObj = new RegionRT(classes,coords,num,c,h,w);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *RegionRTPluginCreator::getPluginName() const NOEXCEPT {
|
||||
return "RegionRT_tkDNN";
|
||||
return REGIONRT_PLUGIN_NAME;
|
||||
}
|
||||
|
||||
const char *RegionRTPluginCreator::getPluginVersion() const NOEXCEPT {
|
||||
return "1";
|
||||
return REGIONRT_PLUGIN_VERSION;
|
||||
}
|
||||
|
||||
const PluginFieldCollection *RegionRTPluginCreator::getFieldNames() NOEXCEPT {
|
||||
@@ -195,4 +233,3 @@ const PluginFieldCollection *RegionRTPluginCreator::getFieldNames() NOEXCEPT {
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
+56
-18
@@ -4,8 +4,14 @@ using namespace nvinfer1;
|
||||
std::vector<PluginField> ReorgRTPluginCreator::mPluginAttributes;
|
||||
PluginFieldCollection ReorgRTPluginCreator::mFC{};
|
||||
|
||||
ReorgRT::ReorgRT(int stride) {
|
||||
static const char* REORGRT_PLUGIN_VERSION{"1"};
|
||||
static const char* REORGRT_PLUGIN_NAME{"ReorgRT_tkDNN"};
|
||||
|
||||
ReorgRT::ReorgRT(int stride,int c,int h,int w) {
|
||||
this->stride = stride;
|
||||
this->c = c;
|
||||
this->h = h;
|
||||
this->w = w;
|
||||
}
|
||||
|
||||
ReorgRT::~ReorgRT() {}
|
||||
@@ -27,11 +33,6 @@ Dims ReorgRT::getOutputDimensions(int index, const Dims *inputs, int nbInputDims
|
||||
return Dims3{inputs[0].d[0]*stride*stride, inputs[0].d[1]/stride, inputs[0].d[2]/stride};
|
||||
}
|
||||
|
||||
void ReorgRT::configureWithFormat(const Dims *inputDims, int nbInputs, const Dims *outputDims, int nbOutputs, DataType type, PluginFormat format, int maxBatchSize) NOEXCEPT {
|
||||
c = inputDims[0].d[0];
|
||||
h = inputDims[0].d[1];
|
||||
w = inputDims[0].d[2];
|
||||
}
|
||||
|
||||
int ReorgRT::initialize() NOEXCEPT {
|
||||
return 0;
|
||||
@@ -50,7 +51,7 @@ int ReorgRT::enqueue(int batchSize, const void *const *inputs, void *const *outp
|
||||
batchSize, c, h, w, stride, stream);
|
||||
return 0;
|
||||
}
|
||||
#elif NV_TENSORRT_MAJOR == 7
|
||||
#elif NV_TENSORRT_MAJOR <= 7
|
||||
int32_t ReorgRT::enqueue(int32_t batchSize, const void *const *inputs, void **outputs, void *workspace, cudaStream_t stream) {
|
||||
reorgForward((dnnType*)reinterpret_cast<const dnnType*>(inputs[0]),
|
||||
reinterpret_cast<dnnType*>(outputs[0]),
|
||||
@@ -78,11 +79,11 @@ bool ReorgRT::supportsFormat(DataType type, PluginFormat format) const NOEXCEPT
|
||||
}
|
||||
|
||||
const char *ReorgRT::getPluginType() const NOEXCEPT {
|
||||
return "ReorgRT_tkDNN";
|
||||
return REORGRT_PLUGIN_NAME;
|
||||
}
|
||||
|
||||
const char *ReorgRT::getPluginVersion() const NOEXCEPT {
|
||||
return "1";
|
||||
return REORGRT_PLUGIN_VERSION;
|
||||
}
|
||||
|
||||
void ReorgRT::destroy() NOEXCEPT {
|
||||
@@ -97,14 +98,45 @@ void ReorgRT::setPluginNamespace(const char *pluginNamespace) NOEXCEPT {
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
IPluginV2 *ReorgRT::clone() const NOEXCEPT {
|
||||
auto *p = new ReorgRT(stride);
|
||||
IPluginV2Ext *ReorgRT::clone() const NOEXCEPT {
|
||||
auto *p = new ReorgRT(stride,c,h,w);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
DataType ReorgRT::getOutputDataType(int index, const nvinfer1::DataType *inputTypes, int nbInputs) const NOEXCEPT {
|
||||
return DataType::kFLOAT;
|
||||
}
|
||||
|
||||
void ReorgRT::attachToContext(cudnnContext *cudnnContext, cublasContext *cublasContext,
|
||||
IGpuAllocator *gpuAllocator) NOEXCEPT {
|
||||
|
||||
}
|
||||
|
||||
bool ReorgRT::isOutputBroadcastAcrossBatch(int outputIndex, const bool *inputIsBroadcasted, int nbInputs) const NOEXCEPT {
|
||||
return false;
|
||||
}
|
||||
|
||||
bool ReorgRT::canBroadcastInputAcrossBatch(int inputIndex) const NOEXCEPT {
|
||||
return false;
|
||||
}
|
||||
|
||||
void ReorgRT::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 ReorgRT::detachFromContext() NOEXCEPT {
|
||||
|
||||
}
|
||||
|
||||
ReorgRTPluginCreator::ReorgRTPluginCreator() {
|
||||
mPluginAttributes.clear();
|
||||
mPluginAttributes.emplace_back(PluginField("stride", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("c", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("h", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("w", nullptr,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
@@ -117,28 +149,34 @@ const char *ReorgRTPluginCreator::getPluginNamespace() const NOEXCEPT {
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *ReorgRTPluginCreator::deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT {
|
||||
IPluginV2Ext *ReorgRTPluginCreator::deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT {
|
||||
auto *pluginObj = new ReorgRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *ReorgRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
IPluginV2Ext *ReorgRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
const PluginField *fields = fc->fields;
|
||||
assert(fc->nbFields == 1);
|
||||
assert(fields[0].type == PluginFieldType::kINT32);
|
||||
assert(fc->nbFields == 4);
|
||||
for(int i=0;i<4;i++){
|
||||
assert(fields[1].type == PluginFieldType::kINT32);
|
||||
}
|
||||
int stride = *(static_cast<const int *>(fields[0].data));
|
||||
auto *pluginObj = new ReorgRT(stride);
|
||||
int c = *(static_cast<const int *>(fields[1].data));
|
||||
int h = *(static_cast<const int *>(fields[2].data));
|
||||
int w = *(static_cast<const int *>(fields[3].data));
|
||||
|
||||
auto *pluginObj = new ReorgRT(stride,c,h,w);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *ReorgRTPluginCreator::getPluginName() const NOEXCEPT {
|
||||
return "ReorgRT_tkDNN";
|
||||
return REORGRT_PLUGIN_NAME;
|
||||
}
|
||||
|
||||
const char *ReorgRTPluginCreator::getPluginVersion() const NOEXCEPT {
|
||||
return "1";
|
||||
return REORGRT_PLUGIN_VERSION;
|
||||
}
|
||||
|
||||
const PluginFieldCollection *ReorgRTPluginCreator::getFieldNames() NOEXCEPT {
|
||||
|
||||
+64
-24
@@ -4,20 +4,22 @@ using namespace nvinfer1;
|
||||
std::vector<PluginField> ReshapeRTPluginCreator::mPluginAttributes;
|
||||
PluginFieldCollection ReshapeRTPluginCreator::mFC{};
|
||||
|
||||
ReshapeRT::ReshapeRT(dataDim_t newDim) {
|
||||
new_dim = newDim;
|
||||
n = new_dim.n;
|
||||
c = new_dim.c;
|
||||
h = new_dim.h;
|
||||
w = new_dim.w;
|
||||
static const char* RESHAPERT_PLUGIN_VERSION{"1"};
|
||||
static const char* RESHAPERT_PLUGIN_NAME{"ReshapeRT_tkDNN"};
|
||||
|
||||
ReshapeRT::ReshapeRT(int n,int c,int h,int w) {
|
||||
this->n = n;
|
||||
this->c = c;
|
||||
this->h = h;
|
||||
this->w = w;
|
||||
}
|
||||
|
||||
ReshapeRT::ReshapeRT(const void *data, size_t length) {
|
||||
const char *buf = reinterpret_cast<const char*>(data),*bufCheck = buf;
|
||||
new_dim.n = readBUF<int>(buf);
|
||||
new_dim.c = readBUF<int>(buf);
|
||||
new_dim.h = readBUF<int>(buf);
|
||||
new_dim.w = readBUF<int>(buf);
|
||||
n = readBUF<int>(buf);
|
||||
c = readBUF<int>(buf);
|
||||
h = readBUF<int>(buf);
|
||||
w = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + length);
|
||||
}
|
||||
|
||||
@@ -31,8 +33,6 @@ Dims ReshapeRT::getOutputDimensions(int index, const Dims *inputs, int nbInputDi
|
||||
return Dims3{ c,h,w} ;
|
||||
}
|
||||
|
||||
void ReshapeRT::configureWithFormat(const Dims *inputDims, int nbInputs, const Dims *outputDims, int nbOutputs,DataType type, PluginFormat format, int maxBatchSize) NOEXCEPT {}
|
||||
|
||||
int ReshapeRT::initialize() NOEXCEPT {
|
||||
return 0;
|
||||
}
|
||||
@@ -48,11 +48,10 @@ int ReshapeRT::enqueue(int batchSize, const void *const *inputs, void *const *ou
|
||||
cudaStream_t stream) NOEXCEPT {
|
||||
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
|
||||
dnnType *dstData = reinterpret_cast<dnnType*>(outputs[0]);
|
||||
|
||||
checkCuda( cudaMemcpyAsync(dstData, srcData, batchSize*c*h*w*sizeof(dnnType), cudaMemcpyDeviceToDevice, stream));
|
||||
return 0;
|
||||
}
|
||||
#elif NV_TENSORRT_MAJOR == 7
|
||||
#elif NV_TENSORRT_MAJOR <= 7
|
||||
int32_t ReshapeRT::enqueue(int32_t batchSize, const void *const *inputs, void **outputs, void *workspace, cudaStream_t stream) {
|
||||
std::cout << new_dim.c << ":" << new_dim.h << std::endl;
|
||||
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
|
||||
@@ -83,11 +82,11 @@ bool ReshapeRT::supportsFormat(DataType type, PluginFormat format) const NOEXCEP
|
||||
}
|
||||
|
||||
const char *ReshapeRT::getPluginType() const NOEXCEPT {
|
||||
return "ReshapeRT_tkDNN";
|
||||
return RESHAPERT_PLUGIN_NAME;
|
||||
}
|
||||
|
||||
const char *ReshapeRT::getPluginVersion() const NOEXCEPT {
|
||||
return "1";
|
||||
return RESHAPERT_PLUGIN_VERSION;
|
||||
}
|
||||
|
||||
void ReshapeRT::destroy() NOEXCEPT {
|
||||
@@ -102,14 +101,47 @@ void ReshapeRT::setPluginNamespace(const char *pluginNamespace) NOEXCEPT {
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
IPluginV2 *ReshapeRT::clone() const NOEXCEPT {
|
||||
auto *p = new ReshapeRT(new_dim);
|
||||
IPluginV2Ext *ReshapeRT::clone() const NOEXCEPT {
|
||||
auto *p = new ReshapeRT(n,c,h,w);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
DataType ReshapeRT::getOutputDataType(int index, const nvinfer1::DataType *inputTypes, int nbInputs) const NOEXCEPT {
|
||||
return DataType::kFLOAT;
|
||||
}
|
||||
|
||||
void ReshapeRT::attachToContext(cudnnContext *cudnnContext, cublasContext *cublasContext,
|
||||
IGpuAllocator *gpuAllocator) NOEXCEPT {
|
||||
|
||||
}
|
||||
|
||||
bool
|
||||
ReshapeRT::isOutputBroadcastAcrossBatch(int outputIndex, const bool *inputIsBroadcasted, int nbInputs) const NOEXCEPT {
|
||||
return false;
|
||||
}
|
||||
|
||||
bool ReshapeRT::canBroadcastInputAcrossBatch(int inputIndex) const NOEXCEPT {
|
||||
return false;
|
||||
}
|
||||
|
||||
void ReshapeRT::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 ReshapeRT::detachFromContext() NOEXCEPT {
|
||||
|
||||
}
|
||||
|
||||
ReshapeRTPluginCreator::ReshapeRTPluginCreator() {
|
||||
mPluginAttributes.clear();
|
||||
mPluginAttributes.emplace_back(PluginField("n", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("c", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("h", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("w", nullptr,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
@@ -122,26 +154,34 @@ const char *ReshapeRTPluginCreator::getPluginNamespace() const NOEXCEPT {
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *ReshapeRTPluginCreator::deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT {
|
||||
IPluginV2Ext *ReshapeRTPluginCreator::deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT {
|
||||
auto *pluginObj = new ReshapeRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *ReshapeRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
IPluginV2Ext *ReshapeRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
const PluginField *fields = fc->fields;
|
||||
dataDim_t newDim = *(static_cast<const dataDim_t *>(fields[0].data));
|
||||
ReshapeRT *pluginObj = new ReshapeRT(newDim);
|
||||
assert(fc->nbFields == 4);
|
||||
for(int i=0;i<4;i++){
|
||||
assert(fields[1].type == PluginFieldType::kINT32);
|
||||
}
|
||||
int n = *(static_cast<const int *>(fields[0].data));
|
||||
int c = *(static_cast<const int *>(fields[1].data));
|
||||
int h = *(static_cast<const int *>(fields[2].data));
|
||||
int w = *(static_cast<const int *>(fields[3].data));
|
||||
|
||||
auto *pluginObj = new ReshapeRT(n,c,h,w);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *ReshapeRTPluginCreator::getPluginName() const NOEXCEPT {
|
||||
return "ReshapeRT_tkDNN";
|
||||
return RESHAPERT_PLUGIN_NAME;
|
||||
}
|
||||
|
||||
const char *ReshapeRTPluginCreator::getPluginVersion() const NOEXCEPT {
|
||||
return "1";
|
||||
return RESHAPERT_PLUGIN_VERSION;
|
||||
}
|
||||
|
||||
const PluginFieldCollection *ReshapeRTPluginCreator::getFieldNames() NOEXCEPT {
|
||||
|
||||
@@ -5,10 +5,13 @@ std::vector<PluginField> ResizeLayerRTPluginCreator::mPluginAttributes;
|
||||
PluginFieldCollection ResizeLayerRTPluginCreator::mFC{};
|
||||
|
||||
|
||||
ResizeLayerRT::ResizeLayerRT(int c, int h, int w) {
|
||||
o_c = c;
|
||||
o_h = h;
|
||||
o_w = w;
|
||||
ResizeLayerRT::ResizeLayerRT(int oc, int oh, int ow,int ic,int ih,int iw) {
|
||||
this->o_c = oc;
|
||||
this->o_h = oh;
|
||||
this->o_w = ow;
|
||||
this->i_c = ic;
|
||||
this->i_h = ih;
|
||||
this->i_w = iw;
|
||||
}
|
||||
|
||||
ResizeLayerRT::ResizeLayerRT(const void *data, size_t length) {
|
||||
@@ -32,13 +35,6 @@ Dims ResizeLayerRT::getOutputDimensions(int index, const Dims *inputs, int nbInp
|
||||
return Dims3{o_c, o_h, o_w};
|
||||
}
|
||||
|
||||
void ResizeLayerRT::configureWithFormat(const Dims *inputDims, int nbInputs, const Dims *outputDims, int nbOutputs,
|
||||
DataType type, PluginFormat format, int maxBatchSize) NOEXCEPT {
|
||||
i_c = inputDims[0].d[0];
|
||||
i_h = inputDims[0].d[1];
|
||||
i_w = inputDims[0].d[2];
|
||||
}
|
||||
|
||||
int ResizeLayerRT::initialize() NOEXCEPT {
|
||||
return 0;
|
||||
}
|
||||
@@ -55,7 +51,7 @@ int ResizeLayerRT::enqueue(int batchSize, const void *const *inputs, void *const
|
||||
batchSize, i_c, i_h, i_w, o_c, o_h, o_w, stream);
|
||||
return 0;
|
||||
}
|
||||
#elif NV_TENSORRT_MAJOR == 7
|
||||
#elif NV_TENSORRT_MAJOR <= 7
|
||||
int32_t ResizeLayerRT::enqueue(int32_t batchSize, const void *const *inputs, void **outputs, void *workspace,
|
||||
cudaStream_t stream) {
|
||||
resizeForward((dnnType*)reinterpret_cast<const dnnType*>(inputs[0]),
|
||||
@@ -105,12 +101,42 @@ void ResizeLayerRT::setPluginNamespace(const char *pluginNamespace) NOEXCEPT {
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
IPluginV2 *ResizeLayerRT::clone() const NOEXCEPT {
|
||||
auto *p = new ResizeLayerRT(o_c,o_h,o_w);
|
||||
IPluginV2Ext *ResizeLayerRT::clone() const NOEXCEPT {
|
||||
auto *p = new ResizeLayerRT(o_c,o_h,o_w,i_c,i_h,i_w);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
DataType
|
||||
ResizeLayerRT::getOutputDataType(int index, const nvinfer1::DataType *inputTypes, int nbInputs) const NOEXCEPT {
|
||||
return DataType::kFLOAT;
|
||||
}
|
||||
|
||||
void ResizeLayerRT::attachToContext(cudnnContext *cudnnContext, cublasContext *cublasContext,
|
||||
IGpuAllocator *gpuAllocator) NOEXCEPT {
|
||||
|
||||
}
|
||||
|
||||
bool ResizeLayerRT::isOutputBroadcastAcrossBatch(int outputIndex, const bool *inputIsBroadcasted,
|
||||
int nbInputs) const NOEXCEPT {
|
||||
return false;
|
||||
}
|
||||
|
||||
bool ResizeLayerRT::canBroadcastInputAcrossBatch(int inputIndex) const NOEXCEPT {
|
||||
return false;
|
||||
}
|
||||
|
||||
void ResizeLayerRT::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 ResizeLayerRT::detachFromContext() NOEXCEPT {
|
||||
|
||||
}
|
||||
|
||||
ResizeLayerRTPluginCreator::ResizeLayerRTPluginCreator() {
|
||||
mPluginAttributes.clear();
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
@@ -125,22 +151,25 @@ const char *ResizeLayerRTPluginCreator::getPluginNamespace() const NOEXCEPT {
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *ResizeLayerRTPluginCreator::deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT {
|
||||
IPluginV2Ext *ResizeLayerRTPluginCreator::deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT {
|
||||
auto *pluginObj = new ResizeLayerRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *ResizeLayerRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
IPluginV2Ext *ResizeLayerRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
const PluginField *fields = fc->fields;
|
||||
assert(fc->nbFields == 3);
|
||||
assert(fields[0].type == PluginFieldType::kINT32);
|
||||
assert(fields[1].type == PluginFieldType::kINT32);
|
||||
assert(fields[2].type == PluginFieldType::kINT32);
|
||||
assert(fc->nbFields == 6);
|
||||
for(int i=0;i<6;i++){
|
||||
assert(fields[i].type == PluginFieldType::kINT32);
|
||||
}
|
||||
int oc = *(static_cast<const int *>(fields[0].data));
|
||||
int oh = *(static_cast<const int *>(fields[1].data));
|
||||
int ow = *(static_cast<const int *>(fields[2].data));
|
||||
auto *pluginObj = new ResizeLayerRT(oc,oh,ow);
|
||||
int ic = *(static_cast<const int *>(fields[3].data));
|
||||
int ih = *(static_cast<const int *>(fields[4].data));
|
||||
int iw = *(static_cast<const int *>(fields[5].data));
|
||||
auto *pluginObj = new ResizeLayerRT(oc,oh,ow,ic,ih,iw);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
@@ -4,22 +4,26 @@ using namespace nvinfer1;
|
||||
std::vector<PluginField> ShortcutRTPluginCreator::mPluginAttributes;
|
||||
PluginFieldCollection ShortcutRTPluginCreator::mFC{};
|
||||
|
||||
ShortcutRT::ShortcutRT(tk::dnn::dataDim_t bdim, bool mul) {
|
||||
bDim = bdim;
|
||||
this->bc = bDim.c;
|
||||
this->bh = bDim.h;
|
||||
this->bw = bDim.w;
|
||||
static const char* SHORTCUTRT_PLUGIN_VERSION{"1"};
|
||||
static const char* SHORTCUTRT_PLUGIN_NAME{"ShortcutRT_tkDNN"};
|
||||
|
||||
ShortcutRT::ShortcutRT(int bc,int bh,int bw,int c,int h,int w,bool mul) {
|
||||
this->bc = bc;
|
||||
this->bh = bh;
|
||||
this->bw = bw;
|
||||
this->mul = mul;
|
||||
this->c = c;
|
||||
this->h = h;
|
||||
this->w = w;
|
||||
}
|
||||
|
||||
ShortcutRT::~ShortcutRT() {}
|
||||
|
||||
ShortcutRT::ShortcutRT(const void *data, size_t length) {
|
||||
const char* buf =reinterpret_cast<const char*>(data),*bufCheck = buf;
|
||||
bDim.c = readBUF<int>(buf);
|
||||
bDim.h = readBUF<int>(buf);
|
||||
bDim.w = readBUF<int>(buf);
|
||||
bDim.l = 1;
|
||||
bc = readBUF<int>(buf);
|
||||
bh = readBUF<int>(buf);
|
||||
bw = readBUF<int>(buf);
|
||||
mul = readBUF<bool>(buf);
|
||||
c = readBUF<int>(buf);
|
||||
h = readBUF<int>(buf);
|
||||
@@ -35,13 +39,6 @@ Dims ShortcutRT::getOutputDimensions(int index, const Dims *inputs, int nbInputD
|
||||
return Dims3{inputs[0].d[0], inputs[0].d[1], inputs[0].d[2]};
|
||||
}
|
||||
|
||||
void ShortcutRT::configureWithFormat(const Dims *inputDims, int nbInputs, const Dims *outputDims, int nbOutputs,
|
||||
DataType type, PluginFormat format, int maxBatchSize) NOEXCEPT {
|
||||
c = inputDims[0].d[0];
|
||||
h = inputDims[0].d[1];
|
||||
w = inputDims[0].d[2];
|
||||
}
|
||||
|
||||
int ShortcutRT::initialize() NOEXCEPT {
|
||||
return 0;
|
||||
}
|
||||
@@ -62,7 +59,7 @@ int ShortcutRT::enqueue(int batchSize, const void *const *inputs, void *const *o
|
||||
|
||||
return 0;
|
||||
}
|
||||
#elif NV_TENSORRT_MAJOR == 7
|
||||
#elif NV_TENSORRT_MAJOR <= 7
|
||||
int32_t ShortcutRT::enqueue(int32_t batchSize, const void *const *inputs, void **outputs, void *workspace,
|
||||
cudaStream_t stream) {
|
||||
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
|
||||
@@ -98,11 +95,11 @@ bool ShortcutRT::supportsFormat(DataType type, PluginFormat format) const NOEXCE
|
||||
}
|
||||
|
||||
const char *ShortcutRT::getPluginType() const NOEXCEPT {
|
||||
return "ShortcutRT_tkDNN";
|
||||
return SHORTCUTRT_PLUGIN_NAME;
|
||||
}
|
||||
|
||||
const char *ShortcutRT::getPluginVersion() const NOEXCEPT {
|
||||
return "1";
|
||||
return SHORTCUTRT_PLUGIN_VERSION;
|
||||
}
|
||||
|
||||
void ShortcutRT::destroy() NOEXCEPT {
|
||||
@@ -117,14 +114,50 @@ void ShortcutRT::setPluginNamespace(const char *pluginNamespace) NOEXCEPT {
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
IPluginV2 *ShortcutRT::clone() const NOEXCEPT {
|
||||
auto *p = new ShortcutRT(bDim,mul);
|
||||
IPluginV2Ext *ShortcutRT::clone() const NOEXCEPT {
|
||||
auto *p = new ShortcutRT(bc,bh,bw,c,h,w,mul);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
void ShortcutRT::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 {
|
||||
|
||||
}
|
||||
|
||||
bool ShortcutRT::isOutputBroadcastAcrossBatch(int32_t outputIndex, const bool *inputIsBroadcasted,
|
||||
int32_t nbInputs) const NOEXCEPT {
|
||||
return false;
|
||||
}
|
||||
|
||||
bool ShortcutRT::canBroadcastInputAcrossBatch(int32_t inputIndex) const NOEXCEPT {
|
||||
return false;
|
||||
}
|
||||
|
||||
void ShortcutRT::attachToContext(cudnnContext *, cublasContext *, IGpuAllocator *) NOEXCEPT {
|
||||
|
||||
}
|
||||
|
||||
void ShortcutRT::detachFromContext() NOEXCEPT {
|
||||
|
||||
}
|
||||
|
||||
DataType ShortcutRT::getOutputDataType(int32_t index, const nvinfer1::DataType *inputTypes, int32_t nbInputs) const NOEXCEPT {
|
||||
return DataType::kFLOAT;
|
||||
}
|
||||
|
||||
|
||||
ShortcutRTPluginCreator::ShortcutRTPluginCreator() {
|
||||
mPluginAttributes.clear();
|
||||
mPluginAttributes.emplace_back(PluginField("bc", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("bh", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("bw", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("mul", nullptr,PluginFieldType::kUNKNOWN,1));
|
||||
mPluginAttributes.emplace_back(PluginField("c", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("h", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("w", nullptr,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
@@ -137,28 +170,33 @@ const char *ShortcutRTPluginCreator::getPluginNamespace() const NOEXCEPT {
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *ShortcutRTPluginCreator::deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT {
|
||||
IPluginV2Ext *ShortcutRTPluginCreator::deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT {
|
||||
auto *pluginObj = new ShortcutRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *ShortcutRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
IPluginV2Ext *ShortcutRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
const PluginField *fields = fc->fields;
|
||||
//todo assert
|
||||
tk::dnn::dataDim_t bdim = *(static_cast<const tk::dnn::dataDim_t *>(fields[0].data));
|
||||
bool mul = *(static_cast<const bool *>(fields[1].data));
|
||||
auto *pluginObj = new ShortcutRT(bdim,mul);
|
||||
assert(fc->nbFields == 7);
|
||||
int bc = *(static_cast<const int *>(fields[0].data));
|
||||
int bh = *(static_cast<const int *>(fields[1].data));
|
||||
int bw = *(static_cast<const int *>(fields[2].data));
|
||||
bool mul = *(static_cast<const bool *>(fields[3].data));
|
||||
int c = *(static_cast<const int *>(fields[4].data));
|
||||
int h = *(static_cast<const int *>(fields[5].data));
|
||||
int w = *(static_cast<const int *>(fields[6].data));
|
||||
auto *pluginObj = new ShortcutRT(bc,bh,bw,c,h,w,mul);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *ShortcutRTPluginCreator::getPluginName() const NOEXCEPT {
|
||||
return "ShortcutRT_tkDNN";
|
||||
return SHORTCUTRT_PLUGIN_NAME;
|
||||
}
|
||||
|
||||
const char *ShortcutRTPluginCreator::getPluginVersion() const NOEXCEPT {
|
||||
return "1";
|
||||
return SHORTCUTRT_PLUGIN_VERSION;
|
||||
}
|
||||
|
||||
const PluginFieldCollection *ShortcutRTPluginCreator::getFieldNames() NOEXCEPT {
|
||||
|
||||
@@ -4,8 +4,14 @@ using namespace nvinfer1;
|
||||
std::vector<PluginField> UpsampleRTPluginCreator::mPluginAttributes;
|
||||
PluginFieldCollection UpsampleRTPluginCreator::mFC{};
|
||||
|
||||
UpsampleRT::UpsampleRT(int stride) {
|
||||
static const char* UPSAMPLERT_PLUGIN_VERSION{"1"};
|
||||
static const char* UPSAMPLERT_PLUGIN_NAME{"UpSample_tkDNN"};
|
||||
|
||||
UpsampleRT::UpsampleRT(int stride,int c,int h,int w) {
|
||||
this->stride = stride;
|
||||
this->h = h;
|
||||
this->c = c;
|
||||
this->w = w;
|
||||
}
|
||||
|
||||
UpsampleRT::UpsampleRT(const void *data, size_t length) {
|
||||
@@ -27,12 +33,7 @@ Dims UpsampleRT::getOutputDimensions(int index, const Dims *inputs, int nbInputD
|
||||
return Dims3(inputs[0].d[0], inputs[0].d[1]*stride, inputs[0].d[2]*stride);
|
||||
}
|
||||
|
||||
void UpsampleRT::configureWithFormat(const Dims *inputDims, int nbInputs, const Dims *outputDims, int nbOutputs,
|
||||
DataType type, PluginFormat format, int maxBatchSize) NOEXCEPT {
|
||||
c = inputDims[0].d[0];
|
||||
h = inputDims[0].d[1];
|
||||
w = inputDims[0].d[2];
|
||||
}
|
||||
|
||||
|
||||
int UpsampleRT::initialize() NOEXCEPT {
|
||||
return 0;
|
||||
@@ -47,14 +48,14 @@ size_t UpsampleRT::getWorkspaceSize(int maxBatchSize) const NOEXCEPT {
|
||||
#if NV_TENSORRT_MAJOR > 7
|
||||
int UpsampleRT::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]);
|
||||
dnnType *dstData = reinterpret_cast<dnnType*>(outputs[0]);
|
||||
auto *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
|
||||
auto *dstData = reinterpret_cast<dnnType*>(outputs[0]);
|
||||
|
||||
fill(dstData, batchSize*c*h*w*stride*stride, 0.0, stream);
|
||||
upsampleForward(srcData, dstData, batchSize, c, h, w, stride, 1, 1, stream);
|
||||
return 0;
|
||||
}
|
||||
#elif NV_TENSORRT_MAJOR == 7
|
||||
#elif NV_TENSORRT_MAJOR <= 7
|
||||
int32_t UpsampleRT::enqueue(int32_t batchSize, const void *const *inputs, void **outputs, void *workspace,
|
||||
cudaStream_t stream) {
|
||||
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
|
||||
@@ -84,11 +85,11 @@ bool UpsampleRT::supportsFormat(DataType type, PluginFormat format) const NOEXCE
|
||||
}
|
||||
|
||||
const char *UpsampleRT::getPluginType() const NOEXCEPT {
|
||||
return "Upsample_tkDNN";
|
||||
return UPSAMPLERT_PLUGIN_NAME;
|
||||
}
|
||||
|
||||
const char *UpsampleRT::getPluginVersion() const NOEXCEPT {
|
||||
return "1";
|
||||
return UPSAMPLERT_PLUGIN_VERSION;
|
||||
}
|
||||
|
||||
void UpsampleRT::destroy() NOEXCEPT {
|
||||
@@ -103,14 +104,45 @@ void UpsampleRT::setPluginNamespace(const char *pluginNamespace) NOEXCEPT {
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
IPluginV2 *UpsampleRT::clone() const NOEXCEPT {
|
||||
auto *p = new UpsampleRT(stride);
|
||||
IPluginV2Ext *UpsampleRT::clone() const NOEXCEPT {
|
||||
auto *p = new UpsampleRT(stride,c,h,w);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
bool UpsampleRT::isOutputBroadcastAcrossBatch(int32_t outputIndex, const bool *inputIsBroadcasted,
|
||||
int32_t nbInputs) const NOEXCEPT {
|
||||
return false;
|
||||
}
|
||||
|
||||
bool UpsampleRT::canBroadcastInputAcrossBatch(int32_t inputIndex) const NOEXCEPT {
|
||||
return false;
|
||||
}
|
||||
|
||||
void UpsampleRT::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 UpsampleRT::attachToContext(cudnnContext *, cublasContext *, IGpuAllocator *) NOEXCEPT {
|
||||
}
|
||||
|
||||
void UpsampleRT::detachFromContext() NOEXCEPT {
|
||||
|
||||
}
|
||||
|
||||
DataType UpsampleRT::getOutputDataType(int32_t index, const nvinfer1::DataType *inputTypes, int32_t nbInputs) const NOEXCEPT {
|
||||
return DataType::kFLOAT;
|
||||
}
|
||||
|
||||
UpsampleRTPluginCreator::UpsampleRTPluginCreator() {
|
||||
mPluginAttributes.clear();
|
||||
mPluginAttributes.emplace_back(PluginField("stride", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("c", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("h", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("w", nullptr,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
@@ -123,26 +155,29 @@ const char *UpsampleRTPluginCreator::getPluginNamespace() const NOEXCEPT {
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *UpsampleRTPluginCreator::deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT {
|
||||
IPluginV2Ext *UpsampleRTPluginCreator::deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT {
|
||||
auto *pluginObj = new UpsampleRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *UpsampleRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
IPluginV2Ext *UpsampleRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
const PluginField *fields = fc->fields;
|
||||
int stride = *(static_cast<const int *>(fields[0].data));
|
||||
auto *pluginObj = new UpsampleRT(stride);
|
||||
int c = *(static_cast<const int*>(fields[1].data));
|
||||
int h = *(static_cast<const int*>(fields[2].data));
|
||||
int w = *(static_cast<const int*>(fields[3].data));
|
||||
auto *pluginObj = new UpsampleRT(stride,c,h,w);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *UpsampleRTPluginCreator::getPluginName() const NOEXCEPT {
|
||||
return "Upsample_tkDNN";
|
||||
return UPSAMPLERT_PLUGIN_NAME;
|
||||
}
|
||||
|
||||
const char *UpsampleRTPluginCreator::getPluginVersion() const NOEXCEPT {
|
||||
return "1";
|
||||
return UPSAMPLERT_PLUGIN_VERSION;
|
||||
}
|
||||
|
||||
const PluginFieldCollection *UpsampleRTPluginCreator::getFieldNames() NOEXCEPT {
|
||||
|
||||
+79
-38
@@ -1,12 +1,21 @@
|
||||
#include <tkDNN/pluginsRT/YoloRT.h>
|
||||
|
||||
#include <utility>
|
||||
using namespace nvinfer1;
|
||||
|
||||
std::vector<PluginField> YoloRTPluginCreator::mPluginAttributes;
|
||||
PluginFieldCollection YoloRTPluginCreator::mFC{};
|
||||
|
||||
YoloRT::YoloRT(int classes, int num, tk::dnn::Yolo *Yolo, int n_masks, float scale_xy, float nms_thresh, int nms_kind,
|
||||
static const char* YOLORT_PLUGIN_VERSION{"1"};
|
||||
static const char* YOLORT_PLUGIN_NAME{"YoloRT_tkDNN"};
|
||||
|
||||
YoloRT::YoloRT(int classes, int num, int c,int h,int w,std::vector<std::string> classNames,
|
||||
std::vector<float> masks_v,std::vector<float> bias_v,int n_masks, float scale_xy,
|
||||
float nms_thresh, int nms_kind,
|
||||
int new_coords) {
|
||||
this->yolo = Yolo;
|
||||
this->c = c;
|
||||
this->h = h;
|
||||
this->w = w;
|
||||
this->classes = classes;
|
||||
this->num = num;
|
||||
this->n_masks = n_masks;
|
||||
@@ -14,14 +23,10 @@ YoloRT::YoloRT(int classes, int num, tk::dnn::Yolo *Yolo, int n_masks, float sca
|
||||
this->nms_thresh = nms_thresh;
|
||||
this->nms_kind = nms_kind;
|
||||
this->new_coords = new_coords;
|
||||
this->classesNames = std::move(classNames);
|
||||
this->mask = std::move(masks_v);
|
||||
this->bias = std::move(bias_v);
|
||||
|
||||
mask = new dnnType[n_masks];
|
||||
bias = new dnnType[num * n_masks * 2];
|
||||
if (yolo != nullptr) {
|
||||
memcpy(mask, yolo->mask_h, sizeof(dnnType) * n_masks);
|
||||
memcpy(bias, yolo->bias_h, sizeof(dnnType) * num * n_masks * 2);
|
||||
classesNames = yolo->classesNames;
|
||||
}
|
||||
}
|
||||
|
||||
YoloRT::YoloRT(const void *data, size_t length) {
|
||||
@@ -38,16 +43,14 @@ YoloRT::YoloRT(const void *data, size_t length) {
|
||||
c = readBUF<int>(buf);
|
||||
h = readBUF<int>(buf);
|
||||
w = readBUF<int>(buf);
|
||||
mask.resize(n_masks);
|
||||
for(int i=0;i<n_masks;i++){
|
||||
maskTemp.push_back(readBUF<dnnType>(buf));
|
||||
std::cout<<maskTemp[i]<<std::endl;
|
||||
mask[i] = readBUF<dnnType>(buf);
|
||||
}
|
||||
bias.resize(n_masks*2*num);
|
||||
for(int i=0;i<n_masks*2*num;i++){
|
||||
biasTemp.push_back(readBUF<dnnType>(buf));
|
||||
std::cout<<biasTemp[i]<<std::endl;
|
||||
bias[i] = readBUF<dnnType>(buf);
|
||||
}
|
||||
mask = maskTemp.data();
|
||||
bias = biasTemp.data();
|
||||
classesNames.resize(classes);
|
||||
for(int i=0;i<classes;i++){
|
||||
char tmp[YOLORT_CLASSNAME_W];
|
||||
@@ -68,12 +71,7 @@ Dims YoloRT::getOutputDimensions(int index, const Dims *inputs, int nbInputDims)
|
||||
return inputs[0];
|
||||
}
|
||||
|
||||
void YoloRT::configureWithFormat(const Dims *inputDims, int nbInputs, const Dims *outputDims, int nbOutputs, DataType type,
|
||||
PluginFormat format, int maxBatchSize) NOEXCEPT {
|
||||
c = inputDims[0].d[0];
|
||||
h = inputDims[0].d[1];
|
||||
w = inputDims[0].d[2];
|
||||
}
|
||||
|
||||
|
||||
int YoloRT::initialize() NOEXCEPT {
|
||||
return 0;
|
||||
@@ -190,11 +188,11 @@ void YoloRT::serialize(void *buffer) const NOEXCEPT {
|
||||
}
|
||||
|
||||
const char *YoloRT::getPluginType() const NOEXCEPT {
|
||||
return "YoloRT_tkDNN";
|
||||
return YOLORT_PLUGIN_NAME;
|
||||
}
|
||||
|
||||
const char *YoloRT::getPluginVersion() const NOEXCEPT {
|
||||
return "1";
|
||||
return YOLORT_PLUGIN_VERSION;
|
||||
}
|
||||
|
||||
void YoloRT::destroy() NOEXCEPT {
|
||||
@@ -209,14 +207,54 @@ void YoloRT::setPluginNamespace(const char *pluginNamespace) NOEXCEPT {
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
IPluginV2 *YoloRT::clone() const NOEXCEPT {
|
||||
auto *p = new YoloRT(classes, num,yolo, n_masks, scaleXY, nms_thresh, nms_kind, new_coords);
|
||||
IPluginV2Ext *YoloRT::clone() const NOEXCEPT {
|
||||
auto *p = new YoloRT(classes, num,c,h,w,classesNames,mask,bias, n_masks, scaleXY, nms_thresh, nms_kind, new_coords);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
DataType YoloRT::getOutputDataType(int index, const nvinfer1::DataType *inputTypes, int nbInputs) const NOEXCEPT {
|
||||
return DataType::kFLOAT;
|
||||
}
|
||||
|
||||
void YoloRT::attachToContext(cudnnContext *cudnnContext, cublasContext *cublasContext,
|
||||
IGpuAllocator *gpuAllocator) NOEXCEPT {
|
||||
|
||||
}
|
||||
|
||||
void YoloRT::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 {
|
||||
|
||||
}
|
||||
|
||||
bool YoloRT::isOutputBroadcastAcrossBatch(int outputIndex, const bool *inputIsBroadcasted, int nbInputs) const NOEXCEPT {
|
||||
return false;
|
||||
}
|
||||
|
||||
bool YoloRT::canBroadcastInputAcrossBatch(int inputIndex) const NOEXCEPT {
|
||||
return false;
|
||||
}
|
||||
|
||||
void YoloRT::detachFromContext() NOEXCEPT {
|
||||
|
||||
}
|
||||
|
||||
YoloRTPluginCreator::YoloRTPluginCreator() {
|
||||
mPluginAttributes.clear();
|
||||
mPluginAttributes.emplace_back(PluginField("classes", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("num", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("c", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("h", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("w", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("classNames", nullptr,PluginFieldType::kUNKNOWN,1));
|
||||
mPluginAttributes.emplace_back(PluginField("mask_v", nullptr,PluginFieldType::kFLOAT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("bias_v", nullptr,PluginFieldType::kFLOAT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("n_masks", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("scaleXy", nullptr,PluginFieldType::kFLOAT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("nms_thresh", nullptr,PluginFieldType::kFLOAT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("nms_kind", nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("new_coords", nullptr,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
@@ -229,34 +267,37 @@ const char *YoloRTPluginCreator::getPluginNamespace() const NOEXCEPT {
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *YoloRTPluginCreator::deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT {
|
||||
IPluginV2Ext *YoloRTPluginCreator::deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT {
|
||||
auto *pluginObj = new YoloRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *YoloRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
IPluginV2Ext *YoloRTPluginCreator::createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT {
|
||||
const PluginField *fields = fc->fields;
|
||||
//todo assert
|
||||
int classes = *(static_cast<const int *>(fields[0].data));
|
||||
int num = *(static_cast<const int *>(fields[1].data));
|
||||
Yolo *yoloTemp = const_cast<Yolo *>(static_cast<const Yolo *>(fields[2].data));
|
||||
int numMasks = *(static_cast<const int*>(fields[3].data));
|
||||
float scaleXY = *(static_cast<const float *>(fields[4].data));
|
||||
float nmsThresh = *(static_cast<const float *>(fields[5].data));
|
||||
int nmsKind = *(static_cast<const int *>(fields[6].data));
|
||||
int newCoords = *(static_cast<const int *>(fields[7].data));
|
||||
YoloRT *pluginObj = new YoloRT(classes,num,yoloTemp,numMasks,scaleXY,nmsThresh,nmsKind,newCoords);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
int c = *(static_cast<const int *>(fields[2].data));
|
||||
int h = *(static_cast<const int *>(fields[3].data));
|
||||
int w = *(static_cast<const int *>(fields[4].data));
|
||||
std::vector<std::string> classNames(static_cast<const std::string *>(fields[5].data),static_cast<const std::string *>(fields[5].data) + fields[5].length);
|
||||
std::vector<dnnType> mask_v(static_cast<const dnnType*>(fields[6].data),static_cast<const dnnType*>(fields[6].data) + fields[6].length);
|
||||
std::vector<dnnType> bias_v(static_cast<const dnnType*>(fields[7].data),static_cast<const dnnType*>(fields[7].data) + fields[7].length);
|
||||
int n_masks = *(static_cast<const int *>(fields[8].data));
|
||||
dnnType scaleXY = *(static_cast<const float*>(fields[9].data));
|
||||
dnnType nmsThresh = *(static_cast<const float*>(fields[10].data));
|
||||
int nms_kind = *(static_cast<const int*>(fields[11].data));
|
||||
int new_coords = *(static_cast<const int*>(fields[12].data));
|
||||
auto *pluginObj = new YoloRT(classes,num,c,h,w,classNames,mask_v,bias_v,n_masks,scaleXY,nmsThresh,nms_kind,new_coords);
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *YoloRTPluginCreator::getPluginName() const NOEXCEPT {
|
||||
return "YoloRT_tkDNN";
|
||||
return YOLORT_PLUGIN_NAME;
|
||||
}
|
||||
|
||||
const char *YoloRTPluginCreator::getPluginVersion() const NOEXCEPT {
|
||||
return "1";
|
||||
return YOLORT_PLUGIN_VERSION;
|
||||
}
|
||||
|
||||
const PluginFieldCollection *YoloRTPluginCreator::getFieldNames() NOEXCEPT {
|
||||
|
||||
Reference in New Issue
Block a user