
NVIDIA TensorRT ScatterElements 插件深度解析从 ONNX 算子到 GPU 原子操作【免费下载链接】TensorRTNVIDIA® TensorRT™ is an SDK for high-performance deep learning inference on NVIDIA GPUs. This repository contains the open source components of TensorRT.项目地址: https://gitcode.com/GitHub_Trending/tens/TensorRT本指南围绕 TensorRT 开源仓库中的 ScatterElements 插件plugin/scatterElementsPlugin/README.md展开系统讲解该插件如何将 ONNXScatterElements算子源自 pytorch_scatter 的 scatter 语义在 TensorRT 中落地为高性能 GPU 内核。你将掌握插件两种版本的接口差异IPluginV3与已弃用的IPluginV2DynamicExt、全部输入输出约定与参数语义、受支持的数据类型矩阵以及其基于 CUDA 原子操作的底层实现原理可直接用于在自有推理图中集成该算子。插件概述ScatterElements 插件实现了 scatter 操作其语义对齐两个权威来源pytorch_scatter原实现与文档以及ONNX 规范中的 ScatterElements 算子ONNX Operators 文档。该操作在需要按索引把一批更新值写入张量指定位置的场景中非常常见典型应用包括稀疏聚合、图神经网络的消息传递、序列对齐等。一个需要特别留意的设计取舍reducenone的 ScatterElements 已由 TensorRT 核心TRT core直接实现不属于本插件职责范围。也就是说本插件只负责带归约语义add / mul / max / min的 scatter 场景。结构同一算子两代实现从源码结构看该插件目录下同时存在两套实现见 plugin/scatterElementsPlugin 下的文件布局版本插件类Creator 类接口插件 name / version最新版v2ScatterElementsPluginV3ScatterElementsPluginV3CreatorIPluginV3含IPluginV3OneCore/IPluginV3OneBuild/IPluginV3OneRuntimeScatterElements/2遗留版v1即将弃用ScatterElementsPluginV2ScatterElementsPluginV2CreatorIPluginV2DynamicExtScatterElements/1类声明分别位于 scatterElementsPlugin.hV3 与其 Creator和 scatterElementsPluginLegacy.hV2 与其 Creator。两者通过 scatterElementsCommon.h 共享同一个ReductionType枚举与字符串映射enum class ReductionType : int32_t { kSUM, kMUL, kMEAN, kMIN, kMAX };注意kMEAN虽然在枚举与字符串映射mean中存在但插件公开参数只接受add / mul / max / min四种归约方式见下文参数表mean并非对外暴露的可选值。两套实现共享同一个 CUDA 内核scatterElementsPluginKernel.cu并通过 CMakeLists.txt 的add_plugin_source(...)一并编入插件库。输入输出约定插件共消费三个输入、产出一个输出布局如下与源码中kDATA_TENSOR_IDX 0、kINDICES_TENSOR_IDX 1、kUPDATES_TENSOR_IDX 2、kOUTPUT_TENSOR_IDX 0的常量定义一一对应#名称类型形状约束说明1dataT秩 r ≥ 1待被 scatter 更新的基础张量2indicesTind仅 INT64秩 r ≥ 1与data同秩目标位置索引沿轴大小为 s 时索引合法范围为[-s, s-1]越界即报错3updatesT秩 r ≥ 1与indices同秩同形状要写入的更新值输出outputT与data形状一致归约后的结果张量上述形状与秩约束在 scatterElementsPlugin.cpp 的onShapeChange中被强制执行data 秩必须 ≥ 1、indices 与 data 同秩、updates 与 indices 同秩且逐维形状相等违反任一条件都会触发PLUGIN_ASSERT/PLUGIN_VALIDATE。输出形状方面getOutputShapes 直接复用data的形状表达式getOutputDataTypes 则令输出类型与data的类型保持一致。同时在enqueue入口处scatterElementsPlugin.cpp会强制校验indices必须是DataType::kINT64。参数详解插件对外暴露两个参数类型参数描述intaxis沿哪个轴执行 scatter默认值 0。负值表示从后往前数维度。合法取值范围为[-r, r-1]其中 r rank(data)。charreduction归约方式add加法、mul乘法、max最大值、min最小值。源码层面的参数解析逻辑位于ScatterElementsPluginV3Creator::createPluginscatterElementsPlugin.cppreduction是必填属性requiredFields{reduction}缺失会触发validateRequiredAttributesExist校验失败axis可选缺省时axisArg默认取 0与文档默认值为 0一致reduction字符串经kREDUCE_STR_TO_ENUM映射为枚举add→kSUM、mul→kMUL、min→kMIN、max→kMAX非法取值直接报错。在 TensorRT 的 Python API 中可通过plugin_creator按如下方式创建伪代码示意属性名与上表严格对应registry trt.get_plugin_registry() creator registry.get_plugin_creator(ScatterElements, 2, ) # 注意非空的 plugin namespace 需要自行传入 reduction trt.PluginField(reduction, np.array([add], dtypenp.bytes_), trt.PluginFieldType.CHAR) axis trt.PluginField(axis, np.array([1], dtypenp.int32), trt.PluginFieldType.INT32) fc trt.PluginFieldCollection([reduction, axis]) plugin creator.create_plugin(my_scatter, fc, trt.TensorRTPhase.BUILD)数据类型支持矩阵插件支持的数据类型组合由 ScatterElementsPlugin_PluginConfig.yaml 明确定义该文件同时是插件自动验证测试的配置来源共有 5 种受支持的组合组合dataindicesupdates1float32int64float322int32int64int323int64int64int644float16int64float165bfloat16int64bfloat16这与supportsFormatCombinationscatterElementsPlugin.cpp中的运行时校验一致格式必须是kLINEARindices 必须是 INT64其余输入/输出类型必须相同且只接受kFLOAT / kHALF / kBF16 / kINT32 / kINT64。一个值得注意的硬件相关细节BF16 支持是有条件的。supportsFormatCombination中 BF16 分支额外要求hasBfloat16AtomicAdd()返回 true而该函数scatterElementsPluginKernel.cu通过cudaGetDeviceProperties检查当前设备major 8即仅在Ampere 及更新架构SM 8.0上允许 BF16。这源于 BF16 的atomicAdd指令需要硬件原生支持老架构上无法使用。内核实现原理从 scatterElementsPluginKernel.cu 的源码结构看内核执行分为两步设备到设备拷贝runScatterElementsKernel首先通过cudaMemcpyAsync(..., cudaMemcpyDeviceToDevice, stream)将data的完整内容拷贝到输出缓冲区确保未被更新的位置保持原值若updates元素数为 0 则提前返回。原子归约写回启动scatterElements_kernel线程数按THREADS 256、BLOCKS(N) (N 255) / 256配置每个线程对应updates中的一个元素根据indices提供的偏移用对应的原子操作把updates值写入输出张量的目标位置。内核把整个 scatter 问题分解为三维循环结构nBaxis 之前的批量维、nEaxis 维度大小、nKaxis 之后的元素个数、nN输出在 axis 上的大小并通过AT_DISPATCH_REDUCTION_TYPES宏reducer.cuh在编译期按归约类型实例化内核从而把ReductionType作为模板参数tReduce内联展开避免运行时分支开销。归约的原子语义由 reducer.cuh 中的ReducerTScalar, tReduce模板提供核心是atomic_write分派add / mean→atomAddmul→atomMulmin→atomMinmax→atomMax原子操作的底层支撑atomics.cuh 为上述四种归约在多种数据类型上提供了完整的原子实现这是理解插件正确性的关键硬件原生路径32 位整型的atomicAdd/atomicMin/atomicMax、SM 7.0 的__half原子加法、SM 8.0 的__nv_bfloat16原子加法直接调用 CUDA 内建函数CAS 模拟路径ATOMIC(NAME)宏为整数与浮点4/8 字节生成基于atomicCAS的读-改-写自旋循环do { assumed old; old atomicCAS(...); } while (assumed ! old)用于 64 位整型加法/乘法、32 位浮点乘法等无硬件指令的场景半精度对齐技巧2 字节类型__half/__nv_bfloat16通过reinterpret_caststd::uint32_t*((char*)address - ((size_t)address 2))先对齐到 4 字节字再在该字的对应半区上做 CAS从而在不引入竞争的前提下实现半精度原子乘法/最值。该文件头部注释表明其源自 pytorch_scatterCopyright (c) 2020 Matthias Fey与 README 中符合 pytorch_scatter 语义的定位相互印证。序列化与生命周期作为IPluginV3插件ScatterElementsPluginV3通过getFieldsToSerialize()scatterElementsPlugin.cpp把两个参数序列化进引擎文件reduction以字符串形式序列化PluginFieldType::kCHAR长度即字符串字节数与attribute_length: reduction: -1的 YAML 约定一致axis以kINT32类型、长度 1 序列化。getWorkspaceSize返回 0即该插件不需要额外 workspaceclone()与attachToContext()通过复制(mReduction, mAxis)构造新实例来支持多 context 场景。构建、测试与验证配置插件的自动化验证由 ScatterElementsPlugin_PluginConfig.yaml 驱动它同时充当插件规范声明与测试用例生成器两种角色接口声明interface: IPluginV3、versions: { 2: {...} }输入data/indices/updates输出output属性规范axisint32长度 1缺省不限范围、reductionchar长度 -1 即变长字符串为可选/必填组合其中reduction是必填属性attributes_required: [reduction]测试矩阵为上述 5 种类型组合分别生成config1~config5每个 config 覆盖axis ∈ {-1, 0, 1}与reduction ∈ {add, mul, min, max}的全部笛卡尔积数值容差abs_tol/rel_tol均为1e-2FP16 与 BF16 的rtol/atol放宽到5e-2bf16_rtol/bf16_atol/fp16_rtol/fp16_atol印证了低精度下对误差的预期管理。已知限制根据 README 的 Known issues 一节使用本插件前需确认以下限制数据类型TBFLOAT16与TINT8当前不受支持。需注意这与上文BF16 出现在支持矩阵中并不矛盾——ScatterElementsPlugin_PluginConfig.yaml是测试配置清单包含 BF16 组合而 README 是相对较早的行为声明以实际运行时的supportsFormatCombination校验为准该校验要求 BF16 必须满足hasBfloat16AtomicAdd()SM 8.0才放行。索引类型ONNX 规范允许Tindint32但本插件只支持 INT64导入含 int32 索引的 ONNX 模型时需要先行转换索引类型。变更记录2024 年 7 月插件第 2 版迁移至IPluginV3接口设计使用IPluginV2DynamicExt的遗留插件第 1 版被标记为弃用。2023 年 10 月本 README 文件首次发布。延伸阅读围绕该插件可在仓库内继续深入以下资源ScatterElements 插件 README本文依据插件规范与测试配置 YAMLV3 插件与 Creator 实现V2 遗留插件实现CUDA 内核实现归约器与原子操作支撑 / [plugin/scatterElementsPlugin/atomics.cuh)License插件的使用、复制与分发条款见 TensorRT Software License Agreement。仓库内各源文件顶部均带有 Apache-2.0 SPDX 许可声明内核文件同时保留了 pytorch_scatterMatthias Fey的原始版权声明。【免费下载链接】TensorRTNVIDIA® TensorRT™ is an SDK for high-performance deep learning inference on NVIDIA GPUs. This repository contains the open source components of TensorRT.项目地址: https://gitcode.com/GitHub_Trending/tens/TensorRT创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考