CANN ops-cv GridSampler2D 算子完全指南:二维网格采样、坐标变换与 NPU 实现解析
【免费下载链接】ops-cv本项目是CANN提供的图像处理、目标检测相关的算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-cv
本文围绕 CANN ops-cv 仓库中的 GridSampler2D 算子文档,系统讲解该算子的功能语义、坐标变换公式、参数与约束、NPU 侧的实现链路,并结合仓库源码(算子定义、InferShape、Tiling、SIMT Kernel 与图模式示例)给出可直接落地使用的调用方式。读完本文,你将掌握在 Ascend 平台上通过 GE IR 构图调用 GridSampler2D 的完整方法,并理解其与 PyTorchgrid_sample对齐的插值、填充与精度处理细节。
产品支持情况
GridSampler2D 算子在当前仓库中面向以下产品提供支持(摘自 README):
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR/Ascend 950DT | √ |
| Atlas A3 训练系列产品/Atlas A3 推理系列产品 | √ |
| Atlas A2 训练系列产品/Atlas A2 推理系列产品 | √ |
| Atlas 200I/500 A2 推理产品 | × |
| Atlas 推理系列产品 | × |
| Atlas 训练系列产品 | × |
从源码结构看,算子的宿主侧实现存放于 op_host/arch35,内核实现位于 op_kernel/arch35,二进制配置注册于 op_host/config/ascend950/grid_sampler2_d_binary.json,可以推断当前版本主要面向 arch35(Ascend 950 系列)编译与验证。
功能说明
算子语义
GridSampler2D 根据grid提供的归一化坐标,对四维输入x进行二维网格采样,支持**双线性(bilinear)、最近邻(nearest)和双三次(bicubic)**三种插值方式。该算子与 PyTorch 的grid_sample语义兼容,这一点在算子原型注释中明确说明(见 grid_sampler2_d_proto.h)。
输入输出尺寸
各张量的 shape 约定如下:
$$ x: (N, C, H_{in}, W_{in}) $$
$$ grid: (N, H_{out}, W_{out}, 2) $$
$$ y: (N, C, H_{out}, W_{out}) $$
grid的最后一维依次存放 x 和 y 坐标,坐标通常归一化到[-1, 1]。实际输入坐标由align_corners决定:
align_corners=true(角像素中心对齐):$$ x' = \frac{grid_x + 1}{2}(W_{in}-1), \quad y' = \frac{grid_y + 1}{2}(H_{in}-1) $$
align_corners=false(像素中心位于半像素偏移处):$$ x' = \frac{(grid_x + 1)W_{in}-1}{2}, \quad y' = \frac{(grid_y + 1)H_{in}-1}{2} $$
越界坐标按照padding_mode指定的zeros、border或reflection方式处理,再按照interpolation_mode指定的bilinear、nearest或bicubic方式计算输出。
内核中对这两种归一化路径有完整实现,UnnormalizeNoClip对应上述两套公式;其中align_corners=false路径使用fmaf(scalingFactor, coord + 1.0f, -0.5f)单次舍入,以对齐 PyTorch x86 编译产物的 FMA 行为(见 grid_sampler2_d_simt.h)。
参数说明
以下参数表完整继承自 GridSampler2D 算子文档:
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| x | 输入 | 输入特征图,shape 为 (N, C, H_in, W_in)。 | FLOAT16、FLOAT32 | NCHW |
| grid | 输入 | 采样网格,shape 为 (N, H_out, W_out, 2),数据类型必须与 x 一致。 | FLOAT16、FLOAT32 | NHWC |
| interpolation_mode | 可选属性 | 插值模式,支持 "bilinear"、"nearest" 和 "bicubic",默认值为 "bilinear"。 | STRING | - |
| padding_mode | 可选属性 | 填充模式,支持 "zeros"、"border" 和 "reflection",默认值为 "zeros"。 | STRING | - |
| align_corners | 可选属性 | 是否将输入和输出的角像素中心对齐,默认值为 false。 | BOOL | - |
| y | 输出 | 采样结果,shape 为 (N, C, H_out, W_out),数据类型与 x 一致。 | FLOAT16、FLOAT32 | NCHW |
上述参数在算子原型注册中均有对应定义(见 grid_sampler2_d_proto.h):x、grid、y均仅支持DT_FLOAT16/DT_FLOAT;三个属性均为可选属性,interpolation_mode默认"bilinear"、padding_mode默认"zeros"、align_corners默认false。
在宿主侧算子定义(grid_sampler2_d_def.cpp)中可以看到更细的格式约定:x为FORMAT_NCHW,grid为FORMAT_ND,输出y为FORMAT_NCHW,三者均声明为AutoContiguous();该定义同时开启了动态 Shape、动态 Rank 支持(DynamicShapeSupportFlag(true)、DynamicRankSupportFlag(true)),并注册了ascend950的 AICore 配置(ExtendCfgInfo("opFile.value", "grid_sampler2_d_apt")、ExtendCfgInfo("opInterface.value", "grid_sampler2_d"))。
约束说明
x与grid均为 4 维张量,数据类型一致且仅支持float16或float32。grid的最后一维必须为 2,且x与grid的 batch 维必须一致。x的H_in和W_in必须大于 0;grid产生空输出时支持空 tensor。x的高宽乘积、grid的 batch 与输出高宽乘积均不能超过INT32_MAX。interpolation_mode仅支持bilinear、nearest、bicubic。padding_mode仅支持zeros、border、reflection。
这些约束在代码中有两层校验:
InferShape 阶段(grid_sampler2_d_infershape.cpp):校验
x/grid均为 4 维、grid最后一维为 2(未知维UNKNOWN_DIM除外)、batch 一致、输入高宽大于 0;随后推导输出 shape——y的N取x的 batch、C取x的通道、H/W取grid的 H/W。数据类型校验则在图推导阶段(grid_sampler2_d_graph_infer.cpp)完成,要求grid与x同类型,且输出数据类型继承自x。Tiling 阶段(grid_sampler2_d_tiling.cpp):再次核对维度数、
grid最后一维为 2、batch 一致、各维非负且输入高宽大于 0,并做INT32_MAX溢出防护——H_in * W_in与N * H_out * W_out均不能超过INT32_MAX。Tiling 同时把三个属性字符串解析为整型枚举:interpolation_mode映射为 0=bilinear、1=nearest、2=bicubic;padding_mode映射为 0=zeros、1=border、2=reflection(见 grid_sampler2_d_tiling.cpp)。
算子实现原理(源码级)
1. Tiling 与核间划分
Tiling 逻辑位于 grid_sampler2_d_tiling.cpp,核心步骤:
- 通过平台接口获取 AIV 核数
coreNum与 UB 内存大小ubSize; - 计算输出像素总数
totalPixels = N * H_out * W_out; - 采用“两步核划分”策略:先按
perCoreElements = CeilDiv(totalPixels, coreNum)估算每核元素数,若小于PER_CORE_MIN(8192)则抬高到 8192,再反推实际需要的核数needCoreNum,从而在小张量场景下避免过度占用核资源; - 预留
DCACHE_SIZE(128 KiB)后设置本地内存大小; - 将
interpolationMode编码进 TilingKey(0/1/2),使内核在编译期通过模板参数分发到 bilinear/nearest/bicubic 三条路径(见 grid_sampler2_d_tiling.cpp)。
Tiling 数据结构定义在 grid_sampler2_d_tiling_data.h,包含N/C/H_in/W_in/H_out/W_out以及三个属性枚举。
2. SIMT Kernel 与坐标处理
内核实现为 SIMT(Single Instruction Multiple Threads)风格,入口见 grid_sampler2_d_apt.cpp,通过DTYPE_X宏自动按数据类型实例化,interpMode由 TilingKey 在编译期确定。核心逻辑在 grid_sampler2_d_simt.h:
- 线程组织:索引宽度模板化——最大地址不超过
INT32_MAX时使用int32_t索引(每 block 1024 线程),否则退化为int64_t索引(512 线程); - 除法优化:
idx / WOut、hw / HOut等常量除法被替换为Simt::UintDiv(magic number + shift),避免除法指令开销; - 坐标计算:
ComputeSourceIndex依次执行反归一化(UnnormalizeNoClip)、按padding_mode处理(zeros不做裁剪、border做ClipCoordinates、reflection先ReflectCoordinates再裁剪)、NaN 坐标归零等; - 三种插值:
BilinearSample:取坐标 floor 后的四个邻域像素,按双线性权重加权;NearestSample:采用BankerRoundToInt(四舍六入五成双,即 PyTorch 默认的 round-half-to-even)确定最近邻;BicubicSample:使用A = -0.75的三次卷积核(CubicConvolution1/2),先在 x 方向对 4×4 邻域做 4 次加权和,再在 y 方向聚合。
- 精度对齐:代码注释详细记录了与 PyTorch CPU 向量化内核逐位对齐的调优过程——对 float32 双三次路径,反归一化与加权和采用
fmaf()匹配 PyTorch x86 的 FMA 链,而三次卷积系数计算则通过volatile变量 + 文件级#pragma clang fp contract(off)禁止 FMA 合并(NPU 的 fmaf 舍入行为与 x86 不同),从而在 NPU 上复现 PyTorch 的舍入语义(见 grid_sampler2_d_simt.h)。
3. 二进制注册
grid_sampler2_d_binary.json 为 ascend950 平台注册了 float16 与 float32 两套high_performance二进制条目,输入输出均使用ND格式、FormatAgnostic匹配模式,三个属性作为可空(value 为 null,采用默认值)的可选属性。
调用说明
当前仓库提供的调用方式为图模式(GE IR):
| 调用方式 | 样例代码 | 说明 |
|---|---|---|
| 图模式 | test_geir_grid_sampler2_d.cpp | 通过算子 IR构图方式调用 GridSampler2D 算子,参见算子调用完成编译和验证。 |
示例程序的核心流程如下(对应 test_geir_grid_sampler2_d.cpp):
- 初始化 GE:调用
ge::GEInitialize,设置ge.exec.deviceId与ge.graphRunMode=1(图模式); - 构图:创建
ge::Graph,用op::Data声明x(NCHW、DT_FLOAT)与grid(NHWC、DT_FLOAT)两个输入节点,并指定set_attr_index(0/1); - 挂接算子节点:
op::GridSampler2D("grid_sampler2_d"),通过set_input_x/set_input_grid连接输入,通过set_attr_interpolation_mode("bilinear")、set_attr_padding_mode("zeros")、set_attr_align_corners(true)设置属性; - 建 Session 并运行:
session->AddGraph后构造输入 Tensor(MakeFloatTensor完成 host 端数据装载),调用session->RunGraph(graphId, inputs, outputs); - 校验结果:
CheckOutput逐元素比对,误差阈值1e-4f。
示例选用 1×1×2×2 的输入x = {1, 2, 3, 4},grid取四个角点坐标(-1/-1、1/-1、-1/1、1/1),在align_corners=true且 bilinear+zeros 配置下,输出应与输入张量完全一致(角像素中心对齐时四个角点恰好落在输入四角像素中心),因此该用例同时验证了坐标变换公式的正确性。
如需完整的编译与验证流程(算子注册、二进制生成、用例编译运行等),请参考 算子调用指南 以及仓库根目录的 CONTRIBUTING.md 与构建脚本说明。
【免费下载链接】ops-cv本项目是CANN提供的图像处理、目标检测相关的算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-cv
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考