Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
NvInferRuntimePlugin.h981 linesDownload Raw Back to include
1/*
2 * SPDX-FileCopyrightText: Copyright (c) 1993-2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
3 * SPDX-License-Identifier: Apache-2.0
4 *
5 * Licensed under the Apache License, Version 2.0 (the "License");
6 * you may not use this file except in compliance with the License.
7 * You may obtain a copy of the License at
8 *
9 * http://www.apache.org/licenses/LICENSE-2.0
10 *
11 * Unless required by applicable law or agreed to in writing, software
12 * distributed under the License is distributed on an "AS IS" BASIS,
13 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
14 * See the License for the specific language governing permissions and
15 * limitations under the License.
16 */
17
18#ifndef NV_INFER_RUNTIME_PLUGIN_H
19#define NV_INFER_RUNTIME_PLUGIN_H
20
21#define NV_INFER_INTERNAL_INCLUDE 1
22#include "NvInferPluginBase.h"
23#undef NV_INFER_INTERNAL_INCLUDE
24
25//!
26//! \file NvInferRuntimePlugin.h
27//!
28//! This file contains common definitions, data structures and interfaces that relate to plugins and are shared
29//! between the standard and safe runtime.
30//!
31//! \warning Do not directly include this file. Instead include NvInferRuntime.h
32//!
33
34//!
35//! \namespace nvinfer1
36//!
37//! \brief The TensorRT API version 1 namespace.
38//!
39namespace nvinfer1
40{
41
42enum class TensorFormat : int32_t;
43namespace v_1_0
44{
45class IGpuAllocator;
46} // namespace v_1_0
47using IGpuAllocator = v_1_0::IGpuAllocator;
48
49//!
50//! \brief PluginFormat is reserved for backward compatibility.
51//!
52//! \see IPluginV2::supportsFormat()
53//!
54using PluginFormat = TensorFormat;
55
56//!
57//! \brief Bit at the plugin version to identify that it is a plugin.
58//!
59static constexpr int32_t kPLUGIN_VERSION_PYTHON_BIT = 0x40;
60
61//!
62//! \struct PluginTensorDesc
63//!
64//! \brief Fields that a plugin might see for an input or output.
65//!
66//! Scale is only valid when data type is DataType::kINT8. TensorRT will set
67//! the value to -1.0F if it is invalid.
68//!
69//! \see IPluginV2IOExt::supportsFormatCombination
70//! \see IPluginV2IOExt::configurePlugin
71//!
72struct PluginTensorDesc
73{
74    //! Dimensions.
75    Dims dims;
76    //! \warning DataType:kBOOL and DataType::kUINT8 are not supported.
77    DataType type;
78    //! Tensor format.
79    TensorFormat format;
80    //! Scale for INT8 data type.
81    float scale;
82};
83
84//!
85//! \struct PluginVersion
86//!
87//! \brief Definition of plugin versions.
88//!
89//! Tag for plug-in versions.  Used in upper byte of getTensorRTVersion().
90//!
91//! \deprecated Deprecated in TensorRT 10.10. PluginVersion is used only in relation to IPluginV2-descendent plugin
92//! interfaces, which are all deprecated.
93//!
94enum class PluginVersion : uint8_t
95{
96    //! IPluginV2
97    kV2 TRT_DEPRECATED_ENUM = 0,
98    //! IPluginV2Ext
99    kV2_EXT TRT_DEPRECATED_ENUM = 1,
100    //! IPluginV2IOExt
101    kV2_IOEXT TRT_DEPRECATED_ENUM = 2,
102    //! IPluginV2DynamicExt
103    kV2_DYNAMICEXT TRT_DEPRECATED_ENUM = 3,
104    //! IPluginV2DynamicExt-based Python plugins
105    kV2_DYNAMICEXT_PYTHON TRT_DEPRECATED_ENUM = kPLUGIN_VERSION_PYTHON_BIT | 3
106};
107
108//!
109//! \enum PluginCreatorVersion
110//!
111//! \brief Enum to identify version of the plugin creator.
112//!
113//! \deprecated Deprecated in TensorRT 10.10. PluginCreatorVersion is used only in relation to plugin creators based
114//! off IPluginCreator, which is deprecated.
115//!
116enum class PluginCreatorVersion : int32_t
117{
118    //! IPluginCreator
119    kV1 TRT_DEPRECATED_ENUM = 0,
120    //! IPluginCreator-based Python plugin creators
121    kV1_PYTHON TRT_DEPRECATED_ENUM = kPLUGIN_VERSION_PYTHON_BIT
122};
123
124//!
125//! \class IPluginV2
126//!
127//! \brief Plugin class for user-implemented layers.
128//!
129//! Plugins are a mechanism for applications to implement custom layers. When
130//! combined with IPluginCreator it provides a mechanism to register plugins and
131//! look up the Plugin Registry during de-serialization.
132//!
133//! \see IPluginCreator
134//! \see IPluginRegistry
135//!
136//! \deprecated Deprecated in TensorRT 8.5. Implement IPluginV3 instead.
137//!
138class TRT_DEPRECATED IPluginV2
139{
140public:
141    //!
142    //! \brief Return the API version with which this plugin was built.
143    //!
144    //! Do not override this method as it is used by the TensorRT library to maintain backwards-compatibility with
145    //! plugins.
146    //!
147    //! \return The TensorRT version in the format (major * 100 + minor) * 100 + patch.
148    //!
149    //! \usage
150    //! - Allowed context for the API call
151    //!   - Thread-safe: Yes, the implementation provided here is safe to call from any thread.
152    //!
153    virtual int32_t getTensorRTVersion() const noexcept
154    {
155        return NV_TENSORRT_VERSION;
156    }
157
158    //!
159    //! \brief Return the plugin type. Should match the plugin name returned by the corresponding plugin creator
160    //!
161    //! \see IPluginCreator::getPluginName()
162    //!
163    //! \warning The string returned must be NULL-terminated and have a length of 1024 bytes or less including the
164    //! NULL terminator.
165    //!
166    //! \usage
167    //! - Allowed context for the API call
168    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
169    //!                  when building networks on multiple devices sharing the same plugin.
170    //!
171    virtual AsciiChar const* getPluginType() const noexcept = 0;
172
173    //!
174    //! \brief Return the plugin version. Should match the plugin version returned by the corresponding plugin creator
175    //!
176    //! \see IPluginCreator::getPluginVersion()
177    //!
178    //! \warning The string returned must be NULL-terminated and have a length of 1024 bytes or less including the
179    //! NULL terminator.
180    //!
181    //! \usage
182    //! - Allowed context for the API call
183    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
184    //!                  when building networks on multiple devices sharing the same plugin.
185    //!
186    virtual AsciiChar const* getPluginVersion() const noexcept = 0;
187
188    //!
189    //! \brief Get the number of outputs from the layer.
190    //!
191    //! \return The number of outputs, which is a positive integer.
192    //!
193    //! This function is called by the implementations of INetworkDefinition and IBuilder. In particular, it is called
194    //! prior to any call to initialize().
195    //!
196    //! \usage
197    //! - Allowed context for the API call
198    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
199    //!                  when building networks on multiple devices sharing the same plugin.
200    //!
201    virtual int32_t getNbOutputs() const noexcept = 0;
202
203    //!
204    //! \brief Get the dimension of an output tensor.
205    //!
206    //! \param index The index of the output tensor. Will lie in the valid range (between 0 and getNbOutputs()-1
207    //! inclusive).
208    //! \param inputs The input tensor dimensions. Will be the start address of a Dims array of length nbInputDims.
209    //! \param nbInputDims The number of input tensors. Will be a non-negative integer.
210    //!
211    //! \return The output tensor dimensions if the index is in the valid range.
212    //!         An invalid value of Dims{-1, {}} must be returned if the index is not in the valid range.
213    //!
214    //! This function is called by the implementations of INetworkDefinition and IBuilder. In particular, it is called
215    //! prior to any call to initialize().
216    //!
217    //! \usage
218    //! - Allowed context for the API call
219    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
220    //!                  when building networks on multiple devices sharing the same plugin.
221    //!
222    //! \note In any non-IPluginV2DynamicExt plugin, batch size must not be included in the returned dimensions,
223    //! even if the plugin is expected to be run in a network with explicit batch mode enabled.
224    //! Please see the TensorRT Developer Guide for more details on how plugin inputs and outputs behave.
225    //!
226    virtual Dims getOutputDimensions(int32_t index, Dims const* inputs, int32_t nbInputDims) noexcept = 0;
227
228    //!
229    //! \brief Check format support.
230    //!
231    //! \param type DataType requested.
232    //! \param format PluginFormat requested.
233    //!
234    //! \return true if the plugin supports the type-format combination.
235    //!
236    //! This function is called by the implementations of INetworkDefinition, IBuilder, and
237    //! safe::ICudaEngine/ICudaEngine. In particular, it is called when creating an engine and when deserializing an
238    //! engine.
239    //!
240    //! \warning for the format field, the values PluginFormat::kCHW4, PluginFormat::kCHW16, and PluginFormat::kCHW32
241    //! will not be passed in, this is to keep backward compatibility with TensorRT 5.x series.  Use PluginV2IOExt
242    //! or PluginV2DynamicExt for other PluginFormats.
243    //!
244    //! \warning DataType:kBOOL and DataType::kUINT8 are not supported.
245    //!
246    //! \usage
247    //! - Allowed context for the API call
248    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
249    //!                  when building networks on multiple devices sharing the same plugin.
250    //!
251    virtual bool supportsFormat(DataType type, PluginFormat format) const noexcept = 0;
252
253    //!
254    //! \brief Configure the layer.
255    //!
256    //! This function is called by the builder prior to initialize(). It provides an opportunity for the layer to make
257    //! algorithm choices on the basis of its weights, dimensions, and maximum batch size.
258    //!
259    //! \param inputDims The input tensor dimensions. Will be the start address of a Dims array of length nbInputs.
260    //! \param nbInputs The number of inputs. Will be a non-negative integer.
261    //! \param outputDims The output tensor dimensions. Will be the start address of a Dims array of length nbOutputs.
262    //! \param nbOutputs The number of outputs. Will be a positive integer identical to the return value of
263    //! getNbOutputs().
264    //! \param type The data type selected for the engine.
265    //! \param format The format selected for the engine.
266    //! \param maxBatchSize The maximum batch size. Will be a positive integer.
267    //!
268    //! The dimensions passed here do not include the outermost batch size (i.e. for 2D image networks, they will be
269    //! 3-dimensional CHW dimensions).
270    //!
271    //! \warning for the format field, the values PluginFormat::kCHW4, PluginFormat::kCHW16, and PluginFormat::kCHW32
272    //! will not be passed in, this is to keep backward compatibility with TensorRT 5.x series.  Use PluginV2IOExt
273    //! or PluginV2DynamicExt for other PluginFormats.
274    //!
275    //! \warning DataType:kBOOL and DataType::kUINT8 are not supported.
276    //!
277    //! \see clone()
278    //!
279    //! \usage
280    //! - Allowed context for the API call
281    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
282    //!                  when building networks on multiple devices sharing the same plugin. However, TensorRT
283    //!                  will not call this method from two threads simultaneously on a given clone of a plugin.
284    //!
285    virtual void configureWithFormat(Dims const* inputDims, int32_t nbInputs, Dims const* outputDims, int32_t nbOutputs,
286        DataType type, PluginFormat format, int32_t maxBatchSize) noexcept
287        = 0;
288
289    //!
290    //! \brief Initialize the layer for execution. This is called when the engine is created.
291    //!
292    //! \return 0 for success, else non-zero (which will cause engine termination).
293    //!
294    //! \usage
295    //! - Allowed context for the API call
296    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
297    //!                  when building networks on multiple devices sharing the same plugin or when using multiple
298    //!                  execution contexts using this plugin.
299    //!
300    virtual int32_t initialize() noexcept = 0;
301
302    //!
303    //! \brief Release resources acquired during plugin layer initialization. This is called when the engine is
304    //! destroyed.
305    //!
306    //! \see initialize()
307    //!
308    //! \usage
309    //! - Allowed context for the API call
310    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
311    //!                  when building networks on multiple devices sharing the same plugin or when using multiple
312    //!                  execution contexts using this plugin. However, TensorRT will not call this method from
313    //!                  two threads simultaneously on a given clone of a plugin.
314    //!
315    virtual void terminate() noexcept = 0;
316
317    //!
318    //! \brief Find the workspace size required by the layer.
319    //!
320    //! This function is called during engine startup, after initialize(). The workspace size returned must be
321    //! sufficient for any batch size up to the maximum.
322    //!
323    //! \param maxBatchSize The maximum batch size, which will be a positive integer.
324    //!
325    //! \return The workspace size in bytes, i.e. the device memory size that the plugin requires for its internal
326    //! computations.
327    //!
328    //! \usage
329    //! - Allowed context for the API call
330    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
331    //!                  when building networks on multiple devices sharing the same plugin. However, TensorRT
332    //!                  will not call this method from two threads simultaneously on a given clone of a plugin.
333    //!
334    virtual size_t getWorkspaceSize(int32_t maxBatchSize) const noexcept = 0;
335
336    //!
337    //! \brief Execute the layer.
338    //!
339    //! \param batchSize The number of inputs in the batch.
340    //! \param inputs The memory for the input tensors. Will be an array of device addresses corresponding to input
341    //!        tensors of length nbInputs, where nbInputs is the second parameter passed to configureWithFormat().
342    //!        The i-th input tensor will have the dimensions inputDims[i], where inputDims is the first parameter
343    //!        that was passed to configureWithFormat().
344    //! \param outputs The memory for the output tensors. Will be an array of device addresses corresponding to output
345    //!        tensors of length getNbOutputs().
346    //! \param workspace Workspace for execution. Will be the start address of a device buffer whose length will be at
347    //!        least getWorkspaceSize(batchSize).
348    //! \param stream The stream in which to execute the kernels. This will be a valid CUDA stream.
349    //!
350    //! \return 0 for success, else non-zero (which will cause engine termination).
351    //!
352    //! \usage
353    //! - Allowed context for the API call
354    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
355    //!                  when multiple execution contexts are used during runtime.
356    //!
357    virtual int32_t enqueue(int32_t batchSize, void const* const* inputs, void* const* outputs, void* workspace,
358        cudaStream_t stream) noexcept
359        = 0;
360
361    //!
362    //! \brief Find the size of the serialization buffer required to store the plugin configuration in a binary file.
363    //!
364    //! \return The size of the serialization buffer in bytes.
365    //!
366    //! \usage
367    //! - Allowed context for the API call
368    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
369    //!                  when building networks on multiple devices sharing the same plugin.
370    //!
371    virtual size_t getSerializationSize() const noexcept = 0;
372
373    //!
374    //! \brief Serialize the layer.
375    //!
376    //! \param buffer A pointer to a host buffer to serialize data. Size of buffer will be at least as large as the
377    //! value returned by getSerializationSize.
378    //!
379    //! \see getSerializationSize()
380    //!
381    //! \usage
382    //! - Allowed context for the API call
383    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
384    //!                  when building networks on multiple devices sharing the same plugin.
385    //!
386    virtual void serialize(void* buffer) const noexcept = 0;
387
388    //!
389    //! \brief Destroy the plugin object. This will be called when the network, builder or engine is destroyed.
390    //!
391    //! \usage
392    //! - Allowed context for the API call
393    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
394    //!                  when building networks on multiple devices sharing the same plugin.
395    //!
396    virtual void destroy() noexcept = 0;
397
398    //!
399    //! \brief Clone the plugin object. This copies over internal plugin parameters and returns a new plugin object with
400    //! these parameters.
401    //!
402    //! The TensorRT runtime calls clone() to clone the plugin when an execution context is created for an engine,
403    //! after the engine has been created.  The runtime does not call initialize() on the cloned plugin,
404    //! so the cloned plugin must be created in an initialized state.
405    //!
406    //! \return A cloned plugin object in an initialized state with the same parameters as the current object.
407    //!         nullptr must be returned if the cloning fails, e.g. because of resource exhaustion.
408    //!
409    //! \usage
410    //! - Allowed context for the API call
411    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
412    //!                  when building networks on multiple devices sharing the same plugin or when creating multiple
413    //!                  execution contexts.
414    //!
415    virtual IPluginV2* clone() const noexcept = 0;
416
417    //!
418    //! \brief Set the namespace that this plugin object belongs to. Ideally, all plugin
419    //! objects from the same plugin library must have the same namespace.
420    //!
421    //! \param pluginNamespace The namespace for the plugin object.
422    //!
423    //! \warning The string pluginNamespace will be NULL-terminated and have a length of 1024 bytes or less including the
424    //! NULL terminator.
425    //!
426    //! \usage
427    //! - Allowed context for the API call
428    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
429    //!                  when building networks on multiple devices sharing the same plugin.
430    //!
431    virtual void setPluginNamespace(AsciiChar const* pluginNamespace) noexcept = 0;
432
433    //!
434    //! \brief Return the namespace of the plugin object.
435    //!
436    //! \return The namespace string that was passed to setPluginNamespace(), possibly after truncation to 1024 bytes
437    //! if a longer string was passed. An empty string must be returned as default value.
438    //!
439    //! \usage
440    //! - Allowed context for the API call
441    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
442    //!                  when building networks on multiple devices sharing the same plugin.
443    //!
444    virtual AsciiChar const* getPluginNamespace() const noexcept = 0;
445
446    // @cond SuppressDoxyWarnings
447    IPluginV2() = default;
448    virtual ~IPluginV2() noexcept = default;
449// @endcond
450
451protected:
452// @cond SuppressDoxyWarnings
453    IPluginV2(IPluginV2 const&) = default;
454    IPluginV2(IPluginV2&&) = default;
455    IPluginV2& operator=(IPluginV2 const&) & = default;
456    IPluginV2& operator=(IPluginV2&&) & = default;
457// @endcond
458};
459
460//!
461//! \class IPluginV2Ext
462//!
463//! \brief Plugin class for user-implemented layers.
464//!
465//! Plugins are a mechanism for applications to implement custom layers. This
466//! interface provides additional capabilities to the IPluginV2 interface by
467//! supporting different output data types and broadcast across batches.
468//!
469//! \see IPluginV2
470//!
471//! \deprecated Deprecated in TensorRT 8.5. Implement IPluginV3 instead.
472//!
473class TRT_DEPRECATED IPluginV2Ext : public IPluginV2
474{
475public:
476    //!
477    //! \brief Return the DataType of the plugin output at the requested index.
478    //!
479    //! \param index The output tensor index in the valid range between 0 and getNbOutputs()-1.
480    //! \param inputTypes The data types of the input tensors, stored in an array of length nbInputs.
481    //! \param nbInputs The number of input tensors. Will be a non-negative integer.
482    //!
483    //! \return The data type of the output tensor with the provided index if the input tensors have the data types
484    //! provided in inputTypes, provided the output tensor index is in the valid range. DataType::kFLOAT must be
485    //! returned if the index is not in the valid range.
486    //!
487    //! The default behavior must be to return the type of the first input, or DataType::kFLOAT if the layer has no
488    //! inputs. The returned data type must have a format that is supported by the plugin.
489    //!
490    //! \see supportsFormat()
491    //!
492    //! \warning DataType:kBOOL and DataType::kUINT8 are not supported.
493    //!
494    //! \usage
495    //! - Allowed context for the API call
496    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
497    //!                  when building networks on multiple devices sharing the same plugin.
498    //!
499    virtual nvinfer1::DataType getOutputDataType(
500        int32_t index, nvinfer1::DataType const* inputTypes, int32_t nbInputs) const noexcept
501        = 0;
502
503    //!
504    //! \brief Return true if the output tensor is broadcast across a batch.
505    //!
506    //! \param outputIndex The index of the output tensor, which will be in the valid range between 0 and
507    //! nbOutputs()-1.
508    //! \param inputIsBroadcasted A boolean array of length nbInputs. The i-th element will be true if and only if
509    //! the tensor for the ith input is broadcast across a batch.
510    //! \param nbInputs The number of inputs. Will be a non-negative integer.
511    //!
512    //! The values in inputIsBroadcasted refer to broadcasting at the semantic level,
513    //! i.e. are unaffected by whether method canBroadcastInputAcrossBatch requests
514    //! physical replication of the values.
515    //!
516    //! \usage
517    //! - Allowed context for the API call
518    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
519    //!                  when building networks on multiple devices sharing the same plugin.
520    //!
521    //! \deprecated Deprecated in TensorRT 10.0. Implicit batch support is removed in TensorRT 10.0.
522    //!
523    TRT_DEPRECATED virtual bool isOutputBroadcastAcrossBatch(
524        int32_t outputIndex, bool const* inputIsBroadcasted, int32_t nbInputs) const noexcept
525        = 0;
526
527    //!
528    //! \brief Return true if the plugin can use an input tensor that is broadcast across batch without replication.
529    //!
530    //! \param inputIndex Index of input that could be broadcast. Will be in the valid range between 0 and
531    //! nbInputs - 1 where nbInputs is the maximum number of input tensors supported by this plugin.
532    //!
533    //! \return true if the index is in the valid range and the plugin is able to broadcast a single copy of this
534    //! input tensor across the batch. False otherwise.
535    //!
536    //! For each input whose tensor is semantically broadcast across a batch,
537    //! TensorRT calls this method before calling configurePlugin.
538    //! If canBroadcastInputAcrossBatch returns true, TensorRT will not replicate the input tensor;
539    //! i.e., there will be a single copy that the plugin must share across the batch.
540    //! If it returns false, TensorRT will replicate the input tensor
541    //! so that it appears like a non-broadcasted tensor.
542    //!
543    //! This method is called only for inputs that can be broadcast.
544    //!
545    //! \usage
546    //! - Allowed context for the API call
547    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
548    //!                  when building networks on multiple devices sharing the same plugin.
549    //!
550    //! \deprecated Deprecated in TensorRT 10.0. Implicit batch support is removed in TensorRT 10.0.
551    //!
552    TRT_DEPRECATED virtual bool canBroadcastInputAcrossBatch(int32_t inputIndex) const noexcept = 0;
553
554    //!
555    //! \brief Configure the layer with input and output data types.
556    //!
557    //! This function is called by the builder prior to initialize(). It provides an opportunity for the layer to make
558    //! algorithm choices on the basis of its weights, dimensions, data types and maximum batch size.
559    //!
560    //! \param inputDims The input tensor dimensions. Will be an array of length nbInputs.
561    //! \param nbInputs The number of inputs. Will be a non-negative integer.
562    //! \param outputDims The output tensor dimensions. Will be an array of length nbOutputs.
563    //! \param nbOutputs The number of outputs. Will be a positive integer.
564    //! \param inputTypes The data types selected for the plugin inputs. Will be an array of length nbInputs.
565    //! \param outputTypes The data types selected for the plugin outputs. Will be an array of length nbOutputs.
566    //! \param inputIsBroadcast True for each input that the plugin must broadcast across the batch.
567    //!                         Will be an array of length nbInputs.
568    //! \param outputIsBroadcast True for each output that TensorRT will broadcast across the batch.
569    //!                          Will be an array of length nbOutputs.
570    //! \param floatFormat The format selected for the engine for the floating point inputs/outputs.
571    //! \param maxBatchSize The maximum batch size. Will be a positive integer.
572    //!
573    //! The dimensions passed here do not include the outermost batch size (i.e. for 2D image networks, they will be
574    //! 3-dimensional CHW dimensions). When inputIsBroadcast or outputIsBroadcast is true, the outermost batch size for
575    //! that input or output must be treated as if it is one.
576    //! Index 'i' of inputIsBroadcast is true only if the input is semantically broadcast across the batch and
577    //! calling canBroadcastInputAcrossBatch with argument 'i' returns true.
578    //! Index 'i' of outputIsBroadcast is true only if calling isOutputBroadcastAcrossBatch with argument 'i'
579    //! returns true.
580    //!
581    //! \warning for the floatFormat field, the values PluginFormat::kCHW4, PluginFormat::kCHW16, and
582    //! PluginFormat::kCHW32 will not be passed in, this is to keep backward compatibility with TensorRT 5.x series. Use
583    //! PluginV2IOExt or PluginV2DynamicExt for other PluginFormats.
584    //!
585    //! \usage
586    //! - Allowed context for the API call
587    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
588    //!                  when building networks on multiple devices sharing the same plugin. However, TensorRT
589    //!                  will not call this method from two threads simultaneously on a given clone of a plugin.
590    //!
591    virtual void configurePlugin(Dims const* inputDims, int32_t nbInputs, Dims const* outputDims, int32_t nbOutputs,
592        DataType const* inputTypes, DataType const* outputTypes, bool const* inputIsBroadcast,
593        bool const* outputIsBroadcast, PluginFormat floatFormat, int32_t maxBatchSize) noexcept
594        = 0;
595
596    IPluginV2Ext() = default;
597    ~IPluginV2Ext() override = default;
598
599    //!
600    //! \brief Attach the plugin object to an execution context and grant the plugin the access to some context
601    //! resources.
602    //!
603    //! \param cudnn The cuDNN context handle of the execution context. Will be a valid cuDNN context handle, or
604    //!              nullptr if TacticSource::kCUDNN is disabled.
605    //! \param cublas The cuBLAS context handle of the execution context. Will be a valid cuBLAS context handle, or
606    //!               nullptr if TacticSource::kCUBLAS is disabled.
607    //! \param allocator The allocator used by the execution context
608    //!
609    //! This function is called automatically for each plugin when a new execution context is created. If the context
610    //! was created without resources, this method is not called until the resources are assigned. It is also called if
611    //! new resources are assigned to the context.
612    //!
613    //! If the plugin needs per-context resource, it can be allocated here.
614    //! The plugin can also get context-owned cuDNN and cuBLAS context here.
615    //!
616    //! \note The TacticSource::kCUDNN and TacticSource::kCUBLAS flag is disabled by default.
617    //! The allocator pointer is unique to each building or execution context instance having overlapping lifetimes.
618    //! It can be used as a key to manage resources across plugin instances sharing the same context.
619    //! Plugins attached to different contexts will have different handles as their execution will not overlap.
620    //!
621    //! \see TacticSources
622    //! \see getPluginCudnnHandle(void* executionContextIdentifier)
623    //! \see getPluginCublasHandle(void* excecutionContextIdentifier)
624    //!
625    //! \note In the automotive safety context, the cuDNN and cuBLAS parameters will be nullptr because cuDNN and cuBLAS
626    //!       are not used by the safe runtime.
627    //!
628    //! \usage
629    //! - Allowed context for the API call
630    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
631    //!                  when building networks on multiple devices sharing the same plugin.
632    //!
633    virtual void attachToContext(
634        cudnnContext* /*cudnn*/, cublasContext* /*cublas*/, IGpuAllocator* /*allocator*/) noexcept
635    {
636    }
637
638    //!
639    //! \brief Detach the plugin object from its execution context.
640    //!
641    //! This function is called automatically for each plugin when an execution context is destroyed or the context
642    //! resources are unassigned from the context.
643    //!
644    //! If the plugin owns per-context resource, it can be released here.
645    //!
646    //! \usage
647    //! - Allowed context for the API call
648    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
649    //!                  when building networks on multiple devices sharing the same plugin.
650    //!
651    virtual void detachFromContext() noexcept {}
652
653    //!
654    //! \brief Clone the plugin object. This copies over internal plugin parameters as well and returns a new plugin
655    //! object with these parameters. If the source plugin is pre-configured with configurePlugin(), the returned object
656    //! must also be pre-configured. The returned object must allow attachToContext() with a new execution context.
657    //! Cloned plugin objects can share the same per-engine immutable resource (e.g. weights) with the source object
658    //! (e.g. via ref-counting) to avoid duplication.
659    //!
660    //! \return A pointer to a cloned plugin object if cloning was successful, otherwise nullptr.
661    //!
662    //! \usage
663    //! - Allowed context for the API call
664    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
665    //!                  when building networks on multiple devices sharing the same plugin.
666    //!
667    IPluginV2Ext* clone() const noexcept override = 0;
668
669protected:
670    // @cond SuppressDoxyWarnings
671    IPluginV2Ext(IPluginV2Ext const&) = default;
672    IPluginV2Ext(IPluginV2Ext&&) = default;
673    IPluginV2Ext& operator=(IPluginV2Ext const&) & = default;
674    IPluginV2Ext& operator=(IPluginV2Ext&&) & = default;
675// @endcond
676
677    //!
678    //! \brief Return the API version with which this plugin was built. The
679    //!  upper byte reserved by TensorRT and is used to differentiate this from IPluginV2.
680    //!
681    //! \return In the lower three bytes, the TensorRT version in the format
682    //!         (major * 100 + minor) * 100 + patch.
683    //!         In the upper byte, the value 1.
684    //!
685    //! Do not override this method as it is used by the TensorRT library to maintain backwards-compatibility with
686    //! plugins.
687    //!
688    //! \usage
689    //! - Allowed context for the API call
690    //!   - Thread-safe: Yes, the implementation provided here is safe to call from any thread.
691    //!
692    int32_t getTensorRTVersion() const noexcept override
693    {
694        return static_cast<int32_t>((static_cast<uint32_t>(PluginVersion::kV2_EXT) << 24U)
695            | (static_cast<uint32_t>(NV_TENSORRT_VERSION) & 0xFFFFFFU));
696    }
697
698    //!
699    //! \brief Derived classes must not implement this. In a C++11 API it would be final.
700    //!
701    //! IPluginV2Ext::configureWithFormat() is a NOP operation for all classes derived from IPluginV2Ext.
702    //! These classes call configurePlugin() instead.
703    //!
704    void configureWithFormat(Dims const* /*inputDims*/, int32_t /*nbInputs*/, Dims const* /*outputDims*/,
705        int32_t /*nbOutputs*/, DataType /*type*/, PluginFormat /*format*/, int32_t /*maxBatchSize*/) noexcept override
706    {
707    }
708};
709
710//!
711//! \class IPluginV2IOExt
712//!
713//! \brief Plugin class for user-implemented layers.
714//!
715//! Plugins are a mechanism for applications to implement custom layers. This interface provides additional
716//! capabilities to the IPluginV2Ext interface by extending different I/O data types and tensor formats.
717//!
718//! \see IPluginV2Ext
719//!
720//! \deprecated Deprecated in TensorRT 10.0. Implement IPluginV3 instead.
721//!
722class TRT_DEPRECATED IPluginV2IOExt : public IPluginV2Ext
723{
724public:
725    //!
726    //! \brief Configure the layer.
727    //!
728    //! This function is called by the builder prior to initialize(). It provides an opportunity for the layer to make
729    //! algorithm choices on the basis of the provided I/O PluginTensorDesc.
730    //!
731    //! \param in The input tensors attributes that are used for configuration.
732    //! \param nbInput Number of input tensors.
733    //! \param out The output tensors attributes that are used for configuration.
734    //! \param nbOutput Number of output tensors.
735    //!
736    //! \usage
737    //! - Allowed context for the API call
738    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
739    //!                  when building networks on multiple devices sharing the same plugin. However, TensorRT
740    //!                  will not call this method from two threads simultaneously on a given clone of a plugin.
741    //!
742    virtual void configurePlugin(
743        PluginTensorDesc const* in, int32_t nbInput, PluginTensorDesc const* out, int32_t nbOutput) noexcept
744        = 0;
745
746    //!
747    //! \brief Return true if plugin supports the format and datatype for the input/output indexed by pos.
748    //!
749    //! For this method inputs are numbered 0..(nbInputs-1) and outputs are numbered nbInputs..(nbInputs+nbOutputs-1).
750    //! Using this numbering, pos is an index into InOut, where 0 <= pos < nbInputs+nbOutputs.
751    //!
752    //! TensorRT invokes this method to ask if the input/output indexed by pos supports the format/datatype specified
753    //! by inOut[pos].format and inOut[pos].type. The override must return true if that format/datatype at inOut[pos]
754    //! are supported by the plugin. If support is conditional on other input/output formats/datatypes, the plugin can
755    //! make its result conditional on the formats/datatypes in inOut[0..pos-1], which will be set to values
756    //! that the plugin supports. The override must not inspect inOut[pos+1..nbInputs+nbOutputs-1],
757    //! which will have invalid values.  In other words, the decision for pos must be based on inOut[0..pos] only.
758    //!
759    //! Some examples:
760    //!
761    //! * A definition for a plugin that supports only FP16 NCHW:
762    //!
763    //!         return inOut.format[pos] == TensorFormat::kLINEAR && inOut.type[pos] == DataType::kHALF;
764    //!
765    //! * A definition for a plugin that supports only FP16 NCHW for its two inputs,
766    //!   and FP32 NCHW for its single output:
767    //!
768    //!         return inOut.format[pos] == TensorFormat::kLINEAR &&
769    //!                (inOut.type[pos] == (pos < 2 ?  DataType::kHALF : DataType::kFLOAT));
770    //!
771    //! * A definition for a "polymorphic" plugin with two inputs and one output that supports
772    //!   any format or type, but the inputs and output must have the same format and type:
773    //!
774    //!         return pos == 0 || (inOut.format[pos] == inOut.format[0] && inOut.type[pos] == inOut.type[0]);
775    //!
776    //! Warning: TensorRT will stop asking for formats once it finds kFORMAT_COMBINATION_LIMIT on combinations.
777    //!
778    //! \usage
779    //! - Allowed context for the API call
780    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
781    //!                  when building networks on multiple devices sharing the same plugin.
782    //!
783    virtual bool supportsFormatCombination(
784        int32_t pos, PluginTensorDesc const* inOut, int32_t nbInputs, int32_t nbOutputs) const noexcept
785        = 0;
786
787    // @cond SuppressDoxyWarnings
788    IPluginV2IOExt() = default;
789    ~IPluginV2IOExt() override = default;
790// @endcond
791
792protected:
793// @cond SuppressDoxyWarnings
794    IPluginV2IOExt(IPluginV2IOExt const&) = default;
795    IPluginV2IOExt(IPluginV2IOExt&&) = default;
796    IPluginV2IOExt& operator=(IPluginV2IOExt const&) & = default;
797    IPluginV2IOExt& operator=(IPluginV2IOExt&&) & = default;
798// @endcond
799
800    //!
801    //! \brief Return the API version with which this plugin was built. The upper byte is reserved by TensorRT and is
802    //! used to differentiate this from IPluginV2 and IPluginV2Ext.
803    //!
804    //! Do not override this method as it is used by the TensorRT library to maintain backwards-compatibility with
805    //! plugins.
806    //!
807    //! \usage
808    //! - Allowed context for the API call
809    //!   - Thread-safe: Yes, the implementation provided here is safe to call from any thread.
810    //!
811    int32_t getTensorRTVersion() const noexcept override
812    {
813        return static_cast<int32_t>((static_cast<uint32_t>(PluginVersion::kV2_IOEXT) << 24U)
814            | (static_cast<uint32_t>(NV_TENSORRT_VERSION) & 0xFFFFFFU));
815    }
816
817private:
818    // Following are obsolete base class methods, and must not be implemented or used.
819
820    //!
821    //! \brief Set plugin configuration.
822    //!
823    void configurePlugin(Dims const*, int32_t, Dims const*, int32_t, DataType const*, DataType const*, bool const*,
824        bool const*, PluginFormat, int32_t) noexcept final
825    {
826    }
827
828    //!
829    //! \brief Check if provided data type is supported.
830    //!
831    bool supportsFormat(DataType, PluginFormat) const noexcept final
832    {
833        return false;
834    }
835};
836
837namespace v_1_0
838{
839class TRT_DEPRECATED IPluginCreator : public IPluginCreatorInterface
840{
841public:
842    //!
843    //! \brief Return the plugin name.
844    //!
845    //! \warning The string returned must be NULL-terminated and have a length of 1024 bytes or less including
846    //! the NULL terminator.
847    //!
848    //! \usage
849    //! - Allowed context for the API call
850    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
851    //!                  when building networks on multiple devices sharing the same plugin or when deserializing
852    //!                  multiple engines concurrently sharing plugins.
853    //!
854    virtual AsciiChar const* getPluginName() const noexcept = 0;
855
856    //!
857    //! \brief Return the plugin version.
858    //!
859    //! \warning The string returned must be NULL-terminated and have a length of 1024 bytes or less including
860    //! the NULL terminator.
861    //!
862    //! \usage
863    //! - Allowed context for the API call
864    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
865    //!                  when building networks on multiple devices sharing the same plugin or when deserializing
866    //!                  multiple engines concurrently sharing plugins.
867    //!
868    virtual AsciiChar const* getPluginVersion() const noexcept = 0;
869
870    //!
871    //! \brief Return a list of fields that need to be passed to createPlugin.
872    //!
873    //! \see PluginFieldCollection
874    //!
875    //! \usage
876    //! - Allowed context for the API call
877    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
878    //!                  when building networks on multiple devices sharing the same plugin or when deserializing
879    //!                  multiple engines concurrently sharing plugins.
880    //!
881    virtual PluginFieldCollection const* getFieldNames() noexcept = 0;
882
883    //!
884    //! \brief Return a plugin object. Return nullptr in case of error.
885    //!
886    //! \param name A NULL-terminated name string of length 1024 or less, including the NULL terminator.
887    //! \param fc A pointer to a collection of fields needed for constructing the plugin.
888    //!
889    //! \usage
890    //! - Allowed context for the API call
891    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
892    //!                  when building networks on multiple devices sharing the same plugin or when deserializing
893    //!                  multiple engines concurrently sharing plugins.
894    //!
895    virtual IPluginV2* createPlugin(AsciiChar const* name, PluginFieldCollection const* fc) noexcept = 0;
896
897    //!
898    //! \brief Called during deserialization of plugin layer. Return a plugin object.
899    //!
900    //! \param name A NULL-terminated name string of length 1024 or less, including the NULL terminator.
901    //! \param serialData The start address of a byte array with the serialized plugin representation.
902    //! \param serialLength The length in bytes of the byte array with the serialized plugin representation.
903    //!
904    //! \return A deserialized plugin object
905    //!
906    //! \usage
907    //! - Allowed context for the API call
908    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
909    //!                  when building networks on multiple devices sharing the same plugin or when deserializing
910    //!                  multiple engines concurrently sharing plugins.
911    //!
912    virtual IPluginV2* deserializePlugin(AsciiChar const* name, void const* serialData, size_t serialLength) noexcept
913        = 0;
914
915    //!
916    //! \brief Set the namespace of the plugin creator based on the plugin
917    //! library it belongs to. This can be set while registering the plugin creator.
918    //!
919    //! \param pluginNamespace A NULL-terminated namespace string of length 1024 or less, including the NULL terminator
920    //!
921    //! \see IPluginRegistry::registerCreator()
922    //!
923    //! \usage
924    //! - Allowed context for the API call
925    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
926    //!                  when building networks on multiple devices sharing the same plugin or when deserializing
927    //!                  multiple engines concurrently sharing plugins.
928    //!
929    virtual void setPluginNamespace(AsciiChar const* pluginNamespace) noexcept = 0;
930
931    //!
932    //! \brief Return the namespace of the plugin creator object.
933    //!
934    //! \warning The string returned must be NULL-terminated and have a length of 1024 bytes or less including the
935    //! NULL terminator.
936    //!
937    //! \usage
938    //! - Allowed context for the API call
939    //!   - Thread-safe: Yes, this method is required to be thread-safe and may be called from multiple threads
940    //!                  when building networks on multiple devices sharing the same plugin or when deserializing
941    //!                  multiple engines concurrently sharing plugins.
942    //!
943    virtual AsciiChar const* getPluginNamespace() const noexcept = 0;
944
945    IPluginCreator() = default;
946    ~IPluginCreator() override = default;
947
948protected:
949    // @cond SuppressDoxyWarnings
950    IPluginCreator(IPluginCreator const&) = default;
951    IPluginCreator(IPluginCreator&&) = default;
952    IPluginCreator& operator=(IPluginCreator const&) & = default;
953    IPluginCreator& operator=(IPluginCreator&&) & = default;
954    // @endcond
955public:
956    //!
957    //! \brief Return version information associated with this interface. Applications must not override this method.
958    //!
959    InterfaceInfo getInterfaceInfo() const noexcept override
960    {
961        return InterfaceInfo{"PLUGIN CREATOR_V1", 1, 0};
962    }
963};
964} // namespace v_1_0
965
966//!
967//! \class IPluginCreator
968//!
969//! \brief Plugin creator class for user implemented layers.
970//!
971//! \see IPlugin and IPluginFactory
972//!
973//! \deprecated Deprecated in TensorRT 10.0. Please implement IPluginCreatorV3One
974//! along with IPluginV3 plugins instead.
975//!
976using IPluginCreator = v_1_0::IPluginCreator;
977
978} // namespace nvinfer1
979
980#endif // NV_INFER_RUNTIME_PLUGIN_H
981 
codekingpro/portable-devtools · Team Ai