/* * 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. */ #ifndef BATCHTILEPLUGIN_H #define BATCHTILEPLUGIN_H #include "NvInferPlugin.h" #include "plugin.h" #include namespace nvinfer1 { namespace plugin { class BatchTilePlugin : public IPluginV2Ext { public: BatchTilePlugin(const std::string name); BatchTilePlugin(const std::string name, size_t copy_size); BatchTilePlugin(const std::string name, const void* data, size_t length); BatchTilePlugin() = delete; int getNbOutputs() const noexcept override; Dims getOutputDimensions(int index, const Dims* inputs, int nbInputDims) noexcept override; int initialize() noexcept override; void terminate() noexcept override; size_t getWorkspaceSize(int) const noexcept override; int enqueue(int batchSize, const void* const* inputs, void* const* outputs, void* workspace, cudaStream_t stream) noexcept override; DataType getOutputDataType(int index, const nvinfer1::DataType* inputTypes, int nbInputs) const noexcept override; size_t getSerializationSize() const noexcept override; void serialize(void* buffer) const noexcept override; bool isOutputBroadcastAcrossBatch(int outputIndex, const bool* inputIsBroadcasted, int nbInputs) const noexcept override; bool canBroadcastInputAcrossBatch(int inputIndex) const noexcept override; void configurePlugin(const Dims* inputDims, int nbInputs, const Dims* outputDims, int nbOutputs, const DataType* inputTypes, const DataType* outputTypes, const bool* inputIsBroadcast, const bool* outputIsBroadcast, PluginFormat floatFormat, int maxBatchSize) noexcept override; bool supportsFormat(DataType type, PluginFormat format) const noexcept override; const char* getPluginType() const noexcept override; const char* getPluginVersion() const noexcept override; void destroy() noexcept override; IPluginV2Ext* clone() const noexcept override; void setPluginNamespace(const char* libNamespace) noexcept override; const char* getPluginNamespace() const noexcept override; private: const std::string mLayerName; size_t mCopySize; std::string mNamespace; }; class BatchTilePluginCreator : public BaseCreator { public: BatchTilePluginCreator(); const char* getPluginName() const noexcept override; const char* getPluginVersion() const noexcept override; const PluginFieldCollection* getFieldNames() noexcept override; IPluginV2Ext* createPlugin(const char* name, const PluginFieldCollection* fc) noexcept override; IPluginV2Ext* deserializePlugin(const char* name, const void* serialData, size_t serialLength) noexcept override; void setPluginNamespace(const char* libNamespace) noexcept override; const char* getPluginNamespace() const noexcept override { return mNamespace.c_str(); } private: static PluginFieldCollection mFC; std::string mNamespace; }; } // namespace plugin } // namespace nvinfer1 #endif