ARTICLE DETAIL

资讯详情

深耕商务建站与企业官网运营的一线实战洞察。

CANN ops-math BallQuery 算子技术指南:PointNet++ 球查询算子的原理、参数与源码实现解析

CANN ops-math BallQuery 算子技术指南:PointNet++ 球查询算子的原理、参数与源码实现解析 CANN ops-math BallQuery 算子技术指南PointNet 球查询算子的原理、参数与源码实现解析【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-math导读BallQuery球查询是 PointNet 系列点云网络中的经典算子用于在给定半径范围内为每个查询中心点收集邻域点的索引。本文以 conversion/ball_query/README.md 为主线完整梳理该算子在 CANN ops-math 中的产品支持情况、数学定义、参数约束与调用方式并深入到算子定义、InferShape、Tiling 与 SIMT Kernel 的实现源码帮助读者既能在实践中正确使用该算子也能理解它在 NPU 上的并行计算原理。一、算子概述BallQuery 在点云网络中的角色Ball Query 是 PointNet 中用于局部区域特征聚合的核心操作。与 KNNK 近邻相比Ball Query 通过固定半径的球形邻域替代 K 近邻从而保证邻域采样在空间上是各向同性isotropic的不受点云密度分布影响因而在点云分割、分类等任务中被广泛采用。在 CANN ops-math 仓库中该算子位于 conversion/ball_query 目录属于 conversion转换/采样类算子集合。它的基本职责是对每个查询中心点center_xyz[m,b]在点集xyz[b]中按k 0..N-1的自然顺序查找满足距离判定条件的点并收集前sample_num个点的索引到输出idx中。产品支持情况该算子当前支持以下产品的 NPU 运行环境产品是否支持Ascend 950PR / Ascend 950DT√Atlas A3 训练系列产品 / Atlas A3 推理系列产品√Atlas A2 训练系列产品 / Atlas A2 推理系列产品√Atlas 200I/500 A2 推理产品×Atlas 推理系列产品√Atlas 训练系列产品√从 算子定义 可以看到源码中通过this-AICore().AddConfig(ascend950, aicoreConfig)显式注册了 ascend950 平台的 AI Core 配置与文档中Ascend 950PR/950DT 支持的声明相互印证同时算子 Kernel 与 Tiling 均位于arch35目录op_kernel/arch35对应 950 系列所采用的架构。二、功能与计算公式2.1 数据布局约定BallQuery 的输入输出采用如下非对称的 3D 布局xyzshape 为(B, 3, N)坐标位于中间维。其中B为 batch 数N为每个 batch 内的点云点数第二维固定为 3 表示 (x, y, z) 坐标。center_xyzshape 为(M, B, 3)坐标位于最后一维。M即 npoint为每个 batch 内的查询中心点数量其第 1 维B必须与xyz的第 0 维B一致。idx输出shape 为(M, B, sample_num)dtype 固定为 INT32。注意两个输入张量的坐标维位置不同xyz 坐标在中间维center_xyz 坐标在末维这是该算子的一个重要记忆点在编写推理脚本或构图代码时极易混淆。2.2 距离计算与命中判定对每个(m, b)取中心点 $(c_x, c_y, c_z) center_xyz[m, b, :]$计算其与xyz[b]中每个点 $(x_k, y_k, z_k) (xyz[b, 0, k], xyz[b, 1, k], xyz[b, 2, k])$ 的距离平方$$ d2 (c_x - x_k)^2 (c_y - y_k)^2 (c_z - z_k)^2 $$距离判定条件为$$ d2 0 \quad 或 \quad min_radius^2 \le d2 max_radius^2 $$即命中点要么与中心点完全重合距离为 0典型场景为查询中心点本身取自点云要么落在以中心点为圆心、min_radius为内半径、max_radius为外半径的球形环带内。当min_radius 0时即退化为经典的实心球邻域查询。2.3 输出收集与填充规则命中点按照k的自然顺序0 到 N-1依次写入idx首个命中点的索引记为first_num当命中点数量不足sample_num时剩余位置用first_num填充当没有任何命中点时剩余位置统一填充 0该行为在 Kernel 实现中体现见下文源码分析。三、参数说明参数名输入/输出/属性描述数据类型数据格式xyz输入所有点的 xyz 坐标shape 为(B,3,N)中间维必须为 3坐标在中间维FLOAT、FLOAT16NDcenter_xyz输入查询中心点坐标shape 为(M,B,3)最后一维必须为 3第 1 维 B 必须与 xyz 的第 0 维一致与 xyz 保持一致NDidx输出采样点在 xyz 中的索引shape 为(M,B,sample_num)INT32NDmin_radius属性最小查询半径环形内边界必须大于等于 0Float-max_radius属性最大查询半径球外边界必须大于 min_radiusFloat-sample_num属性每球最大采样点数必须大于等于 1Int-3.1 算子定义层的参数约束在 算子 IR 定义 中三个属性均被声明为必填REQUIRED_ATTRREG_OP(BallQuery) .INPUT(xyz, TensorType({DT_FLOAT16, DT_FLOAT})) .INPUT(center_xyz, TensorType({DT_FLOAT16, DT_FLOAT})) .OUTPUT(idx, TensorType({DT_INT32})) .REQUIRED_ATTR(min_radius, Float) .REQUIRED_ATTR(max_radius, Float) .REQUIRED_ATTR(sample_num, Int) .OP_END_FACTORY_REG(BallQuery)而在 ball_query_def.cpp 中算子注册了如下关键配置输入xyz、center_xyz均声明AutoContiguous()即框架会对非连续输入自动做连续化处理对应 README 中非连续 Tensor 自动连续化的说明AICore 配置开启DynamicShapeSupportFlag(true)、DynamicRankSupportFlag(true)、DynamicCompileStaticFlag(true)说明该算子支持动态 shape 与静态编译相结合的调度方式ExtendCfgInfo(opFile.value, ball_query_apt)指定了算子内核入口文件为 ball_query_apt.cpp。四、约束说明以下是使用 BallQuery 时必须遵守的约束条件其中多数约束在 InferShape 与 Tiling 源码中都有对应的校验逻辑见第五节xyz为 3 维 Tensorshape 为(B,3,N)中间维必须为 3center_xyz为 3 维 Tensorshape 为(M,B,3)最后一维必须为 3xyz的第 0 维 B 与center_xyz的第 1 维 B 必须一致xyz的第 0 维 B、第 2 维 N以及center_xyz的第 0 维 Mnpoint取值均不得超过 INT32_MAX2147483647xyz与center_xyz的 dtype 必须相同且为 FLOAT 或 FLOAT16输出idx的 dtype 固定为 INT32shape 为(M,B,sample_num)min_radius必须为有限值非 NaN、非 Inf且大于等于 0max_radius必须为有限值非 NaN、非 Inf且大于 0并必须大于min_radiussample_num必须大于等于 1 且不超过 INT32_MAX输入 Tensor 需为连续内存布局contiguous即数据在内存中按行主序紧密排列不含 stride 跳步或间隙概念详见非连续的Tensor。若传入非连续 Tensor如转置、切片得到的视图框架会自动连续化处理产生额外拷贝开销算子仅支持 3D 输入不支持标量、1D 或 8D 场景。五、源码级实现解析5.1 InferShape输出形状推导与校验ball_query_infershape.cpp 实现了形状推导函数InferShape4BallQuery其核心逻辑与 README 约束一一对应坐标维校验对xyz检查dimNum 3 dim(1) 3对center_xyz检查dimNum 3 dim(2) 3。当输入为 unknownRank动态秩时会跳过维数校验Batch 一致性校验在非 unknownRank 场景下要求xyz.dim(0) center_xyz.dim(1)若任一维为动态值 -1 则跳过等值校验dtype 校验要求xyz与center_xyz的 dtype 相同且取值必须为DT_FLOAT或DT_FLOAT16属性校验min_radius必须isfinite且 0max_radius必须isfinite且 0且 min_radiussample_num必须 1输出推导idx.shape (center_xyz.dim(0), center_xyz.dim(1), sample_num)即(M, B, sample_num)。若center_xyz为 unknownRank则 M、B 位置写 -1rank 与 sample_num 仍已知。仓库中的 InferShape 单元测试 覆盖了 FP32 静态 shapexyz(5,3,15)、center_xyz(20,5,3)、sample_num4期望输出idx(20,5,4)与 FP16 动态 shape期望输出idx(-1,2,16)等典型场景可作为理解输出推导规则的参考用例。5.2 Tiling多核切分策略ball_query_tiling.cpp 实现了 host 侧的切分逻辑核心步骤包括平台信息获取通过platform_ascendc::PlatformAscendC获取 AIV 核数coreNum与 UB 内存大小输入与属性校验与 InferShape 相同规格的 shape/dtype/属性合法性检查并将 B、N、npoint 收窄为 int32核数切分以查询点总数为并行单元totalQueryPoints B * npointperCoreElements ceil(totalQueryPoints / coreNum)并设置每核最小处理量PER_CORE_MIN 64同时向上对齐到WARP_SIZE 32最终needCoreNum ceil(totalQueryPoints / perCoreElements)通过context-SetBlockDim()下发TilingData 填充向 BallQueryTilingData 写入B、N、npoint、sampleNum以及预计算的半径平方minRadius2 min_radius²、maxRadius2 max_radius²。将半径平方预计算放到 host 侧可避免 kernel 端重复开方与乘法运算TilingKey 设置通过 ball_query_tiling_key.h 中的BALL_QUERY_SCH_MODE_DEFAULT宏选择默认调度模式并声明为ASCENDC_TPL_AIV_ONLY仅使用 AIV 向量/标量核。5.3 SIMT Kernel并行查询与索引收集ball_query_simt.h 是算子的核心 SIMT 内核实现采用Grid-Stride 循环每个线程以stride blockDim.x * gridDim.x为间隔处理多个查询点适合查询点数远大于线程数的场景。其处理流程为索引分解由一维查询点索引qIdx反解(m, b)qIdx m*B b其中除法采用Simt::UintDiv的 magic/shift 快速除法由 Process 阶段通过GetUintDivMagicAndShift预计算GM 基址计算分别定位xyz[b, 0, 0]、center_xyz[m, b, :]、idx[m, b, :]三个地址坐标读取与精度提升读取中心点坐标后通过CoordToFloat统一提升为 float32——对 half 类型调用__half2float避免 FP16 计算距离平方时的溢出与截断距离判定与收集遍历该 batch 内全部 N 个点按d2 0.0f || (d2 minRadius2 d2 maxRadius2)判定命中命中索引按序写入idxBase[cnt]并记录首个命中点firstResult当cnt达到sampleNum时提前终止内层循环填充与未命中处理命中数不足时用firstResult填充剩余槽位完全未命中时将所有槽位置 0。这与 README 中命中数不足 sample_num 时剩余位置用 first_num 填充的语义一致并补充了零命中时填 0的实现细节。Kernel 入口 ball_query_apt.cpp 通过REGISTER_TILING_DEFAULT/GET_TILING_DATA_WITH_STRUCT读取 TilingData并按schMode BALL_QUERY_SCH_MODE_DEFAULT分派到NsBallQuery::Process。六、调用方式GE 图模式6.1 调用总览调用方式调用样例说明GE 图模式test_geir_ball_query通过算子IR构图方式调用 BallQuery 算子参见算子调用完成算子编译和验证。当前仓库为 BallQuery 提供的官方调用样例是GEGraph Engine图模式即通过ge::op::BallQuery(ball1)在计算图中创建算子节点再经AddGraphRunGraph完成构图、编译与执行。6.2 构图与输入输出声明示例程序 test_geir_ball_query.cpp 中关键构图逻辑位于CreateOppInGraphauto ball1 op::BallQuery(ball1); // 输入xyz(B,3,N) 与 center_xyz(M,B,3) ADD_INPUT_MODE(1, xyz, inDtype, xyzShape, xyzDShape, mode); ADD_INPUT_MODE(2, center_xyz, inDtype, centerXyzShape, centerXyzDShape, mode); // 输出idx(M,B,sample_num)dtype 固定 INT32 std::vectorint64_t idxShape {centerXyzShape[0], centerXyzShape[1], sampleNum}; ADD_OUTPUT_MODE(1, idx, DT_INT32, idxShape, mode); // 属性设置min_radius - max_radius - sample_num ball1.set_attr_min_radius(minRadius); ball1.set_attr_max_radius(maxRadius); ball1.set_attr_sample_num(sampleNum);该示例同时演示了静态S与动态D两种构图模式。特别值得注意的是 D 模式下动态 shape 的设置方式由于 InferShape 会校验坐标维因此不能把所有维都置为 -1而应仅将动态维置 -1、保留坐标维为已知值 3即xyz使用 D shape{-1, 3, -1}center_xyz使用 D shape{-1, -1, 3}。这一细节体现了 BallQuery 坐标维布局约束对动态 shape 构图的影响是实际开发中容易踩坑的点。6.3 执行与验证矩阵示例程序的测试矩阵体现了算子的完整适用边界dtype 矩阵DT_FLOATFP32与DT_FLOAT16FP16两类输入输出固定 INT32shape 场景矩阵共 5 类——常规 3Dxyz(2,3,1024)、center_xyz(512,2,3)、sample_num16、小规模 3D、空点云N0、空查询M0、多 batchB4运行模式每个 (dtype, shape) 组合都同时跑 S静态与 D动态两种模式最终输出逐 case 的 Build / RunGraph / OutputExists 报告与通过率统计。程序启动时会通过全局选项初始化 GE 引擎包括ge.exec.deviceId需通过npu-smi info查询状态为 OK 的设备号、ge.graphRunMode以及ge.exec.precision_mode must_keep_origin_dtype保持原始 dtype 精度。上述 shape 组合可直接作为读者自行编写验证脚本的参考模板。七、总结与使用建议BallQuery 是 CANN ops-math 中面向点云处理的邻域采样算子其核心设计可以概括为四点布局特殊xyz坐标在中间维、center_xyz坐标在末维构图与推理时必须严格区分判定规则d2 0或min_radius² ≤ d2 max_radius²才计入采样天然支持环形空心球邻域半径平方预计算host 侧 Tiling 阶段完成min_radius²、max_radius²计算Kernel 仅做乘加与比较降低片上计算开销并行充分以B × npoint个查询点为并行单元Grid-Stride 循环配合 magic-division 索引反解适配 AIV 多核场景。实际使用时建议先对照约束说明核对 shape/dtype/属性取值再按算子调用指南编译运行示例 test_geir_ball_query.cpp若需深度定制可进一步阅读 Tiling 实现 与 SIMT Kernel 了解切分与并行细节或参考 InferShape UT 中的用例组织方式编写自己的验证用例。【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-math创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表
PREV
查看更多资讯
NEXT
返回资讯列表