CANN ops-cv ThreeInterpolate 算子深度解析:三近邻加权特征插值的定义、约束与 SIMT 实现

发布时间:2026/9/18 11:35:26
CANN ops-cv ThreeInterpolate 算子深度解析:三近邻加权特征插值的定义、约束与 SIMT 实现 CANN ops-cv ThreeInterpolate 算子深度解析三近邻加权特征插值的定义、约束与 SIMT 实现【免费下载链接】ops-cv本项目是CANN提供的图像处理、目标检测相关的算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-cvThreeInterpolate 是 CANN ops-cv 图像算子库中实现基于 3 个最近邻的加权线性特征插值的核心算子典型应用于 PointNet 等点云网络的 feature propagation 阶段。本文以仓库中 image/three_interpolate/README.md 为骨架结合算子定义、InferShape、Tiling 与 SIMT Kernel 源码完整讲解该算子的功能语义、参数约束、产品支持情况、图模式调用方式及底层实现原理帮助读者快速掌握在 NPU 上使用与理解 ThreeInterpolate 的全链路知识。功能说明ThreeInterpolate 算子的功能是基于 3 个最近邻的加权线性特征插值。给定已知特征点的特征描述符features、待插值点的 3 个最近邻索引idx和对应权重weight对每个待插值点的每个特征通道计算 3 个最近邻特征的加权和作为该点的插值特征。其计算公式为$$ y[b, n, c] \sum_{k0}^{2} weight[b, n, k] \times features[b, idx[b, n, k], c] $$其中b表示 batch 维度索引n表示待插值点索引c表示特征通道索引idx[b, n, k]给出第b个 batch 中第n个待插值点的第k个最近邻在已知特征点集合中的位置weight[b, n, k]是对应第k个最近邻的权重通常由距离倒数归一化得到三个权重之和一般为 1。该计算模式是 PointNet 特征传播Feature Propagation阶段的标准做法将上采样点集的特征通过其 3 个最近邻已知点特征加权得到因此本算子在点云分类、分割等三维视觉网络的推理与训练中均有典型应用场景。产品支持情况按照仓库 image/three_interpolate/README.md 的声明该算子在以下产品的支持情况如下产品是否支持Ascend 950PR/Ascend 950DT√Atlas A3 训练系列产品/Atlas A3 推理系列产品√Atlas A2 训练系列产品/Atlas A2 推理系列产品√Atlas 200I/500 A2 推理产品×Atlas 推理系列产品×Atlas 训练系列产品×需要说明的是从源码结构看本算子的实现目前位于arch35目录对应 Ascend 950 系列及后续架构这也与上述Ascend 950PR/Ascend 950DT、Atlas A3/A2 系列受支持其他旧系列产品不支持的声明一致。参数说明ThreeInterpolate 算子共包含 3 个输入和 1 个输出全部为 NDNCHW 无关的通用多维数据格式。详细参数如下参数名输入/输出/属性描述数据类型数据格式features输入已知特征点集合shape(B, M, C)。BbatchM已知点数C通道数。FLOAT、FLOAT16NDidx输入3 个最近邻索引shape(B, N, 3)。N待插值点数每行 3 个索引值取值范围 [0, M)。INT32、INT64NDweight输入3 个最近邻权重shape(B, N, 3)。与 idx 对应通常由距离倒数归一化得到。FLOAT、FLOAT16NDy输出插值后的特征shape(B, N, C)。BbatchN待插值点数C通道数。FLOAT、FLOAT16ND该参数定义在算子原型注册文件 image/three_interpolate/op_host/three_interpolate_def.cpp 中四个 tensor 均被声明为REQUIRED且带AutoContiguous()属性与 README 中所有输入 tensor 必须为连续contiguous格式的约束对应。此外算子注册时通过PrecisionReduceFlag(true)开启精度降低支持并通过DynamicShapeSupportFlag(true)、DynamicRankSupportFlag(true)开启动态 shape 与动态 rank 支持。约束说明使用 ThreeInterpolate 算子时需遵守以下约束features 和 weight 必须同 dtype均为 float32 或均为 float16。该约束在 InferShape 与 Tiling 两处均有强制校验若不一致会报错 weight dtype must be same as features dtype。idx 可独立选择 int32 或 int64与 features/weight 的 dtype 无关。所有输入 tensor 必须为连续contiguous格式。idx 值必须在 [0, M) 范围内由调用方保证不越界。作为兜底Kernel 内部仍会对索引做 clamp 处理见下文源码解析避免越界访问导致的非法内存读取。维度上限B、M、C、N 均不超过 2^32-1且 N×C 不超过 2^32-1B×N×3 与 B×M×C 不超过 2^64-1。支持空 tensorshape 含 0 维此时输出为空 tensor不执行计算。另外当输出非空而 M0没有任何已知特征点可插值时Tiling 阶段会直接报错拒绝执行参见 image/three_interpolate/op_host/arch35/three_interpolate_tiling_arch35.cpp。源码级实现原理1. 算子 IR 原型声明算子的图 IR 原型声明位于 image/three_interpolate/op_graph/three_interpolate_proto.h通过REG_OP(ThreeInterpolate)宏声明了三个输入features、idx、weight与一个输出y数据类型与 README 参数表完全一致REG_OP(ThreeInterpolate) .INPUT(features, TensorType({DT_FLOAT, DT_FLOAT16})) .INPUT(idx, TensorType({DT_INT32, DT_INT64})) .INPUT(weight, TensorType({DT_FLOAT, DT_FLOAT16})) .OUTPUT(y, TensorType({DT_FLOAT, DT_FLOAT16})) .OP_END_FACTORY_REG(ThreeInterpolate)2. InferShape 形状推导InferShape 实现在 image/three_interpolate/op_host/three_interpolate_infershape.cpp核心逻辑包括输入必须均为 3DDIM_NUM 3否则报错校验 features/weight 的 dtype 一致性、idx 的 dtype 合法性校验 B 维度features 与 idx、features 与 weight相等、N 维度idx 与 weight相等校验 idx 与 weight 的最后一维必须等于 3NEIGHBOR_NUM 3输出 shape 推导为(B, N, C)其中 B 取 features 的 BN 取 idx 的 NC 取 features 的 C当输入存在未知 rank动态 rank 场景时输出直接置为 3D 未知 shape(-1, -1, -1)在 Ascend 950short_soc_version Ascend950上走完整校验路径其他平台走简化推导路径。对应单元测试位于 image/three_interpolate/tests/ut/op_host/test_three_interpolate_infershape.cpp。3. Tiling 计算与核数分配Tiling 实现在 image/three_interpolate/op_host/arch35/three_interpolate_tiling_arch35.cpp其工作要点如下重复 InferShape 的结构化校验保证动态 shape 场景下在 Tiling 阶段同样受到约束保护读取平台 AIV 核数GetCoreNumAiv与 UB 内存大小并按128KB DCACHE预留后计算可用的本地内存SetLocalMemorySize(ubSize - DCACHE_SIZE - STATIC_UB_ESTIMATE)以输出总元素数totalElements bs * ns * cs为基准进行负载切分先按核数均分每核最少PER_CORE_MIN 1024个元素再按ALIGN_SIZE 32对齐最终得出实际核数needCoreNum并写入 tiling 数据、通过SetBlockDim设置核数对空 tensortotalElements0与 M0 但输出非空的非法场景做了专门防护申请系统 workspaceGetLibApiWorkSpaceSize供 Kernel 内部使用将bs/ns/cs/ms/needCoreNum/totalElements写入 image/three_interpolate/op_kernel/arch35/three_interpolate_tiling_data.h 定义的ThreeInterpolateTilingData结构体。4. SIMT Kernel 计算核心Kernel 入口为 image/three_interpolate/op_kernel/three_interpolate_apt.cpp实际计算逻辑在 image/three_interpolate/op_kernel/arch35/three_interpolate_simt.h其实现特征包括采用 Grid-Stride 循环模式每线程固定THREAD_NUM 512按index blockDim.x * gridDim.x步进遍历全部输出元素使用UintDivmagic number shift替代除法完成(b, n, c)坐标分解减少整数除法开销magic/shift 参数存放在 UB 中按totalElements是否超过INT32_MAX自动选择 32 位或 64 位索引算术路径兼顾大 shape 场景的溢出安全GM 偏移一律用 64 位计算bs*ns*3与bs*ms*cs可能超出 uint32 范围索引安全处理将 int32/int64 索引统一转成 uint64 后 clamp 到[0, ms-1]防止非法索引导致越界读精度策略对 float16 输入乘加过程先提升到 float32 累加再写回降低精度损失对 float32 输入则行为完全一致直接访问 GM不经过 UB 数据缓冲UB 仅用于存放 UintDiv 的 magic/shift 参数。5. Golden 参考实现测试基准测试基准脚本 image/three_interpolate/tests/assets/golden.py 提供了与 Kernel 完全对齐的参考实现同样以 float32 作为计算 dtype并按w0*f0 w1*f1 w2*f2的单表达式顺序而非逐个累加计算最后再 cast 回原始 dtype。这样从数值上保证与 SIMT Kernel 的浮点舍入行为一致可用作算子精度比对的标准答案。调用说明调用方式调用样例说明图模式test_geir_three_interpolate通过算子IR构图方式调用ThreeInterpolate算子参见算子调用完成编译和验证。图模式样例 image/three_interpolate/examples/arch35/test_geir_three_interpolate.cpp 展示了完整的调用流程其关键步骤包括创建算子节点通过op::ThreeInterpolate(add1)创建算子实例设置输入输出 shape样例中 features 为{2, 4, 3}idx 与 weight 为{2, 2, 3}输出 y 为{2, 2, 3}即 B2、M4、N2、C3构造输入数据通过ADD_INPUT宏为 features、idx、weight 生成占位数据idx 固定为DT_INT32构图与运行将算子加入Graph设置ge.exec.deviceId与ge.graphRunMode全局选项初始化 GEGEInitialize后创建Session通过AddGraph添加计算图最后RunGraph执行并落盘输入输出 bin 文件资源回收GEFinalize()收尾。编译与验证的完整工程流程可参考算子调用文档同时该算子的 Tiling 单测位于 image/three_interpolate/tests/ut/op_host/test_three_interpolate_tiling.cpp。小结ThreeInterpolate 算子以3 个最近邻加权求和这一简洁的数学语义配合连续 tensor、dtype 一致、索引合法等严格约束为点云类网络的 NPU 加速提供了高效的特征插值原语。从 IR 原型、InferShape、Tiling 到 SIMT Kernel仓库源码形成了完整的闭环形状推导保证输出(B, N, C)的合法性Tiling 以 32 字节对齐的最小 1024 元素粒度做负载均衡Kernel 用 float32 累加与索引 clamp 兼顾精度与安全。开发者可直接参考 examples/arch35 样例 以图模式接入该算子并按 README 中的参数表与约束表完成数据准备。【免费下载链接】ops-cv本项目是CANN提供的图像处理、目标检测相关的算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-cv创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考