Files
Rajeev Rao aff45dd565 TensorRT OSS 8.0 release
Signed-off-by: Rajeev Rao <rajeevrao@nvidia.com>
2021-07-02 16:35:44 -07:00

290 lines
7.6 KiB
C++

/*
* Copyright (c) 2021, NVIDIA CORPORATION. All rights reserved.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
#include "scatterPlugin.h"
#include "half.h"
#include <cstring>
#include <cublas_v2.h>
#include <cudnn.h>
#include <iostream>
#include <sstream>
using namespace nvinfer1;
using nvinfer1::plugin::ScatterND;
using nvinfer1::plugin::ScatterNDPluginCreator;
namespace
{
const char* SCATTERND_PLUGIN_VERSION{"1"};
const char* SCATTERND_PLUGIN_NAME{"ScatterND"};
} // namespace
PluginFieldCollection ScatterNDPluginCreator::mFC{};
ScatterND::ScatterND()
{
}
int ScatterND::getNbOutputs() const noexcept
{
// Plugin layer has 1 output
return 1;
}
DimsExprs ScatterND::getOutputDimensions(int32_t outputIndex, const DimsExprs* inputs, int32_t nbInputs, IExprBuilder& exprBuilder) noexcept
{
//output should have same dimensions as data tensor
DimsExprs ret = inputs[dataTensorIdx];
return ret;
}
int ScatterND::initialize() noexcept
{
return 0;
}
void ScatterND::terminate() noexcept
{
}
bool ScatterND::supportsFormatCombination(int32_t pos, const PluginTensorDesc* inOut, int32_t nbInputs, int32_t nbOutputs) noexcept
{
ASSERT(pos < 4);
ASSERT(nbInputs == 3);
ASSERT(nbOutputs == 1);
const PluginTensorDesc& desc = inOut[pos];
bool ret = false;
switch (pos)
{
case dataTensorIdx:
case updateTensorIdx:
ret = ((desc.type == DataType::kFLOAT || desc.type == DataType::kINT32)
&& desc.format == TensorFormat::kLINEAR);
break;
case indexTensorIdx:
ret = (desc.type == DataType::kINT32 && desc.format == TensorFormat::kLINEAR);
break;
case 3:
ret = ((desc.type == DataType::kFLOAT || desc.type == DataType::kINT32) && desc.format == TensorFormat::kLINEAR);
break;
}
return ret;
}
void ScatterND::configurePlugin(const DynamicPluginTensorDesc* in, int32_t nbInputs, const DynamicPluginTensorDesc* out, int32_t nbOutputs) noexcept
{
}
int32_t ScatterND::calculateNumSlices(Dims indexTensorDims) const noexcept
{
int32_t nSlices = 1;
for (int i = 0; i < indexTensorDims.nbDims-1; i++)
{
nSlices *= indexTensorDims.d[i];
}
return nSlices;
}
size_t ScatterND::getWorkspaceSize(const PluginTensorDesc* inputs, int32_t nbInputs, const PluginTensorDesc* outputs,int32_t nbOutputs) const noexcept
{
int32_t nSlices = calculateNumSlices(inputs[indexTensorIdx].dims);
//transformCoeffs + transformed indices
return outputs[0].dims.MAX_DIMS * sizeof(int32_t) + nSlices * sizeof(int32_t);
}
void ScatterND::calculateTransformCoeff(const Dims& dataTensorDims, int indexRank, int32_t* transformCoeff) const noexcept
{
std::vector<int32_t> pitches;
for (int32_t i = indexRank - 1, nIndx = 1; i >= 0 ; i--)
{
pitches.push_back(nIndx);
nIndx *= dataTensorDims.d[i];
}
std::reverse(pitches.begin(), pitches.end()); //last dimension pitch is always one (assuming linear mem)
std::copy(pitches.begin(), pitches.end(), transformCoeff);
}
int32_t ScatterND::calculateCopySize(const Dims& dataDims) const noexcept
{
int32_t copySize = 1;
for (int i = 0; i < dataDims.nbDims; i++)
{
copySize *= dataDims.d[i];
}
copySize *= sizeof(float);
return copySize;
}
int32_t ScatterND::enqueue(const PluginTensorDesc* inputDesc, const PluginTensorDesc* outputDesc,
const void* const* inputs, void* const* outputs, void* workspace, cudaStream_t stream) noexcept
{
int32_t transformCoeff[outputDesc[0].dims.MAX_DIMS];
std::memset(transformCoeff, 0, sizeof(int32_t)*outputDesc[0].dims.MAX_DIMS);
Dims IndexDims = inputDesc[indexTensorIdx].dims;
Dims dataDims = inputDesc[dataTensorIdx].dims;
int32_t indexRank = IndexDims.d[IndexDims.nbDims-1];
ASSERT(indexRank <= dataDims.nbDims);
int32_t nSlices = calculateNumSlices(IndexDims);
int32_t rowSize = 1;
int32_t copySize = calculateCopySize(dataDims);
int32_t elementSizeInBytes = 1;
switch (inputDesc->type)
{
case DataType::kFLOAT:
case DataType::kINT32:
elementSizeInBytes = 4;
break;
case DataType::kHALF:
elementSizeInBytes = 2;
break;
case DataType::kINT8:
case DataType::kBOOL:
elementSizeInBytes = 1;
break;
}
for (int i = indexRank; i < dataDims.nbDims; i++)
{
rowSize *= dataDims.d[i];
}
calculateTransformCoeff(dataDims, indexRank, transformCoeff);
scatterNDInference(stream, transformCoeff,
dataDims.nbDims,
indexRank,
nSlices,
rowSize,
copySize,
elementSizeInBytes,
inputs[indexTensorIdx],
inputs[updateTensorIdx],
inputs[dataTensorIdx],
outputs[0],
workspace );
return 0;
}
size_t ScatterND::getSerializationSize() const noexcept
{
return 0;
}
void ScatterND::serialize(void* buffer) const noexcept
{
return;
}
// Set plugin namespace
void ScatterND::setPluginNamespace(const char* pluginNamespace) noexcept
{
mPluginNamespace = pluginNamespace;
}
const char* ScatterND::getPluginNamespace() const noexcept
{
return mPluginNamespace.c_str();
}
// Return the DataType of the plugin output at the requested index
DataType ScatterND::getOutputDataType(int index, const nvinfer1::DataType* inputTypes, int nbInputs) const noexcept
{
ASSERT(index == 0);
return inputTypes[dataTensorIdx];
}
// Attach the plugin object to an execution context and grant the plugin the access to some context resource.
void ScatterND::attachToContext(cudnnContext* cudnn, cublasContext* cublas, IGpuAllocator* gpuAllocator) noexcept
{
return;
}
// Detach the plugin object from its execution context.
void ScatterND::detachFromContext() noexcept {}
const char* ScatterND::getPluginType() const noexcept
{
return SCATTERND_PLUGIN_NAME;
}
const char* ScatterND::getPluginVersion() const noexcept
{
return SCATTERND_PLUGIN_VERSION;
}
void ScatterND::destroy() noexcept
{
delete this;
}
// Clone the plugin
IPluginV2DynamicExt* ScatterND::clone() const noexcept
{
// Create a new instance
ScatterND* plugin = new ScatterND();
plugin->setPluginNamespace(mPluginNamespace.c_str());
return plugin;
}
ScatterNDPluginCreator::ScatterNDPluginCreator()
{
mFC.nbFields = 0;
}
const char* ScatterNDPluginCreator::getPluginName() const noexcept
{
return SCATTERND_PLUGIN_NAME;
}
const char* ScatterNDPluginCreator::getPluginVersion() const noexcept
{
return SCATTERND_PLUGIN_VERSION;
}
const PluginFieldCollection* ScatterNDPluginCreator::getFieldNames() noexcept
{
return &mFC;
}
IPluginV2Ext* ScatterNDPluginCreator::createPlugin(const char* name, const PluginFieldCollection* fc) noexcept
{
ScatterND* obj = new ScatterND();
obj->setPluginNamespace(mNamespace.c_str());
return obj;
}
IPluginV2Ext* ScatterNDPluginCreator::deserializePlugin(const char* name, const void* serialData, size_t serialLength) noexcept
{
// This object will be deleted when the network is destroyed, which will
// call Normalize::destroy()
ScatterND* obj = new ScatterND();
obj->setPluginNamespace(mNamespace.c_str());
return obj;
}