
CANN ops-nn 算子 SigmoidFocalLossGrad 深度解析Focal Loss 反向梯度计算原理与 GEIR 调用实践【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nnSigmoidFocalLossGrad 是 CANN ops-nn 开源算子库中负责 Sigmoid Focal Loss 反向传播的算子它为pred前向 logits计算 raw selector 反向梯度服务于目标检测等使用 Focal Loss 训练的反向过程。本文将以 loss/sigmoid_focal_loss_grad/README.md 为骨架结合仓库内的 IR 定义、Host Tiling、Kernel 实现、示例与测试代码完整讲解其数学模型、参数约束、产品支持情况以及如何在 GE 图模式下调用并验证该算子帮助开发者在 CANN 环境中正确使用与二次开发这一算子。背景与定位Focal Loss 通过引入聚焦参数gamma与类别平衡权重alpha降低易分类样本的损失贡献使模型训练聚焦于难分样本是目标检测等类别不平衡场景的常用损失函数。SigmoidFocalLossGrad 是该损失函数在反向传播阶段的梯度计算算子给定前向 logitspred与上游梯度dout计算pred对应的梯度grad从而驱动参数更新。从 IR 定义 的注释可以确认该算子与 MMCV 的 SigmoidFocalLoss 算子兼容Compatible with the MMCV operator SigmoidFocalLoss也就是说它对齐的是 MMCV 框架中 sigmoid_focal_loss 的反向语义。需要特别注意的是其target输入是前向 dense target 的补集即 raw backward selector业务标签语义为 0 或 1不能直接当作高层框架中的普通正类标签使用这一点在后续参数说明中会详细展开。产品支持情况根据 README.md 的产品支持矩阵该算子的支持情况如下产品是否支持Ascend 950PR/Ascend 950DT√Atlas A3 训练系列产品/Atlas A3 推理系列产品√Atlas A2 训练系列产品/Atlas A2 推理系列产品√Atlas 200I/500 A2 推理产品×Atlas 推理系列产品×Atlas 训练系列产品√其中Ascend 950 系列的算子配置可以在 算子定义文件 中看到AICore().AddConfig(ascend950, aiCoreConfig)明确注册了 ascend950 的 AICore 配置并且 Kernel 代码也位于op_kernel/arch35/目录下arch35 即 Ascend 950 系列对应的架构目录。功能说明与数学模型公式定义SigmoidFocalLossGrad 计算 Sigmoid Focal Loss 对前向 logitspred的 raw selector 反向梯度。记p sigmoid(pred)、t target且weight缺省时按全 1 处理则梯度计算分为正类项与负类项$$ dpos \alpha\gamma p(1-p)^\gamma\log(p)-\alpha(1-p)^{\gamma1} $$$$ dneg \gamma(\alpha-1)p^\gamma(1-p)\log(1-p)(1-\alpha)p^{\gamma1} $$最终梯度由 raw selectortarget在两项之间选择并乘以上游梯度dout与逐元素权重weight$$ grad (dpos(1-t)dneg,t)\times dout\times weight $$reductionmean时逐元素结果再除以pred的元素数即1/(N*C)系数sum和none不缩放。无论哪种 reduction输出 shape 均与pred一致这一点从 示例代码 中输出 shape 校验gradShape.GetDim(0) ! kShape[0]、GetDim(1) ! kShape[1]即为失败可以得到印证。与 MMCV / PyTorch 语义的对应关系Golden 参考实现 中给出了与上述公式逐句对应的 PyTorch 计算序列可作为理解算子的标准答案probs torch.sigmoid(pred_f)probs_nadd 1 - probsdpos_front -alpha * exp((gamma1) * log(clamp(1-p)))dpos_back exp(gamma*log(clamp(1-p))) * log(clamp(p)) * p * gamma*alpha二者相加得dposdneg_front exp(gamma*log(clamp(p))) * (1-p) * log(clamp(1-p)) * gamma*(alpha-1)dneg_back (1-alpha) * exp((gamma1)*log(clamp(p)))二者相加得dneg最终result (dpos*(1-target) dneg*target) * weight * doutmean 时再乘1/元素数系数。可以看到公式中的p(1-p)^γ、p^γ等幂次项在实现中均通过exp(log(·))的形式计算并对底数做clamp(min1.17549435e-38)即 float32 最小正规数CONST_FP_MIN的下限保护避免对 0 取对数产生 NaN/Inf。这也是 Kernel 端精度设计的核心出发点之一。参数说明算子接口由 IR 定义 通过REG_OP(SigmoidFocalLossGrad)声明各参数细节如下参数名输入/输出/属性描述数据类型数据格式pred输入前向 logits对应公式中 pred。FLOAT16、FLOATNDtarget输入raw backward selector对应公式中 t。INT32NDdout输入上游梯度对应公式中 dout。FLOAT16、FLOATNDweight可选输入逐元素样本权重对应公式中 weight缺省时等价于全 1。FLOAT16、FLOATNDgrad输出pred的梯度数据类型和 shape 均跟随pred。FLOAT16、FLOATNDalpha可选属性类别平衡权重对应公式中 alpha默认值为 0.25。FLOAT-gamma可选属性聚焦指数对应公式中 gamma默认值为 2.0。FLOAT-reduction可选属性结果缩放方式可取 mean、sum 或 none默认值为 mean。STRING-参数在源码中的注册方式IR 层protoIR 定义 中pred、target、dout为必选输入weight通过.OPTIONAL_INPUT声明输出grad限定为DT_FLOAT16、DT_FLOAT三个属性通过.ATTR(alpha, Float, 0.25)、.ATTR(gamma, Float, 2.0)、.ATTR(reduction, String, mean)声明默认值分别对应 0.25、2.0 与 mean。OpDef 层Host算子定义 中对每个输入/输出声明了 8 组(DataType, Format)组合全部限定为FORMAT_ND覆盖 FLOAT16/FLOAT 的不同组合weight为可选输入OPTIONAL并显式设置DynamicShapeSupportFlag(true)、DynamicRankSupportFlag(false)。这与 README 中Ascend 950 图模式支持动态 Shape、不支持动态 Rank的约束完全一致。Tiling 层Tiling 实现 在ReadAttributes中读取alpha、gammaattrs-GetFloat(0)、GetFloat(1)与reductionattrs-GetStr(2)并校验alpha、gamma必须为有限数std::isfinitereduction只接受mean/sum/none否则直接GRAPH_FAILED报错退出。Graph 层图数据类型推断 将输出grad的数据类型直接继承自predSetOutputDataType(kGradIndex, predDtype)从图构建层面保证了grad数据类型跟随pred的约束。参数组合的典型约束从 Tiling 实现 的ValidateTensorContract与ValidateShapesAndFlatten函数可以整理出如下硬性约束违反即报错pred、target、dout、grad与存在时的weight必须是非空二维 ND TensorGetDimNum() 2且两维均大于 0shape 完全相同SameShape逐维比较不支持广播target只支持 INT32targetDesc-GetDataType() ! ge::DT_INT32即失败pred、dout必须是 FLOAT16 或 FLOATweight存在时也必须是 FLOAT16 或 FLOATgrad的数据类型必须与pred相同所有 Tensor 存储格式必须为 NDFORMAT_NDalpha、gamma必须为有限数reduction只能为 mean、sum 或 none若传入weight其 shape 也必须与pred完全一致。关于target的业务语义README 明确说明其业务标签语义为 0 或 1且IR 和 Kernel 不在运行前逐元素检查 Device 侧数值——也就是说框架侧只做 shape/type 契约校验数值合法性由调用方自行保证。这一点与 Golden 参考实现 中raw selector semantics的注释一致target扮演的是在dpos正类项与dneg负类项之间做选择的开关角色。约束说明综合 README 与源码校验逻辑该算子的使用约束可归纳为pred、target、dout、grad与存在时的weight必须是非空二维 ND Tensorshape 完全相同不支持广播。target只支持 INT32业务标签语义为 0 或 1IR 和 Kernel 不在运行前逐元素检查 Device 侧数值。pred、dout和weight可分别使用 FLOAT16 或 FLOATgrad的数据类型必须与pred相同。alpha和gamma必须为有限数reduction只能为 mean、sum 或 none。Ascend 950 图模式支持动态 Shape-1 未知维不支持动态 Rank-2 未知秩。关于动态 Shape 支持动态 Shape 示例 的注释给出了非常明确的工程实践由于 OpDef 只声明了DynamicShapeSupportFlag(true)而未声明动态 Rank因此示例只添加[-1, -1]动态形状图并在同一 Session 中运行三个合法的非空二维 shape同时探测[-2]作为预期被拒绝的用例——当前 Tiling 契约只接受 rank 2动态 Rank 会按预期失败。这为使用该算子做动态 Shape 推理的开发者提供了直接参考。调用说明README 给出了 GE 图模式的调用方式调用方式样例代码说明GE图模式test_geir_sigmoid_focal_loss_grad.cpp通过 SigmoidFocalLossGrad IR 定义 构建算子图。GE 图模式调用示例解析test_geir_sigmoid_focal_loss_grad.cpp 是一个完整的 GEGraph Engine图构建与执行示例核心流程如下准备 Host TensorMakeHostTensor根据ge::Shape({2, 3})、FORMAT_ND与指定数据类型DT_FLOAT/DT_INT32构造ge::TensorkElementCount 6即一个 2×3 的二维张量。构建图创建 4 个ge::op::Data节点pred、target、dout、weight索引分别为 0~3再创建ge::op::SigmoidFocalLossGrad(sigmoid_focal_loss_grad)算子节点通过set_input_pred/set_input_target/set_input_dout/set_input_weight连接输入通过set_attr_alpha(0.25F)、set_attr_gamma(2.0F)、set_attr_reduction(mean)设置属性最后graph.SetInputs(graphInputs).SetOutputs(graphOutputs)完成图定义。初始化并运行GEInitialize时设置全局选项ge.exec.deviceId0、ge.graphRunMode1、ge.jit_compile0创建ge::SessionAddGraph(kGraphId, graph)后以 4 个 Host Tensor 为输入调用RunGraph。校验输出示例中的输入数据为pred {-1.25, -1.0, -0.75, -0.5, -0.25, 0.0}FLOATtarget {0, 1, 0, 1, 0, 1}INT32dout {-1.5, -1.0, -0.5, 0.0, 0.5, 1.0}FLOATweight {1.0, 1.0, 1.0, 1.0, 1.0, 1.0}FLOAT全 1运行后校验输出grad的 shape 必须为 (2, 3)、所有元素必须为有限值std::isfinite并通过%.8f打印梯度结果。这段代码可以直接作为在 CANN 环境中以 GE 图模式验证 SigmoidFocalLossGrad 算子的最小可运行模板。动态 Shape 调用对于需要动态 Shape 的场景参考 动态 Shape 示例构建[-1, -1]动态形状的图在同一 Session 中依次运行多个合法的非空二维 shape例如三个不同的 2D shape每次运行都校验grad的 shape、dtype、字节数与数值CPU 公式对照容忍度atol1.0e-3、rtol2.0e-3并额外探测[-2]动态 Rank 按预期被拒绝。这为算子级动态 Shape 回归提供了完整范式。Kernel 实现原理与数值精度设计算子整体结构该算子的目录结构体现了 CANN 算子开发的标准分层op_graphIR 定义算子原型 图推断数据类型推断op_host算子定义OpDef 注册 InferShape复用逐元素基础 InferShape即Ops::Base::InferShape4Elewise Tilingop_kernelarch35 下的 Kernel 入口、Kernel 模板实现、模板实例化头 与两个结构体头文件examples静态 Shape 与动态 Shape 两个 GEIR 示例testsHost Tiling 单测、Kernel 单测与 Golden 参考。Kernel 入口与模板分派Kernel 入口 是一个以hasWeight为模板参数的__global__ __aicore__函数通过REGISTER_NONE_TILING与GET_TILING_DATA_WITH_STRUCT获取 Tiling 结构体 中的数据。入口逻辑校验pred/target/dout/grad指针非空hasWeight1时还校验weight非空校验dim0 0、blockNum 0、ubFormer 0依据weightDtype0 表示 FLOAT161 表示 FLOAT32实例化不同模板参数组合的SigmoidFocalLossGradKernelPredT, DoutT, WeightT, HAS_WEIGHT并执行。模板实例化方式可以在 模板参数声明 中看到hasWeight作为 8-bit 模板参数取值范围 0~1由编译期展开。Tiling 数据与多核切分Tiling 结构体 描述了 Host 侧传给 Kernel 的关键切分信息dim0展平后的元素总数 N×CcoreNum/blockNum目标 AIV 核数与实际启动 block 数blockFormer每个前序 block 分配的元素量ubFormer每个前序 UB tile 处理的元素量ubLoopOfFormerBlock/ubTailOfFormerBlock、ubLoopOfTailBlock/ubTailOfTailBlock前序 block 与尾 block 各自的循环次数与有效尾元素数alpha/gamma运行期聚焦参数reduceMeanCoef运行期 reduction 系数mean 时为1/(N*C)sum/none 时为 1weightDtype0 为 FLOAT161 为 FLOAT32无 weight 时忽略。Tiling 实现 中定义了kMinBitsPerCore 32768单核最小负载 32K bit等切分常量将 2D shape 展平为一维后按核数与 UB 容量做两层切分核间切分 核内 UB tile 循环并把reduceMeanCoef在 Host 侧预先算好下发给 KernelKernel 侧不再做除法仅做一次乘法缩放。Kernel 内通过GetBlockIdx()获取 block ID用GetBlockWork计算每个 block 的负载与循环参数再以MTE2→V→MTE3的事件同步机制驱动流水线RunTileLoop中显式获取并等待MTE2_V、V_MTE3、MTE3_MTE2三类硬件事件。数值精度设计防下溢、防消减误差这是该算子 Kernel 实现中最值得关注的技术细节。Kernel 模板实现 通过一组 SIMD 向量函数__simd_vf__实现了高精度计算路径稳定的 Sigmoid 求值StableSigmoidExpVf/StableSigmoidVfp sigmoid(pred)在pred为较大负数时1-p接近 1 而p接近 0。实现分别计算p 1/(1exp(-|x|))与1-p 1/(1exp(|x|))两个分支再按x的符号Select选择从而避免直接计算sigmoid(x)时对极端值产生的精度损失随后用kLogEps 1.17549435e-38ffloat32 最小正规数对p、1-p做Maxs下限钳制保证后续Ln的输入恒大于 0。高精度幂运算ComputePowers公式中的p^γ、p^{γ1}、(1-p)^γ、(1-p)^{γ1}全部通过AscendC::Powerfloat, false, kHighPrecisionPowerConfig计算其中kHighPrecisionPowerConfig {PowerAlgo::DOUBLE_FLOAT_TECH}双精度浮点技术即底数拆分为双精度格式参与指数运算以提升精度。DetectClampedPowerBaseVf先检测是否存在被钳制的底数最小底数 kLogEps若有则用CorrectClampedPowerVf以exp(ln(base)*exponent)的等价形式做回退修正确保钳制引入的偏差被纠正。消减误差补偿DPosVfdpos αγ·p·(1-p)^γ·log(p) − α·(1-p)^{γ1}中当p接近 1 时lhs与rhs两个大数相减会产生灾难性消减catastrophic cancellation。实现通过代数恒等式(1-p)^γ·p·log(p) (1-p)^{γ1}/(1-p)·p·log(p)重构出corrected表达式当|original| 0.015625 * max(|lhs|, |rhs|)即相对消减达到 1/64时自动切换到重构路径并附加了maxTerm 最小正规数、q与qGamma1均正常非下溢、有限等保护条件只有所有条件满足才使用修正值否则回退到原始表达式——这是一套仅在必要时修正的保守精度策略。DNegVf同样对1-alpha 1即alpha 0的退化情形做了特殊处理直接取p^{γ1}避免出现(1-alpha)与 1 的差造成的不精确。加权与缩放ScaleStoreVfgrad raw * weight * dout * reduceMeanCoef。当weight与dout的乘积下溢为次正规数abs(weighted) 1.17549435e-38且raw、weight均非零时通过raw / (1/weight)、preciseWeighted / (1/dout)的倒数除法路径恢复精度recoverSubnormal掩码控制随后统一乘reduceMeanCoef最后 FLOAT16 输出路径使用kF32ToB16的CAST_RINT舍入模式并做 32-bit 打包存储。这一整套实现说明该算子不只是简单地套公式而是在 sigmoid 求值、幂运算、大数相减、加权下溢等多个数值薄弱点上做了系统性防护与 Golden 参考实现 中的clamp(minCONST_FP_MIN)下限保护策略相互印证。测试与验证该算子配套了完整的测试体系覆盖 Host 侧与 Kernel 侧Kernel 单测使用 gtest 框架通过TestDtypeProfilePROFILE0~11 共 12 个 dtype 组合剖面覆盖PredT × DoutT × WeightT × hasWeight的各种组合FLOAT16/FLOAT 混合、有无 weight 等。由于 CANN 9.2 的 tikicpulib 无法执行 arch35 RegBase 指令与高级 Power 实现测试通过宏替换把内存拷贝与计算分派替换为标量等价实现SigmoidFocalLossGradUtDataCopyPad、SigmoidFocalLossGradUtPower用std::pow标量计算但真实 Kernel 源码仍被 include 进来所有 dtype/profile/控制流模板都被这些测试实例化从而在 CPU 上完成逻辑等价验证。Tiling 单测验证 Host 侧 Tiling 计算的正确性。Golden 参考提供了sigmoid_focal_loss_grad_golden与SigmoidFocalLossGradKernelSpec其中golden按冻结的 TBE 运算顺序在 PyTorch 中复算先升 float32 计算最后再转回pred的原始载体类型输出third_party {torch: _TorchSigmoidFocalLossGrad}注册了 Torch 小算子组合作为第三方参考tolerance _L0FLOAT16 与 FLOAT32 均采用cross_check交叉校验标准、L0级别——即与第三方参考实现做交叉核对属于最严格的验证等级。同时该文件还注明算子仅注册了 raw 名称sigmoid_focal_loss_grad因为它是通过 Kernel 直接调用与 GEIR 两条路径测试的。总结与使用建议SigmoidFocalLossGrad 是 CANN ops-nn 中一个接口简洁、实现精细的反向损失算子数学本质grad (dpos·(1-target) dneg·target) · dout · weightmean 模式再除N×C输出与pred同 shape 同 dtype调用前必查所有 Tensor 必须为非空二维 ND、shape 完全一致target用 INT32 且语义为 0/1 的 raw selector前向 dense target 的补集alpha/gamma必须有限reduction仅 mean/sum/none直接可用的模板静态 Shape 用 test_geir_sigmoid_focal_loss_grad.cpp动态 Shape 用 动态示例两者都内含完整的 GE 图构建、Session 运行与输出校验逻辑精度设计值得借鉴稳定 sigmoid 分支求值、双精度幂、消减误差代数重构、次正规数恢复为同类含log、pow与多乘积项的损失梯度算子提供了高精度的实现范式验证体系完备Kernel 单测覆盖 12 种 dtype 剖面Golden 以 PyTorch 交叉校验L0 标准兜底可作为算子合入质量的门禁参考。对于要在目标检测训练中接入 Focal Loss 反向、或希望深入理解 CANN 算子数值精度工程实践的开发者本算子及其配套源码是一份结构完整、可直接复用的技术样本。【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nn创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考