Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.
长尾算子开发,MaxPoolV3Grad算子新增950 SIMT实现
vector
op_graph/max_pool_v3_grad_proto.h
REG_OP(MaxPoolV3Grad)
op_host/max_pool_v3_grad_def.cpp
class MaxPoolV3Grad : public OpDef
{DT_FLOAT16, DT_FLOAT}
FORMAT_ND
AutoContiguous()
GetAttrPointer<T>(index)
AddConfig("ascend950", ...)
op_graph/max_pool_v3_grad_graph_infer.cpp
out_grad.dtype = orig_input.dtype
op_host/max_pool_v3_grad_infershape.cpp
IMPL_OP_INFERSHAPE
GetDimNum()==4
out_grad.shape = orig_input.shape
op_host/arch35/max_pool_v3_grad_tiling.cpp 中 MaxPoolV3GradTilingFunc:
op_host/arch35/max_pool_v3_grad_tiling.cpp
MaxPoolV3GradTilingFunc
GetPlatformInfo
ValidateAttrs
check_param
ComputePoolParams
_compute_pool_params
FloorDiv
//
max(.,0)
overlap = (sh<kh)||(sw<kw)
totalOutputPos=N*C*Ho*Wo
needCoreNum
totalInputElements*sizeof(float)
FillTilingData
SetBlockDim
SetScheduleMode(1)
SetTilingKey
DTYPE_
SetLocalMemorySize(ubSize - 128KB DCache)
op_kernel/arch35/max_pool_v3_grad_simt.h 采用 512 线程(constexpr THREAD_NUM=512,__launch_bounds__ 同一常量)的统一 SIMT 架构,纯 GM 计算模式依赖 128KB DCache 加速随机访问,通过 if constexpr (OVERLAP_MODE) 编译期分发非重叠/重叠路径:
op_kernel/arch35/max_pool_v3_grad_simt.h
constexpr THREAD_NUM=512
__launch_bounds__
if constexpr (OVERLAP_MODE)
ProcessOneOutputPos
ComputeOutAddr
ComputeInAddr
matched
RouteGrad
==
outGrad[inAddr]=grad[outAddr]
asc_atomic_add(&wsGrad[inAddr], gradVal)
MaxPoolV3GradSimt
InitPass
__builtin_cce_dcci
ConvertPass
__float2half
op_kernel/arch35/max_pool_v3_grad.cpp 为入口:REGISTER_TILING_DEFAULT + GET_TILING_DATA_WITH_STRUCT,通过 AscendC::GetUserWorkspace(workspace) 获取用户区指针(跳过 ascend950 16MB 系统保留区,避免 overlap 路径踩踏),按 schMode 分发 Process<DTYPE_ORIG_INPUT, 0/1>。
op_kernel/arch35/max_pool_v3_grad.cpp
REGISTER_TILING_DEFAULT
GET_TILING_DATA_WITH_STRUCT
AscendC::GetUserWorkspace(workspace)
Process<DTYPE_ORIG_INPUT, 0/1>
双路径调度:非重叠 = InitPass(out_grad)→SyncAll→直接写;重叠 = InitPass(ws FP32)→SyncAll→atomicAdd ws→SyncAll→ConvertPass(ws→out_grad)。
op_kernel/arch35/max_pool_v3_grad_tiling_data.h
MaxPoolV3GradTilingData
op_kernel/arch35/max_pool_v3_grad_tiling_key.h
ASCENDC_TPL_ARGS_DECL
ASCENDC_TPL_SEL
MAX_POOL_V3_GRAD_TPL_MODE_NON_OVERLAP(0)
MAX_POOL_V3_GRAD_TPL_MODE_OVERLAP(1)
op_host/config/ascend950/max_pool_v3_grad_binary.json
tests/assets/golden.py
golden_max_pool_v3_grad_direct
golden_max_pool_v3_grad_tf
max_pool_v3_grad_golden
tests/ut/op_host/test_max_pool_v3_grad_infershape.cpp
tests/ut/op_host/arch35/test_max_pool_v3_grad_tiling.cpp
Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.
Backgroud(背景信息)
长尾算子开发,MaxPoolV3Grad算子新增950 SIMT实现
Origin(信息来源)
vector
Benefit / Necessity (价值/作用)
Design(设计方案)
算子定义与 Shape 推导
op_graph/max_pool_v3_grad_proto.h:REG_OP(MaxPoolV3Grad),复用 canndev 原型,3 输入(RealNumberType)/ 1 输出,7 个属性(ksize、strides 必选;padding_mode、pads、data_format、global_pooling、ceil_mode 可选)。通过 GEIR 图模式调用,无 aclnn 接口。op_host/max_pool_v3_grad_def.cpp:class MaxPoolV3Grad : public OpDef,输入输出 DataType 收窄为{DT_FLOAT16, DT_FLOAT}、Format 声明FORMAT_ND,全部AutoContiguous()(GE 框架自动转连续,与前向 MaxPoolV3 一致)。属性声明顺序与 tiling/infershape 的GetAttrPointer<T>(index)索引严格一致(0~6)。AddConfig("ascend950", ...)开启动态 shape/rank、PrecisionReduce。op_graph/max_pool_v3_grad_graph_infer.cpp:Dtype 推导out_grad.dtype = orig_input.dtype。op_host/max_pool_v3_grad_infershape.cpp:IMPL_OP_INFERSHAPE,强制 4D 校验(GetDimNum()==4),逐项校验 dtype 一致性/支持 dtype/data_format/padding_mode/global_pooling(必须 false)/ceil_mode(SAME|VALID 下必须 false)/ksize·strides 长度及 N·C 维==1 与 H·W∈[1,255]·[1,63]/pads(CALCULATED 下 ≥0)/grad.shape==orig_output.shape,推导out_grad.shape = orig_input.shape。与 tiling 侧形成双重校验。Host 端 Tiling 实现
op_host/arch35/max_pool_v3_grad_tiling.cpp中MaxPoolV3GradTilingFunc:GetPlatformInfo获取 UB 大小与 AIV 核数。ValidateAttrs按 data_format 解析 H/W 维做范围与 N/C 维==1 校验(对齐 canndevcheck_param)。ComputePoolParams严格对齐 Golden_compute_pool_params计算 VALID/SAME/CALCULATED(含 ceil_mode)三种模式下的 padTop/padLeft 与 Ho/Wo,并用FloorDiv(对齐 Python//向负无穷取整)覆盖空 tensor 场景,Ho/Wo 钳位max(.,0)。overlap = (sh<kh)||(sw<kw),设 overlapMode。totalOutputPos=N*C*Ho*Wo,物理核均分→下限保护 1024→对齐 32→反推needCoreNum;校验不超 INT32_MAX。totalInputElements*sizeof(float)(FP32 累加缓冲),加系统 workspace。FillTilingData写入 16 字段结构体;SetBlockDim/SetScheduleMode(1)(两条路径均需 SyncAll)/SetTilingKey(仅编码 overlapMode,dtype 由DTYPE_宏实例化)/SetLocalMemorySize(ubSize - 128KB DCache)。SIMT 设备端 Kernel
op_kernel/arch35/max_pool_v3_grad_simt.h采用 512 线程(constexpr THREAD_NUM=512,__launch_bounds__同一常量)的统一 SIMT 架构,纯 GM 计算模式依赖 128KB DCache 加速随机访问,通过if constexpr (OVERLAP_MODE)编译期分发非重叠/重叠路径:ProcessOneOutputPos:用 int64_t outPos 索引,分解 nC/hoWo/ho/wo/n/c,ComputeOutAddr/ComputeInAddr按 inputFormat 区分 NCHW(outAddr=outPos,inAddr=nCHW+hiW+wi)与 NHWC(C 交错);维护matched布尔标志按 (ki,kj) 行优先扫描,首匹配位置经RouteGrad路由梯度,越界位置直接 continue(等价 pad_value=-65500 方案且更高效)。==运算符天然对齐 IEEE 754(NaN 不传播、Inf 传播、+0==-0)。RouteGrad:非重叠直接outGrad[inAddr]=grad[outAddr];重叠将 grad 提升为 FP32 后asc_atomic_add(&wsGrad[inAddr], gradVal)。MaxPoolV3GradSimt:Grid-Stride 循环逐输出位置处理。InitPass(末尾__builtin_cce_dcci刷新 VF DCache)/ConvertPass(FP32→输出 dtype,half 走__float2half)。op_kernel/arch35/max_pool_v3_grad.cpp为入口:REGISTER_TILING_DEFAULT+GET_TILING_DATA_WITH_STRUCT,通过AscendC::GetUserWorkspace(workspace)获取用户区指针(跳过 ascend950 16MB 系统保留区,避免 overlap 路径踩踏),按 schMode 分发Process<DTYPE_ORIG_INPUT, 0/1>。双路径调度:非重叠 = InitPass(out_grad)→SyncAll→直接写;重叠 = InitPass(ws FP32)→SyncAll→atomicAdd ws→SyncAll→ConvertPass(ws→out_grad)。
Tiling 数据结构与编译配置
op_kernel/arch35/max_pool_v3_grad_tiling_data.h:MaxPoolV3GradTilingData共 16 字段(totalOutputPos、totalInputElements 为 int64_t;needCoreNum/HoWo/Wo/C/HW/W/inputFormat/kh/kw/sh/sw/padTop/padLeft 为 int64_t, overlapMode为int32_t),不含 threadNum(编译期常量)。op_kernel/arch35/max_pool_v3_grad_tiling_key.h:ASCENDC_TPL_ARGS_DECL/ASCENDC_TPL_SEL声明MAX_POOL_V3_GRAD_TPL_MODE_NON_OVERLAP(0)/MAX_POOL_V3_GRAD_TPL_MODE_OVERLAP(1)两种场景模式,AIV_ONLY,dtype 不入 TilingKey。op_host/config/ascend950/max_pool_v3_grad_binary.json:dtype/format 与二进制映射配置。关键设计决策与精度
__builtin_cce_dcciDCache 一致性三重必要)。Golden 参考实现与单元测试
tests/assets/golden.py:以golden_max_pool_v3_grad_direct(纯 numpy,first-wins,完整支持 VALID/SAME/CALCULATED+ceil_mode,中间提升 float32 累加,Ho/Wo 钳位 0)为默认验收标杆;另提供golden_max_pool_v3_grad_tf(tf.raw_ops.MaxPoolGrad,仅 SAME/VALID+ceil_mode=False 交叉验证);TTK 入口max_pool_v3_grad_golden。tests/ut/op_host/test_max_pool_v3_grad_infershape.cpp:覆盖 dtype 一致性、4D 校验、ksize/strides 约束、grad shape 校验等场景。tests/ut/op_host/arch35/test_max_pool_v3_grad_tiling.cpp:覆盖 NCHW/NHWC、VALID/SAME/CALCULATED、ceil_mode、重叠/非重叠、空 tensor、无效 dtype 等场景。