Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
NvInferPythonPlugin.h595 linesDownload Raw Back to impl
1/*
2 * SPDX-FileCopyrightText: Copyright (c) 2024-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 TRT_PYTHON_IMPL_PLUGIN_H
19#define TRT_PYTHON_IMPL_PLUGIN_H
20
21#include "NvInfer.h"
22
23//!
24//! \file NvInferPythonPlugin.h
25//!
26//! This file contains definitions for supporting the `tensorrt.plugin` Python module
27//!
28//! \warning None of the defintions here are part of the TensorRT C++ API and may not follow semantic versioning rules.
29//! TensorRT clients must not utilize them directly.
30//!
31
32namespace nvinfer1
33{
34
35//! \enum PluginArgType
36//! \brief Numeric type of an extra kernel input argument in an AOT Python plugin
37enum class PluginArgType : int32_t
38{
39    //! Integer argument
40    kINT = 0,
41};
42
43//! \enum PluginArgDataType
44//! \brief Data type of an extra kernel input argument in an AOT Python plugin
45enum class PluginArgDataType : int32_t
46{
47    //! 8-bit signed integer
48    kINT8 = 0,
49    //! 16-bit signed integer
50    kINT16 = 1,
51    //! 32-bit signed integer
52    kINT32 = 2,
53};
54//! \class ISymExpr
55//! \brief Generic interface for a scalar symbolic expression implementable by a Python plugin / TensorRT Python backend
56class ISymExpr
57{
58public:
59    //! \brief Get the type of the symbolic expression
60    virtual PluginArgType getType() const noexcept = 0;
61    //! \brief Get the data type of the symbolic expression
62    virtual PluginArgDataType getDataType() const noexcept = 0;
63    //! \brief Underlying symbolic expression
64    virtual void* getExpr() noexcept = 0;
65};
66
67//! Impl class for ISymExprs
68class ISymExprsImpl
69{
70public:
71    virtual ISymExpr* getSymExpr(int32_t index) const noexcept = 0;
72    virtual bool setSymExpr(int32_t index, ISymExpr* symExpr) noexcept = 0;
73    virtual int32_t getNbSymExprs() const noexcept = 0;
74    virtual bool setNbSymExprs(int32_t count) noexcept = 0;
75
76    virtual ~ISymExprsImpl() noexcept = default;
77};
78
79//! \class ISymExprs
80//! \brief Allows for a sequence of symbolic expressions to be communicated to the TensorRT backend
81//! \note Clients must not implement this class.
82//! \see ISymExpr
83class ISymExprs
84{
85public:
86    //! \brief Get the symbolic expression at the given index
87    //! \return A pointer to the symbolic expression or nullptr if the index is out of range
88    ISymExpr* getSymExpr(int32_t index) const noexcept
89    {
90        return mImpl->getSymExpr(index);
91    }
92
93    //! \brief Set the symbolic expression at the given index
94    //! \return true if the index is in range and the symbolic expression was set successfully, false otherwise
95    bool setSymExpr(int32_t index, ISymExpr* symExpr) noexcept
96    {
97        return mImpl->setSymExpr(index, symExpr);
98    }
99
100    //! \brief Get the number of symbolic expressions
101    int32_t getNbSymExprs() const noexcept
102    {
103        return mImpl->getNbSymExprs();
104    }
105
106    //! \brief Set the number of symbolic expressions
107    //! \return true if the number of symbolic expressions was set successfully, false otherwise
108    bool setNbSymExprs(int32_t count) noexcept
109    {
110        return mImpl->setNbSymExprs(count);
111    }
112
113protected:
114    ISymExprsImpl* mImpl{nullptr};
115    virtual ~ISymExprs() noexcept = default;
116};
117
118//! \enum QuickPluginCreationRequest
119//! \brief Communicates preference when a quickly deployable plugin is to be added to the network
120enum class QuickPluginCreationRequest : int32_t
121{
122    //! No preference specified
123    kUNKNOWN = 0,
124    //! JIT plugin is preferred
125    kPREFER_JIT = 1,
126    //! AOT plugin is preferred
127    kPREFER_AOT = 2,
128    //! JIT plugin must be used. TensorRT should fail if a JIT implementation cannot be found.
129    kSTRICT_JIT = 3,
130    //! AOT plugin must be used. TensorRT should fail if an AOT implementation cannot be found.
131    kSTRICT_AOT = 4,
132};
133
134//! Impl class for IKernelLaunchParams
135class IKernelLaunchParamsImpl
136{
137public:
138    virtual ISymExpr* getGridX() noexcept = 0;
139    virtual bool setGridX(ISymExpr* gridX) noexcept = 0;
140
141    virtual ISymExpr* getGridY() noexcept = 0;
142    virtual bool setGridY(ISymExpr* gridY) noexcept = 0;
143
144    virtual ISymExpr* getGridZ() noexcept = 0;
145    virtual bool setGridZ(ISymExpr* gridZ) noexcept = 0;
146
147    virtual ISymExpr* getBlockX() noexcept = 0;
148    virtual bool setBlockX(ISymExpr* blockX) noexcept = 0;
149
150    virtual ISymExpr* getBlockY() noexcept = 0;
151    virtual bool setBlockY(ISymExpr* blockY) noexcept = 0;
152
153    virtual ISymExpr* getBlockZ() noexcept = 0;
154    virtual bool setBlockZ(ISymExpr* blockZ) noexcept = 0;
155
156    virtual ISymExpr* getSharedMem() noexcept = 0;
157    virtual bool setSharedMem(ISymExpr* sharedMem) noexcept = 0;
158
159    virtual ~IKernelLaunchParamsImpl() noexcept = default;
160};
161
162//! \class IKernelLaunchParams
163//! \brief Allows for kernel launch parameters to be communicated to the TensorRT backend
164//! \note Clients must not implement this class.
165class IKernelLaunchParams
166{
167public:
168    //! Get the X dimension of the grid
169    ISymExpr* getGridX() noexcept
170    {
171        return mImpl->getGridX();
172    }
173
174    //! \brief Set the X dimension of the grid
175    //! \return true if the grid's X dimension was set successfully, false otherwise
176    bool setGridX(ISymExpr* gridX) noexcept
177    {
178        return mImpl->setGridX(gridX);
179    }
180
181    //! Get the Y dimension of the grid
182    ISymExpr* getGridY() noexcept
183    {
184        return mImpl->getGridY();
185    }
186
187    //! \brief Set the Y dimension of the grid
188    //! \return true if the grid's Y dimension was set successfully, false otherwise
189    bool setGridY(ISymExpr* gridY) noexcept
190    {
191        return mImpl->setGridY(gridY);
192    }
193
194    //! Get the Z dimension of the grid
195    ISymExpr* getGridZ() noexcept
196    {
197        return mImpl->getGridZ();
198    }
199
200    //! \brief Set the Z dimension of the grid
201    //! \return true if the grid's Z dimension was set successfully, false otherwise
202    bool setGridZ(ISymExpr* gridZ) noexcept
203    {
204        return mImpl->setGridZ(gridZ);
205    }
206
207    //! \brief Get the X dimension of each thread block
208    ISymExpr* getBlockX() noexcept
209    {
210        return mImpl->getBlockX();
211    }
212
213    //! \brief Set the X dimension of each thread block
214    //! \return true if each thread block's X dimension was set successfully, false otherwise
215    bool setBlockX(ISymExpr* blockX) noexcept
216    {
217        return mImpl->setBlockX(blockX);
218    }
219
220    //! \brief Get the Y dimension of each thread block
221    ISymExpr* getBlockY() noexcept
222    {
223        return mImpl->getBlockY();
224    }
225
226    //! \brief Set the Y dimension of each thread block
227    //! \return true if each thread block's Y dimension was set successfully, false otherwise
228    bool setBlockY(ISymExpr* blockY) noexcept
229    {
230        return mImpl->setBlockY(blockY);
231    }
232
233    //! \brief Get the Z dimension of each thread block
234    ISymExpr* getBlockZ() noexcept
235    {
236        return mImpl->getBlockZ();
237    }
238
239    //! \brief Set the Z dimension of each thread block
240    //! \return true if each thread block's Z dimension was set successfully, false otherwise
241    bool setBlockZ(ISymExpr* blockZ) noexcept
242    {
243        return mImpl->setBlockZ(blockZ);
244    }
245
246    //! \brief Get the dynamic shared-memory per thread block in bytes
247    ISymExpr* getSharedMem() noexcept
248    {
249        return mImpl->getSharedMem();
250    }
251
252    //! \brief Set the dynamic shared-memory per thread block in bytes
253    //! \return true if the dynamic shared-memory per thread block was set successfully, false otherwise
254    bool setSharedMem(ISymExpr* sharedMem) noexcept
255    {
256        return mImpl->setSharedMem(sharedMem);
257    }
258
259protected:
260    IKernelLaunchParamsImpl* mImpl{nullptr};
261    virtual ~IKernelLaunchParams() noexcept = default;
262};
263
264namespace v_1_0
265{
266
267class IPluginV3QuickCore : public IPluginCapability
268{
269public:
270    InterfaceInfo getInterfaceInfo() const noexcept override
271    {
272        return InterfaceInfo{"PLUGIN_V3QUICK_CORE", 1, 0};
273    }
274
275    virtual AsciiChar const* getPluginName() const noexcept = 0;
276
277    virtual AsciiChar const* getPluginVersion() const noexcept = 0;
278
279    virtual AsciiChar const* getPluginNamespace() const noexcept = 0;
280};
281
282class IPluginV3QuickBuild : public IPluginCapability
283{
284public:
285    InterfaceInfo getInterfaceInfo() const noexcept override
286    {
287        return InterfaceInfo{"PLUGIN_V3QUICK_BUILD", 1, 0};
288    }
289
290    //!
291    //! \brief Provide the data types of the plugin outputs if the input tensors have the data types provided.
292    //!
293    //! \param outputTypes Pre-allocated array to which the output data types should be written.
294    //! \param nbOutputs The number of output tensors. This matches the value returned from getNbOutputs().
295    //! \param inputTypes The input data types.
296    //! \param inputRanks Ranks of the input tensors
297    //! \param nbInputs The number of input tensors.
298    //!
299    //! \return 0 for success, else non-zero
300    //!
301    virtual int32_t getOutputDataTypes(DataType* outputTypes, int32_t nbOutputs, DataType const* inputTypes,
302        int32_t const* inputRanks, int32_t nbInputs) const noexcept = 0;
303
304    //!
305    //! \brief Provide expressions for computing dimensions of the output tensors from dimensions of the input tensors.
306    //!
307    //! \param inputs Expressions for dimensions of the input tensors
308    //! \param nbInputs The number of input tensors
309    //! \param shapeInputs Expressions for values of the shape tensor inputs
310    //! \param nbShapeInputs The number of shape tensor inputs
311    //! \param outputs Pre-allocated array to which the output dimensions must be written
312    //! \param exprBuilder Object for generating new dimension expressions
313    //!
314    //! \return 0 for success, else non-zero
315    //!
316    virtual int32_t getOutputShapes(DimsExprs const* inputs, int32_t nbInputs, DimsExprs const* shapeInputs,
317        int32_t nbShapeInputs, DimsExprs* outputs, int32_t nbOutputs, IExprBuilder& exprBuilder) noexcept = 0;
318
319    //!
320    //! \brief Configure the plugin. Behaves similarly to `IPluginV3OneBuild::configurePlugin()`
321    //!
322    //! \return 0 for success, else non-zero
323    //!
324    virtual int32_t configurePlugin(DynamicPluginTensorDesc const* in, int32_t nbInputs,
325        DynamicPluginTensorDesc const* out, int32_t nbOutputs) noexcept = 0;
326
327    //!
328    //! \brief Get number of format combinations supported by the plugin for the I/O characteristics indicated by
329    //! `inOut`.
330    //!
331    virtual int32_t getNbSupportedFormatCombinations(
332        DynamicPluginTensorDesc const* inOut, int32_t nbInputs, int32_t nbOutputs) noexcept = 0;
333
334    //!
335    //! \brief Write all format combinations supported by the plugin for the I/O characteristics indicated by `inOut` to
336    //! `supportedCombinations`. It is guaranteed to have sufficient memory allocated for (nbInputs + nbOutputs) *
337    //! getNbSupportedFormatCombinations() `PluginTensorDesc`s.
338    //!
339    //! \return 0 for success, else non-zero
340    //!
341    virtual int32_t getSupportedFormatCombinations(DynamicPluginTensorDesc const* inOut, int32_t nbInputs,
342        int32_t nbOutputs, PluginTensorDesc* supportedCombinations, int32_t nbFormatCombinations) noexcept = 0;
343
344    //!
345    //! \brief Get the number of outputs from the plugin.
346    //!
347    virtual int32_t getNbOutputs() const noexcept = 0;
348
349    //!
350    //! \brief Communicates to TensorRT that the output at the specified output index is aliased to the input at the
351    //! returned index. Behaves similary to `v_2_0::IPluginV3OneBuild.getAliasedInput()`.
352    //!
353    virtual int32_t getAliasedInput(int32_t outputIndex) noexcept
354    {
355        return -1;
356    }
357
358    //!
359    //! \brief Query for any custom tactics that the plugin intends to use specific to the I/O characteristics indicated
360    //! by the immediately preceding call to `configurePlugin()`.
361    //!
362    //! \return 0 for success, else non-zero
363    //!
364    virtual int32_t getValidTactics(int32_t* tactics, int32_t nbTactics) noexcept
365    {
366        return 0;
367    }
368
369    //!
370    //! \brief Query for number of custom tactics related to the `getValidTactics()` call.
371    //!
372    virtual int32_t getNbTactics() noexcept
373    {
374        return 0;
375    }
376
377    //!
378    //! \brief Called to query the suffix to use for the timing cache ID. May be called anytime after plugin creation.
379    //!
380    virtual char const* getTimingCacheID() noexcept
381    {
382        return nullptr;
383    }
384
385    //!
386    //! \brief Query for a string representing the configuration of the plugin. May be called anytime after
387    //! plugin creation.
388    //!
389    virtual char const* getMetadataString() noexcept
390    {
391        return nullptr;
392    }
393};
394
395class IPluginV3QuickAOTBuild : public IPluginV3QuickBuild
396{
397public:
398    InterfaceInfo getInterfaceInfo() const noexcept override
399    {
400        return InterfaceInfo{"PLUGIN_V3QUICKAOT_BUILD", 1, 0};
401    }
402
403    //! \brief Get the launch parameters for the kernel to be used for the specified input and output types/formats and
404    //! any corresponding custom tactics.
405    //!        If custom tactics are being advertised by the plugin, the corresponding tactic is the one specified by
406    //!        the immediately preceding call to setTactic().
407    //!
408    //! \param inputs Expressions for dimensions of the input tensors
409    //! \param inOut The input and output tensors' attributes
410    //! \param nbInputs The number of input tensors
411    //! \param nbOutputs The number of output tensors
412    //! \param launchParams Interface which allows the specification of kernel launch parameters as symbolic expressions
413    //! of the input dimensions
414    //! \param extraArgs Interface which allows the specification of any scalar arguments to be
415    //! passed to the kernel, as symbolic expressions of the input dimensions
416    //! \param exprBuilder Object for generating new symbolic expressions
417    //!
418    //! \return 0 for success, else non-zero
419    //!
420    virtual int32_t getLaunchParams(DimsExprs const* inputs, DynamicPluginTensorDesc const* inOut, int32_t nbInputs,
421        int32_t nbOutputs, IKernelLaunchParams* launchParams, ISymExprs* extraArgs,
422        IExprBuilder& exprBuilder) noexcept = 0;
423
424    //!
425    //! \brief Get the compiled form for the kernel to be used for the specified input and output types/formats and any
426    //! corresponding custom tactics.
427    //!        If custom tactics are being advertised by the plugin, the corresponding tactic is the one specified by
428    //!        the immediately preceding call to setTactic().
429    //!
430    //! \param in The input tensors' attributes that are used for configuration.
431    //! \param nbInputs Number of input tensors.
432    //! \param out The output tensors' attributes that are used for configuration.
433    //! \param nbOutputs Number of output tensors.
434    //! \param kernelName The name for the kernel.
435    //! \param compiledKernel Compiled form of the kernel.
436    //! \param compiledKernelSize The size of the compiled kernel.
437    //!
438    //! \return 0 for success, else non-zero
439    //!
440    virtual int32_t getKernel(PluginTensorDesc const* in, int32_t nbInputs, PluginTensorDesc const* out,
441        int32_t nbOutputs, const char** kernelName, char** compiledKernel, int32_t* compiledKernelSize) noexcept = 0;
442
443    //!
444    //! \brief Set the tactic to be used in the subsequent call to enqueue(). Behaves similar to
445    //! IPluginV3OneRuntime::setTactic()
446    //!
447    //! \return 0 for success, else non-zero
448    //!
449    virtual int32_t setTactic(int32_t tactic) noexcept
450    {
451        return 0;
452    }
453};
454
455class IPluginV3QuickRuntime : public IPluginCapability
456{
457public:
458    InterfaceInfo getInterfaceInfo() const noexcept override
459    {
460        return InterfaceInfo{"PLUGIN_V3QUICK_RUNTIME", 1, 0};
461    }
462
463    //!
464    //! \brief Set the tactic to be used in the subsequent call to enqueue(). Behaves similar to
465    //! `IPluginV3OneRuntime::setTactic()`.
466    //!
467    //! \return 0 for success, else non-zero
468    //!
469    virtual int32_t setTactic(int32_t tactic) noexcept
470    {
471        return 0;
472    }
473
474    //!
475    //! \brief Execute the plugin.
476    //!
477    //! \param inputDesc how to interpret the memory for the input tensors.
478    //! \param outputDesc how to interpret the memory for the output tensors.
479    //! \param inputs The memory for the input tensors.
480    //! \param inputStrides Strides for input tensors.
481    //! \param outputStrides Strides for output tensors.
482    //! \param outputs The memory for the output tensors.
483    //! \param nbInputs Number of input tensors.
484    //! \param nbOutputs Number of output tensors.
485    //! \param stream The stream in which to execute the kernels.
486    //!
487    //! \return 0 for success, else non-zero
488    //!
489    virtual int32_t enqueue(PluginTensorDesc const* inputDesc, PluginTensorDesc const* outputDesc,
490        void const* const* inputs, void* const* outputs, Dims const* inputStrides, Dims const* outputStrides,
491        int32_t nbInputs, int32_t nbOutputs, cudaStream_t stream) noexcept = 0;
492
493    //!
494    //! \brief Get the plugin fields which should be serialized.
495    //!
496    virtual PluginFieldCollection const* getFieldsToSerialize() noexcept = 0;
497};
498
499class IPluginCreatorV3Quick : public IPluginCreatorInterface
500{
501public:
502    InterfaceInfo getInterfaceInfo() const noexcept override
503    {
504        return InterfaceInfo{"PLUGIN CREATOR_V3QUICK", 1, 0};
505    }
506
507    //!
508    //! \brief Return a plugin object. Return nullptr in case of error.
509    //!
510    //! \param name A NULL-terminated name string of length 1024 or less, including the NULL terminator.
511    //! \param namespace A NULL-terminated name string of length 1024 or less, including the NULL terminator.
512    //! \param fc A pointer to a collection of fields needed for constructing the plugin.
513    //! \param phase The TensorRT phase in which the plugin is being created
514    //! \param quickPluginCreationRequest Whether a JIT or AOT plugin should be created
515    //!
516    virtual IPluginV3* createPlugin(AsciiChar const* name, AsciiChar const* nspace, PluginFieldCollection const* fc,
517        TensorRTPhase phase, QuickPluginCreationRequest quickPluginCreationRequest) noexcept = 0;
518
519    //!
520    //! \brief Return a list of fields that need to be passed to createPlugin() when creating a plugin for use in the
521    //! TensorRT build phase.
522    //!
523    virtual PluginFieldCollection const* getFieldNames() noexcept = 0;
524
525    virtual AsciiChar const* getPluginName() const noexcept = 0;
526
527    virtual AsciiChar const* getPluginVersion() const noexcept = 0;
528
529    virtual AsciiChar const* getPluginNamespace() const noexcept = 0;
530
531    IPluginCreatorV3Quick() = default;
532    virtual ~IPluginCreatorV3Quick() = default;
533
534protected:
535    IPluginCreatorV3Quick(IPluginCreatorV3Quick const&) = default;
536    IPluginCreatorV3Quick(IPluginCreatorV3Quick&&) = default;
537    IPluginCreatorV3Quick& operator=(IPluginCreatorV3Quick const&) & = default;
538    IPluginCreatorV3Quick& operator=(IPluginCreatorV3Quick&&) & = default;
539};
540
541} // namespace v_1_0
542
543//!
544//! \class IPluginV3QuickCore
545//!
546//! \brief Provides core capability (`IPluginCapability::kCORE`) for quickly-deployable TRT plugins
547//!
548//! \warning This class is strictly for the purpose of supporting quickly-deployable TRT Python plugins and is not part
549//! of the public TensorRT C++ API. Users must not inherit from this class.
550//!
551using IPluginV3QuickCore = v_1_0::IPluginV3QuickCore;
552
553//!
554//! \class IPluginV3QuickBuild
555//!
556//! \brief Provides build capability (`IPluginCapability::kBUILD`) for quickly-deployable TRT plugins
557//!
558//! \warning This class is strictly for the purpose of supporting quickly-deployable TRT Python plugins and is not part
559//! of the public TensorRT C++ API. Users must not inherit from this class.
560//!
561using IPluginV3QuickBuild = v_1_0::IPluginV3QuickBuild;
562
563//!
564//! \class IPluginV3QuickAOTBuild
565//!
566//! \brief Provides additional build capabilities for AOT quickly-deployable TRT plugins. Descends from
567//! IPluginV3QuickBuild.
568//!
569//! \warning This class is strictly for the purpose of supporting quickly-deployable TRT Python plugins and is not part
570//! of the public TensorRT C++ API. Users must not inherit from this class.
571//!
572using IPluginV3QuickAOTBuild = v_1_0::IPluginV3QuickAOTBuild;
573
574//!
575//! \class IPluginV3QuickRuntime
576//!
577//! \brief Provides runtime capability (`IPluginCapability::kRUNTIME`) for JIT quickly-deployable TRT plugins
578//!
579//! \warning This class is strictly for the purpose of supporting quickly-deployable TRT Python plugins and is not part
580//! of the public TensorRT C++ API. Users must not inherit from this class.
581//!
582using IPluginV3QuickRuntime = v_1_0::IPluginV3QuickRuntime;
583
584//!
585//! \class IPluginCreatorV3Quick
586//!
587//! \warning This class is strictly for the purpose of supporting quickly-deployable TRT Python plugins and is not part
588//! of the public TensorRT C++ API. Users must not inherit from this class.
589//!
590using IPluginCreatorV3Quick = v_1_0::IPluginCreatorV3Quick;
591
592} // namespace nvinfer1
593
594#endif // TRT_PYTHON_IMPL_PLUGIN_H
595 
codekingpro/portable-devtools · Team Ai