Updates to build libkernel.so under TensorRT 8
This commit is contained in:
@@ -39,7 +39,7 @@ public:
|
||||
float *getLabels() { return mLabels.data(); }
|
||||
int getBatchesRead() const { return mBatchCount; }
|
||||
int getBatchSize() const { return mBatchSize; }
|
||||
nvinfer1::DimsNCHW getDims() const { return mDims; }
|
||||
nvinfer1::Dims4 getDims() const { return mDims; }
|
||||
float* getFileBatch() { return &mFileBatch[0]; }
|
||||
float* getFileLabels() { return &mFileLabels[0]; }
|
||||
void readInListFile(const std::string& dataFilePath, std::vector<std::string>& mListIn);
|
||||
@@ -55,7 +55,7 @@ private:
|
||||
int mFileBatchPos{ 0 };
|
||||
int mImageSize{ 0 };
|
||||
|
||||
nvinfer1::DimsNCHW mDims;
|
||||
nvinfer1::Dims4 mDims;
|
||||
std::vector<float> mBatch;
|
||||
std::vector<float> mLabels;
|
||||
std::vector<float> mFileBatch;
|
||||
|
||||
@@ -30,10 +30,10 @@ public:
|
||||
Int8EntropyCalibrator(BatchStream& stream, int firstBatch, const std::string& calibTableFilePath,
|
||||
const std::string& inputBlobName, bool readCache = true);
|
||||
virtual ~Int8EntropyCalibrator() { checkCuda(cudaFree(mDeviceInput)); }
|
||||
int getBatchSize() const override { return mStream.getBatchSize(); }
|
||||
bool getBatch(void* bindings[], const char* names[], int nbBindings) override;
|
||||
const void* readCalibrationCache(size_t& length) override;
|
||||
void writeCalibrationCache(const void* cache, size_t length) override;
|
||||
int getBatchSize() const NOEXCEPT override { return mStream.getBatchSize(); }
|
||||
bool getBatch(void* bindings[], const char* names[], int nbBindings) NOEXCEPT override;
|
||||
const void* readCalibrationCache(size_t& length) NOEXCEPT override;
|
||||
void writeCalibrationCache(const void* cache, size_t length) NOEXCEPT override;
|
||||
|
||||
private:
|
||||
BatchStream mStream;
|
||||
|
||||
@@ -40,14 +40,16 @@ using namespace nvinfer1;
|
||||
#include "pluginsRT/ReshapeRT.h"
|
||||
#include "pluginsRT/MaxPoolingFixedSizeRT.h"
|
||||
|
||||
class PluginFactory : IPluginFactory
|
||||
/*
|
||||
class PluginFactory : IPlugin
|
||||
{
|
||||
public:
|
||||
YoloRT *yolos[16];
|
||||
int n_yolos;
|
||||
|
||||
virtual IPlugin* createPlugin(const char* layerName, const void* serialData, size_t serialLength);
|
||||
};
|
||||
};*/
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -74,7 +76,6 @@ public:
|
||||
dnnType *output;
|
||||
cudaStream_t stream;
|
||||
|
||||
PluginFactory *pluginFactory;
|
||||
|
||||
NetworkRT(Network *net, const char *name);
|
||||
virtual ~NetworkRT();
|
||||
|
||||
@@ -1,61 +1,147 @@
|
||||
#include<cassert>
|
||||
#include "../kernels.h"
|
||||
#include <cassert>
|
||||
|
||||
class ActivationLeakyRT : public IPlugin {
|
||||
class ActivationLeakyRT : public IPluginV2 {
|
||||
|
||||
public:
|
||||
ActivationLeakyRT(float s) {
|
||||
slope = s;
|
||||
}
|
||||
ActivationLeakyRT(float s) { slope = s; }
|
||||
|
||||
~ActivationLeakyRT(){
|
||||
ActivationLeakyRT(const void *data, size_t length)
|
||||
{
|
||||
const char* buf = reinterpret_cast<const char*>(data),*bufCheck = buf;
|
||||
slope = readBUF<float>(buf);
|
||||
size = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + length);
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
int getNbOutputs() const override {
|
||||
return 1;
|
||||
}
|
||||
~ActivationLeakyRT() {}
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override {
|
||||
return inputs[0];
|
||||
}
|
||||
int getNbOutputs() const NOEXCEPT override { return 1; }
|
||||
|
||||
void configure(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, int maxBatchSize) override {
|
||||
size = 1;
|
||||
for(int i=0; i<outputDims[0].nbDims; i++)
|
||||
size *= outputDims[0].d[i];
|
||||
}
|
||||
Dims getOutputDimensions(int index, const Dims *inputs, int nbInputDims) NOEXCEPT override {
|
||||
return inputs[0];
|
||||
}
|
||||
|
||||
int initialize() override {
|
||||
void configureWithFormat(const Dims *inputDims, int nbInputs, const Dims *outputDims, int nbOutputs,DataType type,PluginFormat format, int maxBatchSize) NOEXCEPT override
|
||||
{
|
||||
assert(type == DataType::kFLOAT && format == PluginFormat::kLINEAR);
|
||||
size = 1;
|
||||
for (int i = 0; i < outputDims[0].nbDims; i++)
|
||||
size *= outputDims[0].d[i];
|
||||
}
|
||||
|
||||
return 0;
|
||||
}
|
||||
int initialize() NOEXCEPT override { return 0; }
|
||||
|
||||
virtual void terminate() override {
|
||||
}
|
||||
virtual void terminate() NOEXCEPT override {}
|
||||
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const override {
|
||||
return 0;
|
||||
}
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const NOEXCEPT override {
|
||||
return 0;
|
||||
}
|
||||
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override {
|
||||
virtual int enqueue(int batchSize, void const *const *inputs, void *const *outputs, void *workspace,
|
||||
cudaStream_t stream) NOEXCEPT override {
|
||||
activationLEAKYForward(
|
||||
(dnnType *) reinterpret_cast<const dnnType *>(inputs[0]),
|
||||
reinterpret_cast<dnnType *>(outputs[0]), batchSize * size, slope,
|
||||
stream);
|
||||
return 0;
|
||||
}
|
||||
|
||||
activationLEAKYForward((dnnType*)reinterpret_cast<const dnnType*>(inputs[0]),
|
||||
reinterpret_cast<dnnType*>(outputs[0]), batchSize*size, slope, stream);
|
||||
return 0;
|
||||
}
|
||||
virtual size_t getSerializationSize() const NOEXCEPT override {
|
||||
return 1 * sizeof(int) + 1 * sizeof(float);
|
||||
}
|
||||
|
||||
virtual void serialize(void *buffer) const NOEXCEPT override {
|
||||
char *buf = reinterpret_cast<char *>(buffer), *a = buf;
|
||||
tk::dnn::writeBUF(buf, size);
|
||||
assert(buf == a + getSerializationSize());
|
||||
}
|
||||
|
||||
virtual size_t getSerializationSize() override {
|
||||
return 1*sizeof(int) + 1*sizeof(float);
|
||||
}
|
||||
bool supportsFormat(DataType type, PluginFormat format) const NOEXCEPT override {
|
||||
return (type == DataType::kFLOAT && format == PluginFormat::kLINEAR);
|
||||
}
|
||||
|
||||
virtual void serialize(void* buffer) override {
|
||||
char *buf = reinterpret_cast<char*>(buffer),*a=buf;
|
||||
tk::dnn::writeBUF(buf, size);
|
||||
assert(buf == a + getSerializationSize());
|
||||
}
|
||||
const char *getPluginType() const NOEXCEPT override {
|
||||
return "ActivationLeakyRT_tkDNN";
|
||||
}
|
||||
|
||||
int size;
|
||||
float slope;
|
||||
const char *getPluginVersion() const NOEXCEPT override {
|
||||
return "1";
|
||||
}
|
||||
|
||||
void destroy() NOEXCEPT override { delete this; }
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override {
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override {
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
IPluginV2* clone() const NOEXCEPT override {
|
||||
ActivationLeakyRT *p = new ActivationLeakyRT(slope);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
int size;
|
||||
float slope;
|
||||
|
||||
private:
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
class ActivationLeakyRTPluginCreator : public IPluginCreator {
|
||||
public:
|
||||
ActivationLeakyRTPluginCreator() {
|
||||
mPluginAttributes.emplace_back(
|
||||
PluginField("slope", nullptr, PluginFieldType::kFLOAT32, 1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
IPluginV2 *deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT override {
|
||||
ActivationLeakyRT *pluginObj = new ActivationLeakyRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override {
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT override {
|
||||
const PluginField *fields = fc->fields;
|
||||
assert(fc->nbFields == 1);
|
||||
assert(fields[0].type == PluginFieldType::kFLOAT32);
|
||||
float slope = *(static_cast<const float *>(fields[0].data));
|
||||
ActivationLeakyRT *pluginObj = new ActivationLeakyRT(slope);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *getPluginName() const NOEXCEPT override{
|
||||
return "ActivationLeakyRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const PluginFieldCollection *getFieldNames() NOEXCEPT override{
|
||||
return &mFC;
|
||||
}
|
||||
|
||||
private:
|
||||
static PluginFieldCollection mFC;
|
||||
static std::vector<PluginField> mPluginAttributes;
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
REGISTER_TENSORRT_PLUGIN(ActivationLeakyRTPluginCreator);
|
||||
@@ -1,11 +1,18 @@
|
||||
#include<cassert>
|
||||
#include "../kernels.h"
|
||||
|
||||
class ActivationLogisticRT : public IPlugin {
|
||||
class ActivationLogisticRT : public IPluginV2 {
|
||||
|
||||
public:
|
||||
ActivationLogisticRT() {
|
||||
|
||||
}
|
||||
|
||||
ActivationLogisticRT(const void *data, size_t length)
|
||||
{
|
||||
const char* buf = reinterpret_cast<const char*>(data),*bufCheck = buf;
|
||||
size = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + length);
|
||||
|
||||
}
|
||||
|
||||
@@ -13,33 +20,33 @@ public:
|
||||
|
||||
}
|
||||
|
||||
int getNbOutputs() const override {
|
||||
int getNbOutputs() const NOEXCEPT override {
|
||||
return 1;
|
||||
}
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override {
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) NOEXCEPT override {
|
||||
return inputs[0];
|
||||
}
|
||||
|
||||
void configure(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, int maxBatchSize) override {
|
||||
void configureWithFormat(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs,DataType type,PluginFormat format, int maxBatchSize) NOEXCEPT override {
|
||||
size = 1;
|
||||
for(int i=0; i<outputDims[0].nbDims; i++)
|
||||
size *= outputDims[0].d[i];
|
||||
}
|
||||
|
||||
int initialize() override {
|
||||
int initialize() NOEXCEPT override {
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
virtual void terminate() override {
|
||||
virtual void terminate() NOEXCEPT override {
|
||||
}
|
||||
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const override {
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const NOEXCEPT override {
|
||||
return 0;
|
||||
}
|
||||
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override {
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void* const* outputs, void* workspace, cudaStream_t stream) NOEXCEPT override {
|
||||
|
||||
activationLOGISTICForward((dnnType*)reinterpret_cast<const dnnType*>(inputs[0]),
|
||||
reinterpret_cast<dnnType*>(outputs[0]), batchSize*size, stream);
|
||||
@@ -47,14 +54,94 @@ public:
|
||||
}
|
||||
|
||||
|
||||
virtual size_t getSerializationSize() override {
|
||||
virtual size_t getSerializationSize() const NOEXCEPT override {
|
||||
return 1*sizeof(int);
|
||||
}
|
||||
|
||||
virtual void serialize(void* buffer) override {
|
||||
virtual void serialize(void* buffer) const NOEXCEPT override {
|
||||
char *buf = reinterpret_cast<char*>(buffer);
|
||||
tk::dnn::writeBUF(buf, size);
|
||||
}
|
||||
|
||||
const char *getPluginType() const NOEXCEPT override {
|
||||
return "ActivationLogisticRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override {
|
||||
return "1";
|
||||
}
|
||||
|
||||
void destroy() NOEXCEPT override { delete this; }
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override {
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override {
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
bool supportsFormat(DataType type,PluginFormat format) const NOEXCEPT override{
|
||||
return true;
|
||||
//todo assert;
|
||||
}
|
||||
|
||||
IPluginV2 *clone() const NOEXCEPT override{
|
||||
ActivationLogisticRT *p = new ActivationLogisticRT();
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
int size;
|
||||
|
||||
private:
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
class ActivationLogisticRTPluginCreator : public IPluginCreator{
|
||||
public:
|
||||
ActivationLogisticRTPluginCreator(){
|
||||
mPluginAttributes.clear();
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
IPluginV2 *deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT override {
|
||||
ActivationLogisticRT *pluginObj = new ActivationLogisticRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override {
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *createPlugin(const char *name, const PluginFieldCollection *fc) NOEXCEPT override {
|
||||
ActivationLogisticRT *pluginObj = new ActivationLogisticRT();
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const PluginFieldCollection *getFieldNames() NOEXCEPT override{
|
||||
return &mFC;
|
||||
}
|
||||
|
||||
const char *getPluginName() const NOEXCEPT override{
|
||||
return "ActivationLogisticRT_tkDNN";
|
||||
}
|
||||
|
||||
private:
|
||||
static PluginFieldCollection mFC;
|
||||
static std::vector<PluginField> mPluginAttributes;
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
REGISTER_TENSORRT_PLUGIN(ActivationLogisticRTPluginCreator);
|
||||
@@ -1,61 +1,134 @@
|
||||
#include<cassert>
|
||||
#include "../kernels.h"
|
||||
|
||||
class ActivationMishRT : public IPlugin {
|
||||
class ActivationMishRT : public IPluginV2 {
|
||||
|
||||
public:
|
||||
ActivationMishRT() {
|
||||
ActivationMishRT() {}
|
||||
|
||||
~ActivationMishRT() {}
|
||||
|
||||
ActivationMishRT(const void *data, size_t length) {
|
||||
const char *buf = reinterpret_cast<const char *>(data), *bufCheck = buf;
|
||||
size = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + length);
|
||||
}
|
||||
|
||||
|
||||
}
|
||||
int getNbOutputs() const NOEXCEPT override { return 1; }
|
||||
|
||||
~ActivationMishRT(){
|
||||
Dims getOutputDimensions(int index, const Dims *inputs, int nbInputDims) NOEXCEPT override { return inputs[0]; }
|
||||
|
||||
}
|
||||
void configureWithFormat(const Dims *inputDims, int nbInputs, const Dims *outputDims, int nbOutputs, DataType type,
|
||||
PluginFormat format, int maxBatchSize) NOEXCEPT override {
|
||||
assert(format == PluginFormat::kLINEAR);
|
||||
size = 1;
|
||||
for (int i = 0; i < outputDims[0].nbDims; i++)
|
||||
size *= outputDims[0].d[i];
|
||||
}
|
||||
|
||||
int getNbOutputs() const override {
|
||||
return 1;
|
||||
}
|
||||
int initialize() NOEXCEPT override { return 0; }
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override {
|
||||
return inputs[0];
|
||||
}
|
||||
virtual void terminate() NOEXCEPT override {}
|
||||
|
||||
void configure(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, int maxBatchSize) override {
|
||||
size = 1;
|
||||
for(int i=0; i<outputDims[0].nbDims; i++)
|
||||
size *= outputDims[0].d[i];
|
||||
}
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const NOEXCEPT override { return 0; }
|
||||
|
||||
int initialize() override {
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
virtual void terminate() override {
|
||||
}
|
||||
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const override {
|
||||
return 0;
|
||||
}
|
||||
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override {
|
||||
|
||||
activationMishForward((dnnType*)reinterpret_cast<const dnnType*>(inputs[0]),
|
||||
reinterpret_cast<dnnType*>(outputs[0]), batchSize*size, stream);
|
||||
return 0;
|
||||
}
|
||||
virtual int enqueue(int batchSize, const void *const *inputs, void *const *outputs, void *workspace,
|
||||
cudaStream_t stream) NOEXCEPT override {
|
||||
activationMishForward((dnnType *) reinterpret_cast<const dnnType *>(inputs[0]),
|
||||
reinterpret_cast<dnnType *>(outputs[0]), batchSize * size, stream);
|
||||
return 0;
|
||||
}
|
||||
|
||||
|
||||
virtual size_t getSerializationSize() override {
|
||||
return 1*sizeof(int);
|
||||
}
|
||||
virtual size_t getSerializationSize() const NOEXCEPT override {
|
||||
return 1 * sizeof(int);
|
||||
}
|
||||
|
||||
virtual void serialize(void* buffer) override {
|
||||
char *buf = reinterpret_cast<char*>(buffer),*a=buf;
|
||||
tk::dnn::writeBUF(buf, size);
|
||||
assert(buf == a + getSerializationSize());
|
||||
}
|
||||
virtual void serialize(void *buffer) const NOEXCEPT override {
|
||||
char *buf = reinterpret_cast<char *>(buffer), *a = buf;
|
||||
tk::dnn::writeBUF(buf, size);
|
||||
assert(buf == a + getSerializationSize());
|
||||
}
|
||||
|
||||
int size;
|
||||
const char *getPluginType() const NOEXCEPT override {
|
||||
return "ActivationMishRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override {
|
||||
return "1";
|
||||
}
|
||||
|
||||
void destroy() NOEXCEPT override { delete this; }
|
||||
|
||||
bool supportsFormat(DataType type,PluginFormat format) const NOEXCEPT override{
|
||||
return true;
|
||||
}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override {
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *plguinNamespace) NOEXCEPT override {
|
||||
mPluginNamespace = plguinNamespace;
|
||||
}
|
||||
|
||||
IPluginV2 *clone() const NOEXCEPT override {
|
||||
ActivationMishRT *p = new ActivationMishRT();
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
int size;
|
||||
private:
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
class ActivationMishRTPluginCreator : public IPluginCreator {
|
||||
public:
|
||||
ActivationMishRTPluginCreator() {
|
||||
mPluginAttributes.clear();
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override {
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
const char* getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *deserializePlugin(const char* name,const void* serialData,size_t serialLength) NOEXCEPT override{
|
||||
ActivationMishRT *pluginObj = new ActivationMishRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *createPlugin(const char* name,const PluginFieldCollection *fc) NOEXCEPT override{
|
||||
const PluginField *fields = fc->fields;
|
||||
ActivationMishRT *pluginObj = new ActivationMishRT();
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *getPluginName() const NOEXCEPT override{
|
||||
return "ActivationMishRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const PluginFieldCollection *getFieldNames() NOEXCEPT override{
|
||||
return &mFC;
|
||||
}
|
||||
|
||||
private:
|
||||
static PluginFieldCollection mFC;
|
||||
static std::vector<PluginField> mPluginAttributes;
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
REGISTER_TENSORRT_PLUGIN(ActivationMishRTPluginCreator);
|
||||
@@ -1,63 +1,149 @@
|
||||
#include<cassert>
|
||||
#include "../kernels.h"
|
||||
|
||||
class ActivationReLUCeiling : public IPlugin {
|
||||
|
||||
class ActivationReLUCeiling : public IPluginV2 {
|
||||
|
||||
public:
|
||||
ActivationReLUCeiling(const float ceiling) {
|
||||
this->ceiling = ceiling;
|
||||
}
|
||||
ActivationReLUCeiling(const float ceiling) {
|
||||
this->ceiling = ceiling;
|
||||
}
|
||||
|
||||
~ActivationReLUCeiling(){
|
||||
~ActivationReLUCeiling() {
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
int getNbOutputs() const override {
|
||||
return 1;
|
||||
}
|
||||
ActivationReLUCeiling(const void *data, size_t length) {
|
||||
const char *buf = reinterpret_cast<const char *>(data), *bufCheck = buf;
|
||||
ceiling = readBUF<float>(buf);
|
||||
size = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + length);
|
||||
}
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override {
|
||||
return inputs[0];
|
||||
}
|
||||
int getNbOutputs() const NOEXCEPT override {
|
||||
return 1;
|
||||
}
|
||||
|
||||
void configure(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, int maxBatchSize) override {
|
||||
size = 1;
|
||||
for(int i=0; i<outputDims[0].nbDims; i++)
|
||||
size *= outputDims[0].d[i];
|
||||
}
|
||||
Dims getOutputDimensions(int index, const Dims *inputs, int nbInputDims) NOEXCEPT override {
|
||||
return inputs[0];
|
||||
}
|
||||
|
||||
int initialize() override {
|
||||
void configureWithFormat(const Dims *inputDims, int nbInputs, const Dims *outputDims, int nbOutputs, DataType type,
|
||||
PluginFormat format, int maxBatchSize) NOEXCEPT override {
|
||||
assert(type == DataType::kFLOAT && format == PluginFormat::kLINEAR);
|
||||
size = 1;
|
||||
for (int i = 0; i < outputDims[0].nbDims; i++)
|
||||
size *= outputDims[0].d[i];
|
||||
}
|
||||
|
||||
return 0;
|
||||
}
|
||||
int initialize() NOEXCEPT override { return 0; }
|
||||
|
||||
virtual void terminate() override {
|
||||
}
|
||||
virtual void terminate() NOEXCEPT override {}
|
||||
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const override {
|
||||
return 0;
|
||||
}
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const NOEXCEPT override {
|
||||
return 0;
|
||||
}
|
||||
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override {
|
||||
|
||||
activationReLUCeilingForward((dnnType*)reinterpret_cast<const dnnType*>(inputs[0]),
|
||||
reinterpret_cast<dnnType*>(outputs[0]), batchSize*size, ceiling, stream);
|
||||
return 0;
|
||||
}
|
||||
virtual int enqueue(int batchSize, const void *const *inputs, void *const *outputs, void *workspace,
|
||||
cudaStream_t stream) NOEXCEPT override {
|
||||
activationReLUCeilingForward((dnnType *) reinterpret_cast<const dnnType *>(inputs[0]),
|
||||
reinterpret_cast<dnnType *>(outputs[0]), batchSize * size, ceiling, stream);
|
||||
return 0;
|
||||
}
|
||||
|
||||
|
||||
virtual size_t getSerializationSize() override {
|
||||
return 1*sizeof(int) + 1*sizeof(float);
|
||||
}
|
||||
virtual size_t getSerializationSize() const NOEXCEPT override {
|
||||
return 1 * sizeof(int) + 1 * sizeof(float);
|
||||
}
|
||||
|
||||
virtual void serialize(void* buffer) override {
|
||||
char *buf = reinterpret_cast<char*>(buffer),*a=buf;
|
||||
tk::dnn::writeBUF(buf, ceiling);
|
||||
tk::dnn::writeBUF(buf, size);
|
||||
assert(buf = a + getSerializationSize());
|
||||
|
||||
}
|
||||
virtual void serialize(void *buffer) const NOEXCEPT override {
|
||||
char *buf = reinterpret_cast<char *>(buffer), *a = buf;
|
||||
tk::dnn::writeBUF(buf, ceiling);
|
||||
tk::dnn::writeBUF(buf, size);
|
||||
assert(buf = a + getSerializationSize());
|
||||
|
||||
int size;
|
||||
float ceiling;
|
||||
}
|
||||
|
||||
IPluginV2 *clone() const NOEXCEPT override {
|
||||
ActivationReLUCeiling *p = new ActivationReLUCeiling(ceiling);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
bool supportsFormat(DataType type, PluginFormat format) const NOEXCEPT override {
|
||||
return (type == DataType::kFLOAT && format == PluginFormat::kLINEAR);
|
||||
}
|
||||
|
||||
void destroy() NOEXCEPT override { delete this; };
|
||||
|
||||
const char *getPluginType() const NOEXCEPT override {
|
||||
return "ActivationReLUCeilingRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override {
|
||||
return "1";
|
||||
}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override {
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override {
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
int size;
|
||||
float ceiling;
|
||||
private:
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
class ActivationReLUCeilingPluginCreator : public IPluginCreator {
|
||||
public:
|
||||
ActivationReLUCeilingPluginCreator() {
|
||||
mPluginAttributes.emplace_back(PluginField("ceiling", nullptr, PluginFieldType::kFLOAT32, 1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override {
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override {
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *deserializePlugin(const char *name, const void *serialData, size_t serialLength) NOEXCEPT override {
|
||||
ActivationReLUCeiling *pluginObj = new ActivationReLUCeiling(serialData, serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *createPlugin(const char *name,const PluginFieldCollection *fc) NOEXCEPT override{
|
||||
const PluginField *fields = fc->fields;
|
||||
float ceiling = *(static_cast<const float *>(fields[0].data));
|
||||
ActivationReLUCeiling *pluginObj = new ActivationReLUCeiling(ceiling);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *getPluginName() const NOEXCEPT override{
|
||||
return "ActivationReLUCeilingRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const PluginFieldCollection *getFieldNames() NOEXCEPT override{
|
||||
return &mFC;
|
||||
}
|
||||
|
||||
private:
|
||||
static PluginFieldCollection mFC;
|
||||
static std::vector<PluginField> mPluginAttributes;
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
REGISTER_TENSORRT_PLUGIN(ActivationReLUCeilingPluginCreator);
|
||||
|
||||
@@ -1,8 +1,7 @@
|
||||
#include<cassert>
|
||||
#include "../kernels.h"
|
||||
|
||||
|
||||
class DeformableConvRT : public IPlugin {
|
||||
class DeformableConvRT : public IPluginV2 {
|
||||
|
||||
|
||||
|
||||
@@ -27,6 +26,8 @@ public:
|
||||
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;
|
||||
@@ -38,7 +39,6 @@ public:
|
||||
checkCuda( cudaMalloc(&mask, chunk_dim*sizeof(dnnType)));
|
||||
checkCuda( cudaMalloc(&ones_d2, dim_ones*sizeof(dnnType)));
|
||||
if(deformable != nullptr) {
|
||||
this->defRT = deformable;
|
||||
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) );
|
||||
@@ -61,27 +61,78 @@ public:
|
||||
cublasDestroy(handle);
|
||||
}
|
||||
|
||||
int getNbOutputs() const override {
|
||||
return 1;
|
||||
DeformableConvRT(const void *data,size_t length){
|
||||
const char* buf = reinterpret_cast<const char*>(data),*bufCheck = buf;
|
||||
chunk_dim = readBUF<int>(buf);
|
||||
kh = readBUF<int>(buf);
|
||||
kw = readBUF<int>(buf);
|
||||
sh = readBUF<int>(buf);
|
||||
sw = readBUF<int>(buf);
|
||||
ph = readBUF<int>(buf);
|
||||
pw = readBUF<int>(buf);
|
||||
deformableGroup = readBUF<int>(buf);
|
||||
i_n = readBUF<int>(buf);
|
||||
i_c = readBUF<int>(buf);
|
||||
i_h = readBUF<int>(buf);
|
||||
i_w = readBUF<int>(buf);
|
||||
o_n = readBUF<int>(buf);
|
||||
o_c = readBUF<int>(buf);
|
||||
o_h = readBUF<int>(buf);
|
||||
o_w = readBUF<int>(buf);
|
||||
dnnType *aus = new dnnType[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];
|
||||
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];
|
||||
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];
|
||||
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];
|
||||
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];
|
||||
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);
|
||||
|
||||
assert(buf == bufCheck + length);
|
||||
}
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override {
|
||||
return DimsCHW{defRT->output_dim.c, defRT->output_dim.h, defRT->output_dim.w};
|
||||
int getNbOutputs() const NOEXCEPT override {return 1;}
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) NOEXCEPT override {
|
||||
return Dims3{defRT->output_dim.c, defRT->output_dim.h, defRT->output_dim.w};
|
||||
}
|
||||
|
||||
void configure(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, int maxBatchSize) override { }
|
||||
void configureWithFormat(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs,DataType type,PluginFormat format,int maxBatchSize) NOEXCEPT override { }
|
||||
|
||||
int initialize() override {
|
||||
return 0;
|
||||
}
|
||||
int initialize() NOEXCEPT override {return 0;}
|
||||
|
||||
virtual void terminate() override { }
|
||||
virtual void terminate() NOEXCEPT override { }
|
||||
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const override {
|
||||
return 0;
|
||||
}
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const NOEXCEPT override { return 0;}
|
||||
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override {
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void* const* outputs, void* workspace, cudaStream_t stream) NOEXCEPT override {
|
||||
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
|
||||
dnnType *output_conv = (dnnType*)reinterpret_cast<const dnnType*>(inputs[1]);
|
||||
|
||||
@@ -109,13 +160,12 @@ public:
|
||||
return 0;
|
||||
}
|
||||
|
||||
|
||||
virtual size_t getSerializationSize() override {
|
||||
virtual size_t getSerializationSize() const NOEXCEPT override {
|
||||
return 16 * 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);
|
||||
}
|
||||
|
||||
virtual void serialize(void* buffer) override {
|
||||
virtual void serialize(void* buffer) const NOEXCEPT override {
|
||||
char *buf = reinterpret_cast<char*>(buffer),*a=buf;
|
||||
tk::dnn::writeBUF(buf, chunk_dim);
|
||||
tk::dnn::writeBUF(buf, kh);
|
||||
@@ -166,6 +216,35 @@ public:
|
||||
assert(buf == a + getSerializationSize());
|
||||
}
|
||||
|
||||
void destroy() NOEXCEPT override {delete this;}
|
||||
|
||||
bool supportsFormat(DataType type,PluginFormat format) const NOEXCEPT override{
|
||||
return true;
|
||||
//todo assert
|
||||
}
|
||||
const char *getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
const char *getPluginType() const NOEXCEPT override{
|
||||
return "DeformableConvRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
IPluginV2* clone() const NOEXCEPT override{
|
||||
DeformableConvRT *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);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
|
||||
cublasStatus_t stat;
|
||||
cublasHandle_t handle;
|
||||
int i_n, i_c, i_h, i_w;
|
||||
@@ -193,4 +272,90 @@ public:
|
||||
|
||||
|
||||
tk::dnn::DeformConv2d *defRT;
|
||||
|
||||
private:
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
class DeformableConvRTPluginCreator : public IPluginCreator{
|
||||
public:
|
||||
DeformableConvRTPluginCreator(){
|
||||
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("deformableGroup",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("defRT",nullptr,PluginFieldType::kUNKNOWN,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override {
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *deserializePlugin(const char *name,const void *serialData,size_t serialLength) NOEXCEPT override{
|
||||
DeformableConvRT *pluginObj = new DeformableConvRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *createPlugin(const char* name,const PluginFieldCollection *fc) NOEXCEPT override{
|
||||
const PluginField *fields = fc->fields;
|
||||
int chunk_dim = *(static_cast<const int *>(fields[0].data));
|
||||
int kh = *(static_cast<const int *>(fields[1].data));
|
||||
int kw = *(static_cast<const int *>(fields[2].data));
|
||||
int sh = *(static_cast<const int *>(fields[3].data));
|
||||
int sw = *(static_cast<const int *>(fields[4].data));
|
||||
int ph = *(static_cast<const int *>(fields[5].data));
|
||||
int pw = *(static_cast<const int *>(fields[6].data));
|
||||
int deformableGroup = *(static_cast<const int *>(fields[7].data));
|
||||
int i_n = *(static_cast<const int *>(fields[8].data));
|
||||
int i_c = *(static_cast<const int *>(fields[9].data));
|
||||
int i_h = *(static_cast<const int *>(fields[10].data));
|
||||
int i_w = *(static_cast<const int *>(fields[11].data));
|
||||
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));
|
||||
DeformConv2d *defRT = const_cast<DeformConv2d *>(static_cast<const DeformConv2d *>(fields[15].data));
|
||||
DeformableConvRT *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);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *getPluginName() const NOEXCEPT override{
|
||||
return "DeformableConvRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const PluginFieldCollection *getFieldNames() NOEXCEPT override{
|
||||
return &mFC;
|
||||
}
|
||||
|
||||
private:
|
||||
static PluginFieldCollection mFC;
|
||||
static std::vector<PluginField> mPluginAttributes;
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
REGISTER_TENSORRT_PLUGIN(DeformableConvRTPluginCreator);
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
#include<cassert>
|
||||
|
||||
class FlattenConcatRT : public IPlugin {
|
||||
class FlattenConcatRT : public IPluginV2 {
|
||||
|
||||
public:
|
||||
FlattenConcatRT() {
|
||||
@@ -11,19 +11,29 @@ public:
|
||||
}
|
||||
}
|
||||
|
||||
FlattenConcatRT(const void *data,size_t length){
|
||||
const char *buf = reinterpret_cast<const char *>(data),*bufCheck=buf;
|
||||
c = readBUF<int>(buf);
|
||||
h = readBUF<int>(buf);
|
||||
w = readBUF<int>(buf);
|
||||
rows = readBUF<int>(buf);
|
||||
cols = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + length);
|
||||
}
|
||||
|
||||
~FlattenConcatRT(){
|
||||
|
||||
}
|
||||
|
||||
int getNbOutputs() const override {
|
||||
int getNbOutputs() const NOEXCEPT override {
|
||||
return 1;
|
||||
}
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override {
|
||||
return DimsCHW{ inputs[0].d[0] * inputs[0].d[1] * inputs[0].d[2], 1, 1};
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) NOEXCEPT override {
|
||||
return Dims3{ inputs[0].d[0] * inputs[0].d[1] * inputs[0].d[2], 1, 1};
|
||||
}
|
||||
|
||||
void configure(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, int maxBatchSize) override {
|
||||
void configureWithFormat(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs,DataType type,PluginFormat format,int maxBatchSize) NOEXCEPT override {
|
||||
assert(nbOutputs == 1 && nbInputs ==1);
|
||||
rows = inputDims[0].d[0];
|
||||
cols = inputDims[0].d[1] * inputDims[0].d[2];
|
||||
@@ -32,19 +42,13 @@ public:
|
||||
w = 1;
|
||||
}
|
||||
|
||||
int initialize() override {
|
||||
return 0;
|
||||
}
|
||||
int initialize() NOEXCEPT override {return 0;}
|
||||
|
||||
virtual void terminate() override {
|
||||
checkERROR(cublasDestroy(handle));
|
||||
}
|
||||
virtual void terminate() NOEXCEPT override { checkERROR(cublasDestroy(handle));}
|
||||
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const override {
|
||||
return 0;
|
||||
}
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const NOEXCEPT override {return 0;}
|
||||
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override {
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void* const* outputs, void* workspace, cudaStream_t stream) NOEXCEPT override {
|
||||
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
|
||||
dnnType *dstData = reinterpret_cast<dnnType*>(outputs[0]);
|
||||
checkCuda( cudaMemcpyAsync(dstData, srcData, batchSize*rows*cols*sizeof(dnnType), cudaMemcpyDeviceToDevice, stream));
|
||||
@@ -59,12 +63,11 @@ public:
|
||||
return 0;
|
||||
}
|
||||
|
||||
|
||||
virtual size_t getSerializationSize() override {
|
||||
virtual size_t getSerializationSize() const NOEXCEPT override {
|
||||
return 5*sizeof(int);
|
||||
}
|
||||
|
||||
virtual void serialize(void* buffer) override {
|
||||
virtual void serialize(void* buffer) const NOEXCEPT override {
|
||||
char *buf = reinterpret_cast<char*>(buffer),*a = buf;
|
||||
tk::dnn::writeBUF(buf, c);
|
||||
tk::dnn::writeBUF(buf, h);
|
||||
@@ -74,8 +77,86 @@ public:
|
||||
assert(buf == a + getSerializationSize());
|
||||
}
|
||||
|
||||
void destroy() NOEXCEPT override{delete this;}
|
||||
|
||||
bool supportsFormat(DataType type,PluginFormat format) const NOEXCEPT override{
|
||||
return true;
|
||||
}
|
||||
|
||||
const char *getPluginType() const NOEXCEPT override{
|
||||
return "FlattenConcatRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
IPluginV2 *clone() const NOEXCEPT override {
|
||||
FlattenConcatRT *p = new FlattenConcatRT();
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
int c, h, w;
|
||||
int rows, cols;
|
||||
cublasStatus_t stat;
|
||||
cublasHandle_t handle;
|
||||
cublasHandle_t handle;
|
||||
private:
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
class FlattenConcatRTPluginCreator : public IPluginCreator{
|
||||
public:
|
||||
FlattenConcatRTPluginCreator(){
|
||||
mPluginAttributes.clear();
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *deserializePlugin(const char *name,const void *serialData,size_t serialLength) NOEXCEPT override{
|
||||
FlattenConcatRT *pluginObj = new FlattenConcatRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *createPlugin(const char *name,const PluginFieldCollection *fc) NOEXCEPT override{
|
||||
FlattenConcatRT *pluginObj = new FlattenConcatRT();
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *getPluginName() const NOEXCEPT override{
|
||||
return "FlattenConcatRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const PluginFieldCollection *getFieldNames() NOEXCEPT override{
|
||||
return &mFC;
|
||||
}
|
||||
|
||||
private:
|
||||
static PluginFieldCollection mFC;
|
||||
static std::vector<PluginField> mPluginAttributes;
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
REGISTER_TENSORRT_PLUGIN(FlattenConcatRTPluginCreator);
|
||||
@@ -1,7 +1,8 @@
|
||||
#include<cassert>
|
||||
#include "../kernels.h"
|
||||
|
||||
class MaxPoolFixedSizeRT : public IPlugin {
|
||||
|
||||
class MaxPoolFixedSizeRT : public IPluginV2 {
|
||||
|
||||
public:
|
||||
MaxPoolFixedSizeRT(int c, int h, int w, int n, int strideH, int strideW, int winSize, int padding) {
|
||||
@@ -15,32 +16,40 @@ public:
|
||||
this->padding = padding;
|
||||
}
|
||||
|
||||
MaxPoolFixedSizeRT(const void *data,size_t length){
|
||||
const char *buf = reinterpret_cast<const char*>(data),*bufCheck = buf;
|
||||
c = readBUF<int>(buf);
|
||||
h = readBUF<int>(buf);
|
||||
w = readBUF<int>(buf);
|
||||
n = readBUF<int>(buf);
|
||||
stride_H = readBUF<int>(buf);
|
||||
stride_W = readBUF<int>(buf);
|
||||
winSize = readBUF<int>(buf);
|
||||
padding = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + length);
|
||||
}
|
||||
|
||||
~MaxPoolFixedSizeRT(){
|
||||
}
|
||||
|
||||
int getNbOutputs() const override {
|
||||
int getNbOutputs() const NOEXCEPT override {
|
||||
return 1;
|
||||
}
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override {
|
||||
return DimsCHW{this->c, this->h, this->w};
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) NOEXCEPT override {
|
||||
return Dims3{this->c, this->h, this->w};
|
||||
}
|
||||
|
||||
void configure(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, int maxBatchSize) override {
|
||||
void configureWithFormat(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs,DataType type,PluginFormat format,int maxBatchSize) NOEXCEPT override {
|
||||
}
|
||||
|
||||
int initialize() override {
|
||||
return 0;
|
||||
}
|
||||
int initialize() NOEXCEPT override {return 0;}
|
||||
|
||||
virtual void terminate() override {
|
||||
}
|
||||
virtual void terminate() NOEXCEPT override {}
|
||||
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const override {
|
||||
return 0;
|
||||
}
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const NOEXCEPT override { return 0;}
|
||||
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override {
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void* const* outputs, void* workspace, cudaStream_t stream) NOEXCEPT override {
|
||||
|
||||
//std::cout<<this->n<<" "<<this->c<<" "<<this->h<<" "<<this->w<<" "<<this->stride_H<<" "<<this->stride_W<<" "<<this->winSize<<" "<<this->padding<<std::endl;
|
||||
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
|
||||
@@ -50,11 +59,11 @@ public:
|
||||
}
|
||||
|
||||
|
||||
virtual size_t getSerializationSize() override {
|
||||
virtual size_t getSerializationSize() const NOEXCEPT override {
|
||||
return 8*sizeof(int);
|
||||
}
|
||||
|
||||
virtual void serialize(void* buffer) override {
|
||||
virtual void serialize(void* buffer) const NOEXCEPT override {
|
||||
char *buf = reinterpret_cast<char*>(buffer),*a=buf;
|
||||
|
||||
tk::dnn::writeBUF(buf, this->c);
|
||||
@@ -68,8 +77,106 @@ public:
|
||||
assert(buf == a + getSerializationSize());
|
||||
}
|
||||
|
||||
void destroy() NOEXCEPT override{delete this;}
|
||||
|
||||
bool supportsFormat(DataType type,PluginFormat format) const NOEXCEPT override{
|
||||
return true;
|
||||
//todo assert
|
||||
}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
const char *getPluginType() const NOEXCEPT override{
|
||||
return "MaxPoolingFixedSizeRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
IPluginV2 *clone() const NOEXCEPT override{
|
||||
MaxPoolFixedSizeRT *p = new MaxPoolFixedSizeRT(c,h,w,n,stride_H,stride_W,winSize,padding);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
|
||||
int n, c, h, w;
|
||||
int stride_H, stride_W;
|
||||
int winSize;
|
||||
int padding;
|
||||
|
||||
private:
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
class MaxPoolFixedSizeRTPluginCreator : public IPluginCreator{
|
||||
public:
|
||||
MaxPoolFixedSizeRTPluginCreator(){
|
||||
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("n",nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("stride_H",nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("stride_W",nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("winSize",nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("padding",nullptr,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *deserializePlugin(const char *name,const void *serialData,size_t serialLength) NOEXCEPT override{
|
||||
MaxPoolFixedSizeRT *pluginObj = new MaxPoolFixedSizeRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *createPlugin(const char *name,const PluginFieldCollection *fc) NOEXCEPT override{
|
||||
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));
|
||||
int n = *(static_cast<const int *>(fields[3].data));
|
||||
int stride_H = *(static_cast<const int *>(fields[4].data));
|
||||
int stride_W = *(static_cast<const int *>(fields[5].data));
|
||||
int winSize = *(static_cast<const int *>(fields[6].data));
|
||||
int padding = *(static_cast<const int *>(fields[7].data));
|
||||
MaxPoolFixedSizeRT *pluginObj = new MaxPoolFixedSizeRT(c,h,w,n,stride_H,stride_W,winSize,padding);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *getPluginName() const NOEXCEPT override{
|
||||
return "MaxPoolingFixedSizeRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const PluginFieldCollection *getFieldNames() NOEXCEPT override{
|
||||
return &mFC;
|
||||
}
|
||||
|
||||
private:
|
||||
static PluginFieldCollection mFC;
|
||||
static std::vector<PluginField> mPluginAttributes;
|
||||
std::string mPluginNamespace;
|
||||
|
||||
};
|
||||
|
||||
REGISTER_TENSORRT_PLUGIN(MaxPoolFixedSizeRTPluginCreator);
|
||||
|
||||
@@ -1,11 +1,10 @@
|
||||
#include<cassert>
|
||||
#include "../kernels.h"
|
||||
|
||||
class RegionRT : public IPlugin {
|
||||
class RegionRT : public IPluginV2 {
|
||||
|
||||
public:
|
||||
RegionRT(int classes, int coords, int num) {
|
||||
|
||||
this->classes = classes;
|
||||
this->coords = coords;
|
||||
this->num = num;
|
||||
@@ -15,33 +14,39 @@ public:
|
||||
|
||||
}
|
||||
|
||||
int getNbOutputs() const override {
|
||||
RegionRT(const void *data,size_t length){
|
||||
const char *buf = reinterpret_cast<const char*>(data),*bufCheck=buf;
|
||||
classes = readBUF<int>(buf);
|
||||
coords = readBUF<int>(buf);
|
||||
num = readBUF<int>(buf);
|
||||
c = readBUF<int>(buf);
|
||||
h = readBUF<int>(buf);
|
||||
w = readBUF<int>(buf);
|
||||
assert(buf == bufCheck+length);
|
||||
}
|
||||
|
||||
int getNbOutputs() const NOEXCEPT override {
|
||||
return 1;
|
||||
}
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override {
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) NOEXCEPT override {
|
||||
return inputs[0];
|
||||
}
|
||||
|
||||
void configure(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, int maxBatchSize) override {
|
||||
void configureWithFormat(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs,DataType type,PluginFormat format, int maxBatchSize) NOEXCEPT override {
|
||||
c = inputDims[0].d[0];
|
||||
h = inputDims[0].d[1];
|
||||
w = inputDims[0].d[2];
|
||||
}
|
||||
|
||||
int initialize() override {
|
||||
int initialize() NOEXCEPT override { return 0; }
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
virtual void terminate() override {
|
||||
}
|
||||
virtual void terminate() NOEXCEPT override { }
|
||||
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const override {
|
||||
return 0;
|
||||
}
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const NOEXCEPT override { return 0; }
|
||||
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override {
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void* const* outputs, void* workspace, cudaStream_t stream) NOEXCEPT override {
|
||||
|
||||
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
|
||||
dnnType *dstData = reinterpret_cast<dnnType*>(outputs[0]);
|
||||
@@ -68,11 +73,11 @@ public:
|
||||
}
|
||||
|
||||
|
||||
virtual size_t getSerializationSize() override {
|
||||
virtual size_t getSerializationSize() const NOEXCEPT override {
|
||||
return 6*sizeof(int);
|
||||
}
|
||||
|
||||
virtual void serialize(void* buffer) override {
|
||||
virtual void serialize(void* buffer) const NOEXCEPT override {
|
||||
char *buf = reinterpret_cast<char*>(buffer),*a=buf;
|
||||
tk::dnn::writeBUF(buf, classes);
|
||||
tk::dnn::writeBUF(buf, coords);
|
||||
@@ -83,6 +88,34 @@ public:
|
||||
assert(buf == a + getSerializationSize());
|
||||
}
|
||||
|
||||
const char *getPluginType() const NOEXCEPT override{
|
||||
return "RegionRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
void destroy() NOEXCEPT override {delete this;}
|
||||
|
||||
const char* getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char* pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
bool supportsFormat(DataType type,PluginFormat format) const NOEXCEPT override{
|
||||
return true;
|
||||
}
|
||||
|
||||
IPluginV2* clone() const NOEXCEPT override{
|
||||
RegionRT *p = new RegionRT(classes,coords,num);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
int c, h, w;
|
||||
int classes, coords, num;
|
||||
|
||||
@@ -92,4 +125,64 @@ public:
|
||||
return batch*c*h*w + n*w*h*(coords+classes+1) + entry*w*h + loc;
|
||||
}
|
||||
|
||||
private:
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
class RegionRTPluginCreator : public IPluginCreator{
|
||||
public:
|
||||
RegionRTPluginCreator(){
|
||||
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));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char* pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
IPluginV2 *deserializePlugin(const char* name,const void *serialData,size_t serialLength) NOEXCEPT override{
|
||||
RegionRT *pluginObj = new RegionRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *createPlugin(const char* name,const PluginFieldCollection *fc) NOEXCEPT override{
|
||||
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);
|
||||
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);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
|
||||
}
|
||||
|
||||
const char *getPluginName() const NOEXCEPT override{
|
||||
return "RegionRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const PluginFieldCollection *getFieldNames() NOEXCEPT override{
|
||||
return &mFC;
|
||||
}
|
||||
private:
|
||||
static PluginFieldCollection mFC;
|
||||
static std::vector<PluginField> mPluginAttributes;
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
REGISTER_TENSORRT_PLUGIN(RegionRTPluginCreator);
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
#include<cassert>
|
||||
#include "../kernels.h"
|
||||
|
||||
class ReorgRT : public IPlugin {
|
||||
class ReorgRT : public IPluginV2 {
|
||||
|
||||
public:
|
||||
ReorgRT(int stride) {
|
||||
@@ -12,33 +12,34 @@ public:
|
||||
|
||||
}
|
||||
|
||||
int getNbOutputs() const override {
|
||||
return 1;
|
||||
ReorgRT(const void* data,size_t length){
|
||||
const char* buf = reinterpret_cast<const char*>(data),*bufCheck = buf;
|
||||
stride = readBUF<int>(buf);
|
||||
c = readBUF<int>(buf);
|
||||
h = readBUF<int>(buf);
|
||||
w = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + length);
|
||||
}
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override {
|
||||
return DimsCHW{inputs[0].d[0]*stride*stride, inputs[0].d[1]/stride, inputs[0].d[2]/stride};
|
||||
int getNbOutputs() const NOEXCEPT override {return 1;}
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) NOEXCEPT override {
|
||||
return Dims3{inputs[0].d[0]*stride*stride, inputs[0].d[1]/stride, inputs[0].d[2]/stride};
|
||||
}
|
||||
|
||||
void configure(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, int maxBatchSize) override {
|
||||
void configureWithFormat(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs,DataType type,PluginFormat format, int maxBatchSize) NOEXCEPT override {
|
||||
c = inputDims[0].d[0];
|
||||
h = inputDims[0].d[1];
|
||||
w = inputDims[0].d[2];
|
||||
}
|
||||
|
||||
int initialize() override {
|
||||
int initialize() NOEXCEPT override { return 0;}
|
||||
|
||||
return 0;
|
||||
}
|
||||
virtual void terminate() NOEXCEPT override {}
|
||||
|
||||
virtual void terminate() override {
|
||||
}
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const NOEXCEPT override { return 0;}
|
||||
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const override {
|
||||
return 0;
|
||||
}
|
||||
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override {
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void* const* outputs, void* workspace, cudaStream_t stream) NOEXCEPT override {
|
||||
|
||||
reorgForward((dnnType*)reinterpret_cast<const dnnType*>(inputs[0]),
|
||||
reinterpret_cast<dnnType*>(outputs[0]),
|
||||
@@ -47,11 +48,11 @@ public:
|
||||
}
|
||||
|
||||
|
||||
virtual size_t getSerializationSize() override {
|
||||
virtual size_t getSerializationSize() const NOEXCEPT override {
|
||||
return 4*sizeof(int);
|
||||
}
|
||||
|
||||
virtual void serialize(void* buffer) override {
|
||||
virtual void serialize(void* buffer) const NOEXCEPT override {
|
||||
char *buf = reinterpret_cast<char*>(buffer),*a=buf;
|
||||
tk::dnn::writeBUF(buf, stride);
|
||||
tk::dnn::writeBUF(buf, c);
|
||||
@@ -59,6 +60,84 @@ public:
|
||||
tk::dnn::writeBUF(buf, w);
|
||||
assert(buf == a + getSerializationSize());
|
||||
}
|
||||
bool supportsFormat(DataType type,PluginFormat format) const NOEXCEPT override{return true;}
|
||||
|
||||
const char *getPluginType() const NOEXCEPT override{
|
||||
return "ReorgRT_tkDNN";
|
||||
}
|
||||
|
||||
const char* getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
void destroy() NOEXCEPT override{ delete this;}
|
||||
|
||||
const char* getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char* pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
IPluginV2* clone() const NOEXCEPT override{
|
||||
ReorgRT *p = new ReorgRT(stride);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
int c, h, w, stride;
|
||||
private:
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
class ReorgRTPluginCreator : public IPluginCreator{
|
||||
public:
|
||||
ReorgRTPluginCreator(){
|
||||
mPluginAttributes.emplace_back(PluginField("stride",nullptr,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char* pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
const char* getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2* deserializePlugin(const char* name,const void* serialData,size_t serialLength) NOEXCEPT override{
|
||||
ReorgRT *pluginObj = new ReorgRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *createPlugin(const char* name,const PluginFieldCollection* fc) NOEXCEPT override{
|
||||
const PluginField *fields = fc->fields;
|
||||
assert(fc->nbFields == 1);
|
||||
assert(fields[0].type == PluginFieldType::kINT32);
|
||||
int stride = *(static_cast<const int *>(fields[0].data));
|
||||
ReorgRT *pluginObj = new ReorgRT(stride);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *getPluginName() const NOEXCEPT override{
|
||||
return "ReorgRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const PluginFieldCollection *getFieldNames() NOEXCEPT override{
|
||||
return &mFC;
|
||||
}
|
||||
private:
|
||||
static PluginFieldCollection mFC;
|
||||
static std::vector<PluginField> mPluginAttributes;
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
REGISTER_TENSORRT_PLUGIN(ReorgRTPluginCreator);
|
||||
|
||||
|
||||
@@ -1,42 +1,47 @@
|
||||
#include<cassert>
|
||||
|
||||
class ReshapeRT : public IPlugin {
|
||||
class ReshapeRT : public IPluginV2 {
|
||||
|
||||
public:
|
||||
ReshapeRT(dataDim_t new_dim) {
|
||||
ReshapeRT(dataDim_t newDim) {
|
||||
new_dim = newDim;
|
||||
n = new_dim.n;
|
||||
c = new_dim.c;
|
||||
h = new_dim.h;
|
||||
w = new_dim.w;
|
||||
}
|
||||
|
||||
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);
|
||||
assert(buf == bufCheck + length);
|
||||
}
|
||||
|
||||
~ReshapeRT(){
|
||||
|
||||
}
|
||||
|
||||
int getNbOutputs() const override {
|
||||
int getNbOutputs() const NOEXCEPT override {
|
||||
return 1;
|
||||
}
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override {
|
||||
return DimsCHW{ c,h,w};
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) NOEXCEPT override {
|
||||
return Dims3{ c,h,w};
|
||||
}
|
||||
|
||||
void configure(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, int maxBatchSize) override {
|
||||
void configureWithFormat (const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, DataType type,PluginFormat format, int maxBatchSize) NOEXCEPT override {
|
||||
}
|
||||
|
||||
int initialize() override {
|
||||
return 0;
|
||||
}
|
||||
int initialize() NOEXCEPT override {return 0;}
|
||||
|
||||
virtual void terminate() override {
|
||||
}
|
||||
virtual void terminate() NOEXCEPT override {}
|
||||
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const override {
|
||||
return 0;
|
||||
}
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const NOEXCEPT override { return 0;}
|
||||
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override {
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void* const* outputs, void* workspace, cudaStream_t stream) NOEXCEPT override {
|
||||
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
|
||||
dnnType *dstData = reinterpret_cast<dnnType*>(outputs[0]);
|
||||
|
||||
@@ -44,12 +49,11 @@ public:
|
||||
return 0;
|
||||
}
|
||||
|
||||
|
||||
virtual size_t getSerializationSize() override {
|
||||
virtual size_t getSerializationSize() const NOEXCEPT override {
|
||||
return 4*sizeof(int);
|
||||
}
|
||||
|
||||
virtual void serialize(void* buffer) override {
|
||||
virtual void serialize(void* buffer) const NOEXCEPT override {
|
||||
char *buf = reinterpret_cast<char*>(buffer),*a = buf;
|
||||
tk::dnn::writeBUF(buf, n);
|
||||
tk::dnn::writeBUF(buf, c);
|
||||
@@ -58,5 +62,87 @@ public:
|
||||
assert(buf == a + getSerializationSize());
|
||||
}
|
||||
|
||||
bool supportsFormat(DataType type,PluginFormat format) const NOEXCEPT override{
|
||||
return true;
|
||||
//todo assert
|
||||
}
|
||||
|
||||
const char *getPluginType() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "ReshapeRT_tkDNN";
|
||||
}
|
||||
|
||||
void destroy() NOEXCEPT override{delete this;}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
IPluginV2 *clone() const NOEXCEPT override{
|
||||
ReshapeRT *p = new ReshapeRT(new_dim);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
int n, c, h, w;
|
||||
dataDim_t new_dim;
|
||||
private:
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
class ReshapeRTPluginCreator : public IPluginCreator{
|
||||
public:
|
||||
ReshapeRTPluginCreator(){
|
||||
mPluginAttributes.emplace_back(PluginField("new_dim",nullptr,PluginFieldType::kUNKNOWN,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *deserializePlugin(const char* name,const void *serialData,size_t serialLength) NOEXCEPT override{
|
||||
ReshapeRT *pluginObj = new ReshapeRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *createPlugin(const char* name,const PluginFieldCollection *fc) NOEXCEPT override{
|
||||
const PluginField *fields = fc->fields;
|
||||
dataDim_t newDim = *(static_cast<const dataDim_t *>(fields[0].data));
|
||||
ReshapeRT *pluginObj = new ReshapeRT(newDim);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *getPluginName() const NOEXCEPT override{
|
||||
return "ReshapeRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const PluginFieldCollection *getFieldNames() NOEXCEPT override{
|
||||
return &mFC;
|
||||
}
|
||||
|
||||
private:
|
||||
static PluginFieldCollection mFC;
|
||||
static std::vector<PluginField> mPluginAttributes;
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
REGISTER_TENSORRT_PLUGIN(ReshapeRTPluginCreator);
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
#include<cassert>
|
||||
#include "../kernels.h"
|
||||
|
||||
class ResizeLayerRT : public IPlugin {
|
||||
class ResizeLayerRT : public IPluginV2 {
|
||||
|
||||
public:
|
||||
ResizeLayerRT(int c, int h, int w) {
|
||||
@@ -10,35 +10,41 @@ public:
|
||||
o_w = w;
|
||||
}
|
||||
|
||||
ResizeLayerRT(const void *data,size_t length){
|
||||
const char *buf = reinterpret_cast<const char*>(data),*bufCheck = buf;
|
||||
o_c = readBUF<int>(buf);
|
||||
o_h = readBUF<int>(buf);
|
||||
o_w = readBUF<int>(buf);
|
||||
i_c = readBUF<int>(buf);
|
||||
i_h = readBUF<int>(buf);
|
||||
i_w = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + length);
|
||||
}
|
||||
|
||||
~ResizeLayerRT(){
|
||||
}
|
||||
|
||||
int getNbOutputs() const override {
|
||||
int getNbOutputs() const NOEXCEPT override {
|
||||
return 1;
|
||||
}
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override {
|
||||
return DimsCHW{o_c, o_h, o_w};
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) NOEXCEPT override {
|
||||
return Dims3{o_c, o_h, o_w};
|
||||
}
|
||||
|
||||
void configure(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, int maxBatchSize) override {
|
||||
void configureWithFormat(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs,DataType type,PluginFormat format,int maxBatchSize) NOEXCEPT override {
|
||||
i_c = inputDims[0].d[0];
|
||||
i_h = inputDims[0].d[1];
|
||||
i_w = inputDims[0].d[2];
|
||||
}
|
||||
|
||||
int initialize() override {
|
||||
return 0;
|
||||
}
|
||||
int initialize() NOEXCEPT override {return 0;}
|
||||
|
||||
virtual void terminate() override {
|
||||
}
|
||||
virtual void terminate() NOEXCEPT override {}
|
||||
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const override {
|
||||
return 0;
|
||||
}
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const NOEXCEPT override { return 0;}
|
||||
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override {
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void* const* outputs, void* workspace, cudaStream_t stream) NOEXCEPT override {
|
||||
// printf("%d %d %d %d %d %d\n", i_c, i_w, i_h, o_c, o_w, o_h);
|
||||
resizeForward((dnnType*)reinterpret_cast<const dnnType*>(inputs[0]),
|
||||
reinterpret_cast<dnnType*>(outputs[0]),
|
||||
@@ -47,11 +53,11 @@ public:
|
||||
}
|
||||
|
||||
|
||||
virtual size_t getSerializationSize() override {
|
||||
virtual size_t getSerializationSize() const NOEXCEPT override {
|
||||
return 6*sizeof(int);
|
||||
}
|
||||
|
||||
virtual void serialize(void* buffer) override {
|
||||
virtual void serialize(void* buffer) const NOEXCEPT override {
|
||||
char *buf = reinterpret_cast<char*>(buffer),*a=buf;
|
||||
|
||||
tk::dnn::writeBUF(buf, o_c);
|
||||
@@ -64,5 +70,96 @@ public:
|
||||
assert(buf == a + getSerializationSize());
|
||||
}
|
||||
|
||||
bool supportsFormat(DataType type,PluginFormat format) const NOEXCEPT override{
|
||||
return true;
|
||||
//todo assert
|
||||
}
|
||||
|
||||
const char *getPluginType() const NOEXCEPT override{
|
||||
return "ResizeLayerRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
void destroy() NOEXCEPT override{delete this;}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
IPluginV2 *clone() const NOEXCEPT override{
|
||||
ResizeLayerRT *p = new ResizeLayerRT(o_c,o_h,o_w);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
int i_c, i_h, i_w, o_c, o_h, o_w;
|
||||
|
||||
private:
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
class ResizeLayerRTPluginCreator : public IPluginCreator{
|
||||
public:
|
||||
ResizeLayerRTPluginCreator(){
|
||||
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));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *deserializePlugin(const char *name,const void *serialData,size_t serialLength) NOEXCEPT override{
|
||||
ResizeLayerRT *pluginObj = new ResizeLayerRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *createPlugin(const char *name,const PluginFieldCollection *fc) NOEXCEPT override{
|
||||
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);
|
||||
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));
|
||||
ResizeLayerRT *pluginObj = new ResizeLayerRT(oc,oh,ow);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *getPluginName() const NOEXCEPT override{
|
||||
return "ResizeLayerRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const PluginFieldCollection *getFieldNames() NOEXCEPT override{
|
||||
return &mFC;
|
||||
}
|
||||
|
||||
|
||||
private:
|
||||
static PluginFieldCollection mFC;
|
||||
static std::vector<PluginField> mPluginAttributes;
|
||||
std::string mPluginNamespace;
|
||||
|
||||
};
|
||||
|
||||
REGISTER_TENSORRT_PLUGIN(ResizeLayerRTPluginCreator);
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
#include<cassert>
|
||||
#include "../kernels.h"
|
||||
|
||||
class RouteRT : public IPlugin {
|
||||
class RouteRT : public IPluginV2 {
|
||||
|
||||
/**
|
||||
THIS IS NOT USED ANYMORE
|
||||
@@ -17,17 +17,31 @@ public:
|
||||
|
||||
}
|
||||
|
||||
int getNbOutputs() const override {
|
||||
RouteRT(const void* data,size_t length){
|
||||
const char* buf = reinterpret_cast<const char*>(data),*bufCheck = buf;
|
||||
groups = readBUF<int>(buf);
|
||||
group_id = readBUF<int>(buf);
|
||||
in = readBUF<int>(buf);
|
||||
for(int i=0;i <MAX_INPUTS;i++){
|
||||
c_in[i] = readBUF<int>(buf);
|
||||
}
|
||||
c= readBUF<int>(buf);
|
||||
h = readBUF<int>(buf);
|
||||
w = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + length);
|
||||
}
|
||||
|
||||
int getNbOutputs() const NOEXCEPT override {
|
||||
return 1;
|
||||
}
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override {
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) NOEXCEPT override {
|
||||
int out_c = 0;
|
||||
for(int i=0; i<nbInputDims; i++) out_c += inputs[i].d[0];
|
||||
return DimsCHW{out_c/groups, inputs[0].d[1], inputs[0].d[2]};
|
||||
return Dims3{out_c/groups, inputs[0].d[1], inputs[0].d[2]};
|
||||
}
|
||||
|
||||
void configure(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, int maxBatchSize) override {
|
||||
void configureWithFormat(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs,DataType type,PluginFormat format,int maxBatchSize) NOEXCEPT override {
|
||||
in = nbInputs;
|
||||
c = 0;
|
||||
for(int i=0; i<nbInputs; i++) {
|
||||
@@ -39,22 +53,14 @@ public:
|
||||
c /= groups;
|
||||
}
|
||||
|
||||
int initialize() override {
|
||||
int initialize() NOEXCEPT override { return 0;}
|
||||
|
||||
return 0;
|
||||
}
|
||||
virtual void terminate() NOEXCEPT override {}
|
||||
|
||||
virtual void terminate() override {
|
||||
}
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const NOEXCEPT override {return 0;}
|
||||
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const override {
|
||||
return 0;
|
||||
}
|
||||
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override {
|
||||
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void* const* outputs, void* workspace, cudaStream_t stream) NOEXCEPT override {
|
||||
dnnType *dstData = reinterpret_cast<dnnType*>(outputs[0]);
|
||||
|
||||
for(int b=0; b<batchSize; b++) {
|
||||
int offset = 0;
|
||||
for(int i=0; i<in; i++) {
|
||||
@@ -65,16 +71,14 @@ public:
|
||||
offset += part_in_dim;
|
||||
}
|
||||
}
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
|
||||
virtual size_t getSerializationSize() override {
|
||||
virtual size_t getSerializationSize() const NOEXCEPT override {
|
||||
return (6+MAX_INPUTS)*sizeof(int);
|
||||
}
|
||||
|
||||
virtual void serialize(void* buffer) override {
|
||||
virtual void serialize(void* buffer) const NOEXCEPT override {
|
||||
char *buf = reinterpret_cast<char*>(buffer),*a=buf;
|
||||
tk::dnn::writeBUF(buf, groups);
|
||||
tk::dnn::writeBUF(buf, group_id);
|
||||
@@ -88,9 +92,90 @@ public:
|
||||
assert(buf == a + getSerializationSize());
|
||||
}
|
||||
|
||||
const char *getPluginType() const NOEXCEPT override{
|
||||
return "RouteRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
void destroy() NOEXCEPT override {delete this; }
|
||||
|
||||
const char* getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
bool supportsFormat(DataType type,PluginFormat format) const NOEXCEPT override { return true;}
|
||||
|
||||
IPluginV2* clone() const NOEXCEPT override{
|
||||
RouteRT *p = new RouteRT(groups,group_id);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
static const int MAX_INPUTS = 4;
|
||||
int in;
|
||||
int c_in[MAX_INPUTS];
|
||||
int c, h, w;
|
||||
int groups, group_id;
|
||||
private:
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
class RouteRTPluginCreator : public IPluginCreator{
|
||||
public:
|
||||
RouteRTPluginCreator(){
|
||||
mPluginAttributes.emplace_back(PluginField("groups",nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("group_id",nullptr,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char* pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *deserializePlugin(const char* name,const void* serialData,size_t serialLength) NOEXCEPT override{
|
||||
RouteRT *pluginObj = new RouteRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *createPlugin(const char* name,const PluginFieldCollection *fc) NOEXCEPT override{
|
||||
const PluginField *fields = fc->fields;
|
||||
assert(fc->nbFields == 2);
|
||||
assert(fields[0].type == PluginFieldType::kINT32);
|
||||
assert(fields[1].type == PluginFieldType::kINT32);
|
||||
int groups = *(static_cast<const int *>(fields[0].data));
|
||||
int group_id = *(static_cast<const int *>(fields[1].data));
|
||||
RouteRT *pluginObj = new RouteRT(groups,group_id);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *getPluginName() const NOEXCEPT override{
|
||||
return "RouteRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const PluginFieldCollection *getFieldNames() NOEXCEPT override{
|
||||
return &mFC;
|
||||
}
|
||||
private:
|
||||
static PluginFieldCollection mFC;
|
||||
static std::vector<PluginField> mPluginAttributes;
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
REGISTER_TENSORRT_PLUGIN(RouteRTPluginCreator);
|
||||
|
||||
@@ -1,47 +1,52 @@
|
||||
#include<cassert>
|
||||
#include "../kernels.h"
|
||||
|
||||
class ShortcutRT : public IPlugin {
|
||||
|
||||
class ShortcutRT : public IPluginV2 {
|
||||
|
||||
public:
|
||||
ShortcutRT(tk::dnn::dataDim_t bdim, bool mul) {
|
||||
this->bc = bdim.c;
|
||||
this->bh = bdim.h;
|
||||
this->bw = bdim.w;
|
||||
bDim = bdim;
|
||||
this->bc = bDim.c;
|
||||
this->bh = bDim.h;
|
||||
this->bw = bDim.w;
|
||||
this->mul = mul;
|
||||
}
|
||||
|
||||
~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;
|
||||
mul = readBUF<bool>(buf);
|
||||
c = readBUF<int>(buf);
|
||||
h = readBUF<int>(buf);
|
||||
w = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + length);
|
||||
}
|
||||
|
||||
int getNbOutputs() const override {
|
||||
return 1;
|
||||
int getNbOutputs() const NOEXCEPT override {return 1;}
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) NOEXCEPT override {
|
||||
return Dims3{inputs[0].d[0], inputs[0].d[1], inputs[0].d[2]};
|
||||
}
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override {
|
||||
return DimsCHW{inputs[0].d[0], inputs[0].d[1], inputs[0].d[2]};
|
||||
}
|
||||
|
||||
void configure(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, int maxBatchSize) override {
|
||||
void configureWithFormat(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs,DataType type,PluginFormat format,int maxBatchSize) NOEXCEPT override {
|
||||
c = inputDims[0].d[0];
|
||||
h = inputDims[0].d[1];
|
||||
w = inputDims[0].d[2];
|
||||
}
|
||||
|
||||
int initialize() override {
|
||||
int initialize() NOEXCEPT override {return 0;}
|
||||
|
||||
return 0;
|
||||
}
|
||||
virtual void terminate() NOEXCEPT override {}
|
||||
|
||||
virtual void terminate() override {
|
||||
}
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const NOEXCEPT override { return 0;}
|
||||
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const override {
|
||||
return 0;
|
||||
}
|
||||
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override {
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void* const* outputs, void* workspace, cudaStream_t stream) NOEXCEPT override {
|
||||
|
||||
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
|
||||
dnnType *srcDataBack = (dnnType*)reinterpret_cast<const dnnType*>(inputs[1]);
|
||||
@@ -54,11 +59,11 @@ public:
|
||||
}
|
||||
|
||||
|
||||
virtual size_t getSerializationSize() override {
|
||||
virtual size_t getSerializationSize() const NOEXCEPT override {
|
||||
return 6*sizeof(int) + sizeof(bool);
|
||||
}
|
||||
|
||||
virtual void serialize(void* buffer) override {
|
||||
virtual void serialize(void* buffer) const NOEXCEPT override {
|
||||
char *buf = reinterpret_cast<char*>(buffer),*a=buf;
|
||||
tk::dnn::writeBUF(buf, bc);
|
||||
tk::dnn::writeBUF(buf, bh);
|
||||
@@ -71,7 +76,91 @@ public:
|
||||
|
||||
}
|
||||
|
||||
bool supportsFormat(DataType type,PluginFormat format) const NOEXCEPT override{
|
||||
return true;
|
||||
}
|
||||
|
||||
const char* getPluginType() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const char* getPluginVersion() const NOEXCEPT override{
|
||||
return "ShortcutRT_tkDNN";
|
||||
}
|
||||
|
||||
void destroy() NOEXCEPT override{delete this;}
|
||||
|
||||
const char* getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char* pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
IPluginV2 *clone() const NOEXCEPT override{
|
||||
ShortcutRT *p = new ShortcutRT(bDim,mul);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
int c, h, w;
|
||||
int bc, bh, bw;
|
||||
bool mul;
|
||||
tk::dnn::dataDim_t bDim;
|
||||
private:
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
|
||||
class ShortcutRTPluginCreator : public IPluginCreator {
|
||||
public:
|
||||
ShortcutRTPluginCreator(){
|
||||
mPluginAttributes.emplace_back(PluginField("bDim",nullptr,PluginFieldType::kUNKNOWN,1));
|
||||
mPluginAttributes.emplace_back(PluginField("mul",nullptr,PluginFieldType::kUNKNOWN,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char* pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *deserializePlugin(const char *name,const void *serialData,size_t serialLength) NOEXCEPT override{
|
||||
ShortcutRT *pluginObj = new ShortcutRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *createPlugin(const char *name,const PluginFieldCollection *fc) NOEXCEPT override{
|
||||
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));
|
||||
ShortcutRT *pluginObj = new ShortcutRT(bdim,mul);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *getPluginName() const NOEXCEPT override{
|
||||
return "ShortcutRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const PluginFieldCollection *getFieldNames() NOEXCEPT override{
|
||||
return &mFC;
|
||||
}
|
||||
private:
|
||||
static PluginFieldCollection mFC;
|
||||
static std::vector<PluginField> mPluginAttributes;
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
REGISTER_TENSORRT_PLUGIN(ShortcutRTPluginCreator);
|
||||
@@ -1,44 +1,47 @@
|
||||
#include<cassert>
|
||||
#include "../kernels.h"
|
||||
|
||||
class UpsampleRT : public IPlugin {
|
||||
|
||||
class UpsampleRT : public IPluginV2 {
|
||||
|
||||
public:
|
||||
UpsampleRT(int stride) {
|
||||
this->stride = stride;
|
||||
}
|
||||
|
||||
~UpsampleRT(){
|
||||
|
||||
UpsampleRT(const void *data,size_t length){
|
||||
const char* buf = reinterpret_cast<const char*>(data),*bufCheck=buf;
|
||||
stride = readBUF<int>(buf);
|
||||
c = readBUF<int>(buf);
|
||||
h = readBUF<int>(buf);
|
||||
w = readBUF<int>(buf);
|
||||
assert(buf == bufCheck + length);
|
||||
}
|
||||
|
||||
int getNbOutputs() const override {
|
||||
|
||||
~UpsampleRT(){}
|
||||
|
||||
int getNbOutputs() const NOEXCEPT override {
|
||||
return 1;
|
||||
}
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override {
|
||||
return DimsCHW(inputs[0].d[0], inputs[0].d[1]*stride, inputs[0].d[2]*stride);
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) NOEXCEPT override {
|
||||
return Dims3(inputs[0].d[0], inputs[0].d[1]*stride, inputs[0].d[2]*stride);
|
||||
}
|
||||
|
||||
void configure(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, int maxBatchSize) override {
|
||||
void configureWithFormat (const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs,DataType type,PluginFormat format,int maxBatchSize) NOEXCEPT override {
|
||||
c = inputDims[0].d[0];
|
||||
h = inputDims[0].d[1];
|
||||
w = inputDims[0].d[2];
|
||||
}
|
||||
|
||||
int initialize() override {
|
||||
int initialize() NOEXCEPT override {return 0;}
|
||||
|
||||
return 0;
|
||||
}
|
||||
virtual void terminate() NOEXCEPT override {}
|
||||
|
||||
virtual void terminate() override {
|
||||
}
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const NOEXCEPT override { return 0;}
|
||||
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const override {
|
||||
return 0;
|
||||
}
|
||||
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override {
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void* const* outputs, void* workspace, cudaStream_t stream) NOEXCEPT override {
|
||||
|
||||
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
|
||||
dnnType *dstData = reinterpret_cast<dnnType*>(outputs[0]);
|
||||
@@ -49,11 +52,9 @@ public:
|
||||
}
|
||||
|
||||
|
||||
virtual size_t getSerializationSize() override {
|
||||
return 4*sizeof(int);
|
||||
}
|
||||
virtual size_t getSerializationSize() const NOEXCEPT override { return 4*sizeof(int);}
|
||||
|
||||
virtual void serialize(void* buffer) override {
|
||||
virtual void serialize(void* buffer) const NOEXCEPT override {
|
||||
char *buf = reinterpret_cast<char*>(buffer),*a=buf;
|
||||
tk::dnn::writeBUF(buf, stride);
|
||||
tk::dnn::writeBUF(buf, c);
|
||||
@@ -62,5 +63,85 @@ public:
|
||||
assert(buf == a + getSerializationSize());
|
||||
}
|
||||
|
||||
bool supportsFormat(DataType type,PluginFormat format) const NOEXCEPT override{
|
||||
//todo assert
|
||||
return true;
|
||||
}
|
||||
|
||||
const char *getPluginType() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "UpsampleRT_tkDNN";
|
||||
}
|
||||
|
||||
void destroy() NOEXCEPT override{delete this;}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
IPluginV2* clone() const NOEXCEPT override{
|
||||
UpsampleRT *p = new UpsampleRT(stride);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
int c, h, w, stride;
|
||||
private:
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
class UpsampleRTPluginCreator : public IPluginCreator{
|
||||
public:
|
||||
UpsampleRTPluginCreator(){
|
||||
mPluginAttributes.emplace_back(PluginField("stride",nullptr,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char* pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *deserializePlugin(const char* name,const void* serialData,size_t serialLength) NOEXCEPT override{
|
||||
UpsampleRT *pluginObj = new UpsampleRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *createPlugin(const char* name,const PluginFieldCollection *fc) NOEXCEPT override{
|
||||
const PluginField *fields = fc->fields;
|
||||
int stride = *(static_cast<const int *>(fields[0].data));
|
||||
UpsampleRT *pluginObj = new UpsampleRT(stride);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *getPluginName() const NOEXCEPT override{
|
||||
return "UpsampleRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const PluginFieldCollection *getFieldNames() NOEXCEPT override{
|
||||
return &mFC;
|
||||
}
|
||||
private:
|
||||
static PluginFieldCollection mFC;
|
||||
static std::vector<PluginField> mPluginAttributes;
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
REGISTER_TENSORRT_PLUGIN(UpsampleRTPluginCreator);
|
||||
|
||||
+236
-111
@@ -1,143 +1,268 @@
|
||||
#include<cassert>
|
||||
#include "../kernels.h"
|
||||
|
||||
#define YOLORT_CLASSNAME_W 256
|
||||
|
||||
class YoloRT : public IPlugin {
|
||||
|
||||
|
||||
class YoloRT : public IPluginV2 {
|
||||
|
||||
public:
|
||||
YoloRT(int classes, int num, tk::dnn::Yolo *yolo = nullptr, int n_masks=3, float scale_xy=1, float nms_thresh=0.45, int nms_kind=0, int new_coords=0) {
|
||||
|
||||
this->classes = classes;
|
||||
this->num = num;
|
||||
this->n_masks = n_masks;
|
||||
this->scaleXY = scale_xy;
|
||||
this->nms_thresh = nms_thresh;
|
||||
this->nms_kind = nms_kind;
|
||||
this->new_coords = new_coords;
|
||||
YoloRT(int classes, int num, tk::dnn::Yolo *Yolo = nullptr, int n_masks = 3, float scale_xy = 1,
|
||||
float nms_thresh = 0.45, int nms_kind = 0, int new_coords = 0) {
|
||||
this->yolo = Yolo;
|
||||
this->classes = classes;
|
||||
this->num = num;
|
||||
this->n_masks = n_masks;
|
||||
this->scaleXY = scale_xy;
|
||||
this->nms_thresh = nms_thresh;
|
||||
this->nms_kind = nms_kind;
|
||||
this->new_coords = new_coords;
|
||||
|
||||
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;
|
||||
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){
|
||||
const char* buf = reinterpret_cast<const char*>(data),*bufCheck = buf;
|
||||
classes = readBUF<int>(buf);
|
||||
num = readBUF<int>(buf);
|
||||
n_masks = readBUF<int>(buf);
|
||||
scaleXY = readBUF<float>(buf);
|
||||
nms_thresh = readBUF<float>(buf);
|
||||
nms_kind = readBUF<int>(buf);
|
||||
new_coords = readBUF<int>(buf);
|
||||
c = readBUF<int>(buf);
|
||||
h = readBUF<int>(buf);
|
||||
w = readBUF<int>(buf);
|
||||
for(int i=0;i<n_masks;i++)
|
||||
mask[i] = readBUF<dnnType>(buf);
|
||||
for(int i=0;i<n_masks*2*num;i++)
|
||||
bias[i] = readBUF<dnnType>(buf);
|
||||
classesNames.resize(classes);
|
||||
for(int i=0;i<classes;i++){
|
||||
char tmp[YOLORT_CLASSNAME_W];
|
||||
for(int j=0;j<YOLORT_CLASSNAME_W;j++)
|
||||
tmp[j] = readBUF<char>(buf);
|
||||
classesNames[1] = std::string(tmp);
|
||||
}
|
||||
assert(buf == bufCheck + length);
|
||||
}
|
||||
|
||||
}
|
||||
~YoloRT() {
|
||||
|
||||
int getNbOutputs() const override {
|
||||
return 1;
|
||||
}
|
||||
|
||||
Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) override {
|
||||
return inputs[0];
|
||||
}
|
||||
|
||||
void configure(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, int maxBatchSize) override {
|
||||
c = inputDims[0].d[0];
|
||||
h = inputDims[0].d[1];
|
||||
w = inputDims[0].d[2];
|
||||
}
|
||||
|
||||
int initialize() override {
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
virtual void terminate() override {
|
||||
}
|
||||
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const override {
|
||||
return 0;
|
||||
}
|
||||
|
||||
virtual int enqueue(int batchSize, const void*const * inputs, void** outputs, void* workspace, cudaStream_t stream) override {
|
||||
|
||||
dnnType *srcData = (dnnType*)reinterpret_cast<const dnnType*>(inputs[0]);
|
||||
dnnType *dstData = reinterpret_cast<dnnType*>(outputs[0]);
|
||||
|
||||
checkCuda( cudaMemcpyAsync(dstData, srcData, batchSize*c*h*w*sizeof(dnnType), cudaMemcpyDeviceToDevice, stream));
|
||||
}
|
||||
|
||||
|
||||
for (int b = 0; b < batchSize; ++b){
|
||||
for(int n = 0; n < n_masks; ++n){
|
||||
int index = entry_index(b, n*w*h, 0);
|
||||
if (new_coords == 1){
|
||||
if (this->scaleXY != 1) scalAdd(dstData + index, 2 * w*h, this->scaleXY, -0.5*(this->scaleXY - 1), 1);
|
||||
}
|
||||
else{
|
||||
activationLOGISTICForward(srcData + index, dstData + index, 2*w*h, stream); //x,y
|
||||
int getNbOutputs() const NOEXCEPT override {
|
||||
return 1;
|
||||
}
|
||||
|
||||
if (this->scaleXY != 1) scalAdd(dstData + index, 2 * w*h, this->scaleXY, -0.5*(this->scaleXY - 1), 1);
|
||||
Dims getOutputDimensions(int index, const Dims *inputs, int nbInputDims) NOEXCEPT override {
|
||||
return inputs[0];
|
||||
}
|
||||
|
||||
index = entry_index(b, n*w*h, 4);
|
||||
activationLOGISTICForward(srcData + index, dstData + index, (1+classes)*w*h, stream);
|
||||
void configureWithFormat(const Dims *inputDims, int nbInputs, const Dims *outputDims, int nbOutputs, DataType type,
|
||||
PluginFormat format, int maxBatchSize) NOEXCEPT override {
|
||||
c = inputDims[0].d[0];
|
||||
h = inputDims[0].d[1];
|
||||
w = inputDims[0].d[2];
|
||||
}
|
||||
|
||||
int initialize() NOEXCEPT override {
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
virtual void terminate() NOEXCEPT override {
|
||||
}
|
||||
|
||||
virtual size_t getWorkspaceSize(int maxBatchSize) const NOEXCEPT override {
|
||||
return 0;
|
||||
}
|
||||
|
||||
virtual int enqueue(int batchSize, const void *const *inputs, void *const *outputs, void *workspace,
|
||||
cudaStream_t stream) NOEXCEPT override {
|
||||
|
||||
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));
|
||||
|
||||
|
||||
for (int b = 0; b < batchSize; ++b) {
|
||||
for (int n = 0; n < n_masks; ++n) {
|
||||
int index = entry_index(b, n * w * h, 0);
|
||||
if (new_coords == 1) {
|
||||
if (this->scaleXY != 1)
|
||||
scalAdd(dstData + index, 2 * w * h, this->scaleXY, -0.5 * (this->scaleXY - 1), 1);
|
||||
} else {
|
||||
activationLOGISTICForward(srcData + index, dstData + index, 2 * w * h, stream); //x,y
|
||||
|
||||
if (this->scaleXY != 1)
|
||||
scalAdd(dstData + index, 2 * w * h, this->scaleXY, -0.5 * (this->scaleXY - 1), 1);
|
||||
|
||||
index = entry_index(b, n * w * h, 4);
|
||||
activationLOGISTICForward(srcData + index, dstData + index, (1 + classes) * w * h, stream);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
//std::cout<<"YOLO END\n";
|
||||
return 0;
|
||||
}
|
||||
//std::cout<<"YOLO END\n";
|
||||
return 0;
|
||||
}
|
||||
|
||||
|
||||
virtual size_t getSerializationSize() override {
|
||||
return 8*sizeof(int) + 2*sizeof(float)+ n_masks*sizeof(dnnType) + num*n_masks*2*sizeof(dnnType) + YOLORT_CLASSNAME_W*classes*sizeof(char);
|
||||
}
|
||||
virtual size_t getSerializationSize() const NOEXCEPT override {
|
||||
return 8 * sizeof(int) + 2 * sizeof(float) + n_masks * sizeof(dnnType) + num * n_masks * 2 * sizeof(dnnType) +
|
||||
YOLORT_CLASSNAME_W * classes * sizeof(char);
|
||||
}
|
||||
|
||||
virtual void serialize(void* buffer) override {
|
||||
char *buf = reinterpret_cast<char*>(buffer),*a=buf;
|
||||
tk::dnn::writeBUF(buf, classes); //std::cout << "Classes :" << classes << std::endl;
|
||||
tk::dnn::writeBUF(buf, num); //std::cout << "Num : " << num << std::endl;
|
||||
tk::dnn::writeBUF(buf, n_masks); //std::cout << "N_Masks" << n_masks << std::endl;
|
||||
tk::dnn::writeBUF(buf, scaleXY); //std::cout << "ScaleXY :" << scaleXY << std::endl;
|
||||
tk::dnn::writeBUF(buf, nms_thresh); //std::cout << "nms_thresh :" << nms_thresh << std::endl;
|
||||
tk::dnn::writeBUF(buf, nms_kind); //std::cout << "nms_kind : " << nms_kind << std::endl;
|
||||
tk::dnn::writeBUF(buf, new_coords); //std::cout << "new_coords : " << new_coords << std::endl;
|
||||
tk::dnn::writeBUF(buf, c); //std::cout << "C : " << c << std::endl;
|
||||
tk::dnn::writeBUF(buf, h); //std::cout << "H : " << h << std::endl;
|
||||
tk::dnn::writeBUF(buf, w); //std::cout << "C : " << c << std::endl;
|
||||
for (int i = 0; i < n_masks; i++)
|
||||
{
|
||||
tk::dnn::writeBUF(buf, mask[i]); //std::cout << "mask[i] : " << mask[i] << std::endl;
|
||||
}
|
||||
for (int i = 0; i < n_masks * 2 * num; i++)
|
||||
{
|
||||
tk::dnn::writeBUF(buf, bias[i]); //std::cout << "bias[i] : " << bias[i] << std::endl;
|
||||
}
|
||||
bool supportsFormat(DataType type, PluginFormat format) const NOEXCEPT override {
|
||||
return true; //todo implement proper supportsFormat
|
||||
}
|
||||
|
||||
// save classes names
|
||||
for(int i=0; i<classes; i++) {
|
||||
char tmp[YOLORT_CLASSNAME_W];
|
||||
strcpy(tmp, classesNames[i].c_str());
|
||||
for(int j=0; j<YOLORT_CLASSNAME_W; j++) {
|
||||
tk::dnn::writeBUF(buf, tmp[j]);
|
||||
}
|
||||
}
|
||||
assert(buf == a + getSerializationSize());
|
||||
}
|
||||
virtual void serialize(void *buffer) const NOEXCEPT override {
|
||||
char *buf = reinterpret_cast<char *>(buffer), *a = buf;
|
||||
tk::dnn::writeBUF(buf, classes); //std::cout << "Classes :" << classes << std::endl;
|
||||
tk::dnn::writeBUF(buf, num); //std::cout << "Num : " << num << std::endl;
|
||||
tk::dnn::writeBUF(buf, n_masks); //std::cout << "N_Masks" << n_masks << std::endl;
|
||||
tk::dnn::writeBUF(buf, scaleXY); //std::cout << "ScaleXY :" << scaleXY << std::endl;
|
||||
tk::dnn::writeBUF(buf, nms_thresh); //std::cout << "nms_thresh :" << nms_thresh << std::endl;
|
||||
tk::dnn::writeBUF(buf, nms_kind); //std::cout << "nms_kind : " << nms_kind << std::endl;
|
||||
tk::dnn::writeBUF(buf, new_coords); //std::cout << "new_coords : " << new_coords << std::endl;
|
||||
tk::dnn::writeBUF(buf, c); //std::cout << "C : " << c << std::endl;
|
||||
tk::dnn::writeBUF(buf, h); //std::cout << "H : " << h << std::endl;
|
||||
tk::dnn::writeBUF(buf, w); //std::cout << "C : " << c << std::endl;
|
||||
for (int i = 0; i < n_masks; i++) {
|
||||
tk::dnn::writeBUF(buf, mask[i]); //std::cout << "mask[i] : " << mask[i] << std::endl;
|
||||
}
|
||||
for (int i = 0; i < n_masks * 2 * num; i++) {
|
||||
tk::dnn::writeBUF(buf, bias[i]); //std::cout << "bias[i] : " << bias[i] << std::endl;
|
||||
}
|
||||
|
||||
int c, h, w;
|
||||
// save classes names
|
||||
for (int i = 0; i < classes; i++) {
|
||||
char tmp[YOLORT_CLASSNAME_W];
|
||||
strcpy(tmp, classesNames[i].c_str());
|
||||
for (int j = 0; j < YOLORT_CLASSNAME_W; j++) {
|
||||
tk::dnn::writeBUF(buf, tmp[j]);
|
||||
}
|
||||
}
|
||||
assert(buf == a + getSerializationSize());
|
||||
}
|
||||
|
||||
const char *getPluginType() const NOEXCEPT override {
|
||||
return "YoloRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override {
|
||||
return "1";
|
||||
}
|
||||
|
||||
void destroy() NOEXCEPT override { delete this; }
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override {
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override {
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
IPluginV2 *clone() const NOEXCEPT override {
|
||||
YoloRT *p = new YoloRT(classes, num,yolo, n_masks, scaleXY, nms_thresh, nms_kind, new_coords);
|
||||
p->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return p;
|
||||
}
|
||||
|
||||
Yolo *yolo;
|
||||
int c, h, w;
|
||||
int classes, num, n_masks;
|
||||
float scaleXY;
|
||||
float nms_thresh;
|
||||
int nms_kind;
|
||||
int new_coords;
|
||||
std::vector<std::string> classesNames;
|
||||
float scaleXY;
|
||||
float nms_thresh;
|
||||
int nms_kind;
|
||||
int new_coords;
|
||||
std::vector<std::string> classesNames;
|
||||
|
||||
dnnType *mask;
|
||||
dnnType *bias;
|
||||
|
||||
int entry_index(int batch, int location, int entry) {
|
||||
int n = location / (w*h);
|
||||
int loc = location % (w*h);
|
||||
return batch*c*h*w + n*w*h*(4+classes+1) + entry*w*h + loc;
|
||||
}
|
||||
int entry_index(int batch, int location, int entry) {
|
||||
int n = location / (w * h);
|
||||
int loc = location % (w * h);
|
||||
return batch * c * h * w + n * w * h * (4 + classes + 1) + entry * w * h + loc;
|
||||
}
|
||||
|
||||
private:
|
||||
std::string mPluginNamespace;
|
||||
|
||||
};
|
||||
|
||||
class YoloRTPluginCreator : public IPluginCreator{
|
||||
public:
|
||||
YoloRTPluginCreator(){
|
||||
mPluginAttributes.emplace_back(PluginField("classes",nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("num",nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("yolo",nullptr,PluginFieldType::kUNKNOWN,1));
|
||||
mPluginAttributes.emplace_back(PluginField("numMasks",nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("scaleXY",nullptr,PluginFieldType::kFLOAT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("nmsThresh",nullptr,PluginFieldType::kFLOAT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("nmsKind",nullptr,PluginFieldType::kINT32,1));
|
||||
mPluginAttributes.emplace_back(PluginField("newCoords",nullptr,PluginFieldType::kINT32,1));
|
||||
mFC.nbFields = mPluginAttributes.size();
|
||||
mFC.fields = mPluginAttributes.data();
|
||||
}
|
||||
|
||||
void setPluginNamespace(const char *pluginNamespace) NOEXCEPT override{
|
||||
mPluginNamespace = pluginNamespace;
|
||||
}
|
||||
|
||||
const char *getPluginNamespace() const NOEXCEPT override{
|
||||
return mPluginNamespace.c_str();
|
||||
}
|
||||
|
||||
IPluginV2 *deserializePlugin(const char *name,const void *serialData,size_t serialLength) NOEXCEPT override{
|
||||
YoloRT *pluginObj = new YoloRT(serialData,serialLength);
|
||||
pluginObj->setPluginNamespace(mPluginNamespace.c_str());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
IPluginV2 *createPlugin(const char* name,const PluginFieldCollection *fc) NOEXCEPT override{
|
||||
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());
|
||||
return pluginObj;
|
||||
}
|
||||
|
||||
const char *getPluginName() const NOEXCEPT override{
|
||||
return "YoloRT_tkDNN";
|
||||
}
|
||||
|
||||
const char *getPluginVersion() const NOEXCEPT override{
|
||||
return "1";
|
||||
}
|
||||
|
||||
const PluginFieldCollection *getFieldNames() NOEXCEPT override{
|
||||
return &mFC;
|
||||
}
|
||||
|
||||
private:
|
||||
static PluginFieldCollection mFC;
|
||||
static std::vector<PluginField> mPluginAttributes;
|
||||
std::string mPluginNamespace;
|
||||
};
|
||||
|
||||
REGISTER_TENSORRT_PLUGIN(YoloRTPluginCreator);
|
||||
|
||||
@@ -11,6 +11,7 @@
|
||||
#include "cuda_runtime_api.h"
|
||||
#include <cublas_v2.h>
|
||||
#include <cudnn.h>
|
||||
#include <NvInferVersion.h>
|
||||
|
||||
|
||||
#ifdef __linux__
|
||||
@@ -22,6 +23,15 @@
|
||||
#include <chrono>
|
||||
|
||||
|
||||
|
||||
|
||||
#if NV_TENSORRT_MAJOR > 7
|
||||
#define NOEXCEPT noexcept
|
||||
#else
|
||||
#define NOEXCEPT
|
||||
#endif
|
||||
|
||||
|
||||
#define dnnType float
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user