CANN ops-nn 算子解读SigmoidCrossEntropyWithLogits 的数学原理、接口约束与 NPU 实现【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nnSigmoidCrossEntropyWithLogits 是 CANN ops-nn 神经网络算子库loss/sigmoid_cross_entropy_with_logits中用于二分类场景的逐元素损失算子它直接在 NPU 上计算 logits 与标签之间的 Sigmoid Cross Entropy 损失与 TensorFlow 同名算子语义兼容。本文以该算子官方文档为主体结合仓库中的算子原型、Shape 推导、Tiling 与 Kernel 源码系统讲解其数学公式、参数与约束、调用方式以及从图构图到 AIV 向量核执行的全链路实现原理帮助读者在 CANN 环境下正确使用与二次开发该算子。一、算子功能与数学原理1.1 功能定位SigmoidCrossEntropyWithLogits 算子接收两个输入预测值predictlogits与标签值target输出逐元素计算的 Sigmoid Cross Entropy 损失loss。它常用于多标签二分类、多任务分类等场景对每个 logit 独立应用 Sigmoid 激活再计算与对应标签的交叉熵最后对整批结果做归约即可得到训练损失。从算子原型注册文件 op_graph/sigmoid_cross_entropy_with_logits_proto.h 的注释可以确认其框架兼容性Compatible with TensorFlow operator SigmoidCrossEntropyWithLogits.也就是说在 CANN 图中可直接承接 TensorFlow 模型导出图中的同名算子实现无损迁移。1.2 计算公式官方文档给出的计算公式为$$ \text{loss} \max(\text{predict}, 0) - \text{predict} \times \text{target} \log(1 \exp(-|\text{predict}|)) $$其中predict为输入的 logits 值target为标签值loss为计算得到的损失值。该公式在数学上等价于先对 logits 求 Sigmoid 概率 $p \frac{1}{1e^{-x}}$再计算二元交叉熵 $-(t \cdot \log p (1-t) \cdot \log(1-p))$但采用了数值稳定的写法用 $\max(x,0)$ 与 $\log(1e^{-|x|})$ 组合避免 $e^{-x}$ 在 $x$ 为较大负数时下溢、在 $x$ 为较大正数时上溢的问题这也是 TensorFlow 官方实现的同款稳定化处理。1.3 Kernel 层的指令级实现数值稳定的公式并不是只在文档层面描述在 NPU Kernel 中得到了逐条指令的落实。核心计算位于 op_kernel/arch35/sigmoid_cross_entropy_with_logits_dag.h 的CalcSigmoidCrossEntropyWithLogits结构中其向量计算序列为Reg::Maxs(vregMaxPredict, vregPredict, 0.0f)—— 计算 $\max(\text{predict}, 0)$Reg::Abs(vregAbsPredict, vregPredict)→Reg::Neg(vregNegAbsPredict, ...)→Reg::Exp(vregExpNegAbs, ...)—— 计算 $e^{-|\text{predict}|}$Reg::Duplicate(vregOneAddExp, 1.0f)Reg::Add(...)—— 计算 $1 e^{-|\text{predict}|}$Reg::Log(vregLog, vregOneAddExp)—— 计算 $\log(1 e^{-|\text{predict}|})$Reg::Mul(vregMulPredictTarget, vregPredict, vregTarget)—— 计算 $\text{predict} \times \text{target}$Reg::Sub(vregOutput, vregMaxPredict, vregMulPredictTarget)Reg::Add(vregOutput, vregOutput, vregLog)—— 组合出最终损失。可见文档公式与 Kernel 指令序列一一对应逐元素、逐向量地完成稳定化损失计算。二、产品支持情况官方文档中的产品支持矩阵如下产品是否支持Ascend 950PR / Ascend 950DT√Atlas A3 训练系列产品 / Atlas A3 推理系列产品√Atlas A2 训练系列产品 / Atlas A2 推理系列产品√Atlas 200I/500 A2 推理产品×Atlas 推理系列产品√Atlas 训练系列产品√从源码结构看该算子的 Tiling 与 Kernel 均位于arch35目录op_host/arch35 与 op_kernel/arch35算子定义中通过OpAICoreConfig为ascend950平台添加了 AICore 配置见 op_host/sigmoid_cross_entropy_with_logits_def.cpp与文档中 Ascend 950 系列产品支持情况一致。文档明确标注 Atlas 200I/500 A2 推理产品不支持部署前请务必核对目标产品型号。三、参数说明官方文档参数表如下参数名输入/输出/属性描述数据类型数据格式predict输入预测值 logits。FLOAT16、FLOAT、BFLOAT16NDtarget输入标签值。FLOAT16、FLOAT、BFLOAT16NDloss输出损失值。shape 和输入 predict 一致。FLOAT16、FLOAT、BFLOAT16ND3.1 原型层的类型与格式约束上述参数表在算子原型中有完全对应的硬性约束op_graph/sigmoid_cross_entropy_with_logits_proto.h 中的注册信息为REG_OP(SigmoidCrossEntropyWithLogits) .INPUT(predict, TensorType({DT_FLOAT16, DT_FLOAT, DT_BF16})) .INPUT(target, TensorType({DT_FLOAT16, DT_FLOAT, DT_BF16})) .OUTPUT(loss, TensorType({DT_FLOAT16, DT_FLOAT, DT_BF16})) .OP_END_FACTORY_REG(SigmoidCrossEntropyWithLogits)在算子定义文件 op_host/sigmoid_cross_entropy_with_logits_def.cpp 中两个输入与一个输出均被声明为REQUIRED必选数据类型限定为{DT_FLOAT16, DT_FLOAT, DT_BF16}三种数据格式限定为{FORMAT_ND, FORMAT_ND, FORMAT_ND}即仅支持 ND 格式。3.2 动态能力配置同一文件中的 AICore 配置还揭示了算子的动态特性sigmoid_cross_entropy_with_logits_def.cppOpAICoreConfig aicoreConfig; aicoreConfig.DynamicCompileStaticFlag(true) .DynamicRankSupportFlag(true) .DynamicShapeSupportFlag(true) .PrecisionReduceFlag(false); this-AICore().AddConfig(ascend950, aicoreConfig);DynamicCompileStaticFlag(true)支持编译期静态确定部分信息后的动态编译DynamicRankSupportFlag(true)支持动态 Rank维度数量在运行期确定DynamicShapeSupportFlag(true)支持动态 ShapePrecisionReduceFlag(false)不做精度降级处理FLOAT32 输入按原生精度计算。这解释了算子为何能以shape: [-2]的形式出现在二进制配置中见下文并支持运行期任意维度的张量。3.3 输出与输入的一致性推导输出loss的 Shape 和数据类型并非独立指定而是由 Shape 推导InferShape逻辑从输入推导而来。op_host/sigmoid_cross_entropy_with_logits_infershape.cpp 中的实现为static ge::graphStatus InferShapeForSigmoidCrossEntropyWithLogits(gert::InferShapeContext* context) { ge::graphStatus ret Ops::Base::InferShape4Elewise(context); ... } static graphStatus InferDataTypeForSigmoidCrossEntropyWithLogits(gert::InferDataTypeContext* context) { context-SetOutputDataType(LOSS_INDEX, context-GetInputDataType(PREDICT_INDEX)); return GRAPH_SUCCESS; }Shape 推导复用InferShape4Elewise逐元素通用推导工具保证输出与输入同 Shape数据类型推导直接将输出类型设置为predict输入的类型保证loss与predict数据类型一致。这一点也被单元测试覆盖tests/ut/op_host/test_sigmoid_cross_entropy_with_logits_infershape.cpp 中分别以二维形状{96, 256}和一维形状{1024}构造输入输出并断言推导成功。四、约束说明官方文档明确了两条约束且这些约束在 Tiling 源码中有对应的强制校验逻辑op_host/arch35/sigmoid_cross_entropy_with_logits_tiling.cpppredict 和 target 必须具有相同的数据类型和形状。CalcInputDtype()校验输入类型必须是 FLOAT16/BF16/FLOAT 三者之一且target的 dtype 必须与predict相同CheckShape()校验target的存储 Shape 与predict一致且输出loss的存储 Shape 也与predict一致同时拒绝空张量shape size 为 0 时报错。支持 FLOAT16、FLOAT、BFLOAT16 数据类型。除 dtype 校验外CalcOutputDtype()还会校验输出loss的 dtype 与输入predict相同。从实现看标量输入Shape 为 0 维会被EnsureNotScalar统一视为{1}形状参与一致性比较属于对边界情况的兼容处理。五、调用说明官方文档给出的调用方式为图模式通过算子 IR 构图方式调用即引用 op_graph/sigmoid_cross_entropy_with_logits_proto.h 中的算子声明完成构图。仓库提供了完整可编译的构图示例 examples/test_geir_sigmoid_cross_entropy_with_logits.cpp其核心构图流程如下// 1. 创建算子实例 auto sigmoidCrossEntropyWithLogits1 op::SigmoidCrossEntropyWithLogits(sigmoidCrossEntropyWithLogits1); // 2. 构造 predict 输入占位符 形状 {4, 2}ND 格式DT_FLOAT std::vectorint64_t predictShape {4, 2}; auto placeholder1 op::Data(placeholder1).set_attr_index(0); TensorDesc placeholder1_desc TensorDesc(ge::Shape(predictShape), FORMAT_ND, inDtype); sigmoidCrossEntropyWithLogits1.set_input_predict(placeholder1); // 3. 构造 target 输入形状同样为 {4, 2} auto placeholder2 op::Data(placeholder2).set_attr_index(1); sigmoidCrossEntropyWithLogits1.set_input_target(placeholder2); // 4. 声明输出并建图运行 TensorDesc loss_desc TensorDesc(ge::Shape(predictShape), FORMAT_ND, inDtype); sigmoidCrossEntropyWithLogits1.update_output_desc_loss(loss_desc); graph.SetInputs(inputs).SetOutputs(outputs); session-AddGraph(graph_id, graph, graph_options); session-RunGraph(graph_id, input, output);构图要点可归纳为使用op::Data创建输入占位符并通过set_attr_index指定输入索引通过set_input_predict/set_input_target将两个占位符挂到算子输入通过update_output_desc_loss预先声明输出张量描述Shape、格式、类型将算子加入Graph设置图输入输出后经Session::AddGraph与Session::RunGraph完成 NPU 上的执行示例中ge.exec.deviceId指定运行设备ge.graphRunMode指定图运行模式运行结果会 dump 为tc_ge_irrun_test_0008_npu_input_*.bin/_output_*.bin文件并打印每个输出元素。示例中使用的输入形状为{4, 2}、数据类型DT_FLOAT与算子支持的数据类型完全匹配由于算子支持动态 Shape实际业务中可将{4, 2}替换为任意形状的 predict/target 对。六、NPU 执行链路从 Tiling 到 Kernel6.1 Tiling 准备与数据切分在图编译阶段Tiling 负责根据输入形状、平台资源决定 Kernel 的切分参数。op_host/arch35/sigmoid_cross_entropy_with_logits_tiling.cpp 的执行链路为TilingPrepareForSigmoidCrossEntropyWithLogits通过platform_ascendc::PlatformAscendC获取 AIV 核数GetCoreNumAiv与 UB 内存大小GetCoreMemSize(CoreMemType::UB, ...)写入编译信息Tiling4SigmoidCrossEntropyWithLogits→RunTiling()依次执行CalcInputDtype输入类型校验→CalcOutputDtype输出类型校验→CheckShape形状一致性校验→DoElewiseTilingDoElewiseTiling按 dtype 选择对应的模板参数FP16/BF16/FP32调用通用逐元素 Tiling 工具ElewiseBaseTiling的DoTiling生成EleBaseTilingData16B类型的 Tiling 数据并设置tilingKey与blockDim核数申请固定大小 workspaceconst size_t ASCEND_WORKSPACE 16777216;即 16 MB 的工作内存。6.2 Kernel 分派与模板实例化Kernel 入口 op_kernel/sigmoid_cross_entropy_with_logits.cpp 通过编译期常量dType做三路分派template uint64_t dType __global__ __aicore__ void sigmoid_cross_entropy_with_logits(GM_ADDR predict, GM_ADDR target, GM_ADDR loss, GM_ADDR workspace, GM_ADDR tiling) { REGISTER_TILING_DEFAULT(EleBaseTilingData16B); GET_TILING_DATA_PTR_WITH_STRUCT(EleBaseTilingData16B, tilingData, tiling); KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY); if constexpr (dType TPL_FP16) { ElementwiseSch16B1, SigmoidCrossEntropyWithLogitsDagWithCasthalf::OpDag sch(tilingData); sch.Init(predict, target, loss); sch.Process(); } else if constexpr (dType TPL_BF16) { ElementwiseSch16B1, SigmoidCrossEntropyWithLogitsDagWithCastbfloat16_t::OpDag sch(tilingData); ... } else if constexpr (dType TPL_FP32) { ElementwiseSch16B1, SigmoidCrossEntropyWithLogitsDagNoCastfloat::OpDag sch(tilingData); ... } }任务类型为KERNEL_TYPE_AIV_ONLY即由 AIV 向量核执行逐元素计算三种 dtype 的模板参数在 op_kernel/arch35/sigmoid_cross_entropy_with_logits_struct.h 中定义TPL_FP161、TPL_BF162、TPL_FP323调度器统一使用ElementwiseSch16B1, OpDag16 字节对齐的逐元素调度器实际计算逻辑封装在 DAG 描述中。6.3 DAG 计算图FP16/BF16 的精度提升策略op_kernel/arch35/sigmoid_cross_entropy_with_logits_dag.h 中定义了两套计算 DAGSigmoidCrossEntropyWithLogitsDagNoCastTFP32 专用输入输出直接以 float 参与计算不做类型转换SigmoidCrossEntropyWithLogitsDagWithCastU, T floatFP16/BF16 专用先将 16 位输入Cast提升为 float 计算CAST_MODE_NONE完成损失计算后再以CAST_MODE_RINT四舍五入模式将结果 Cast 回 16 位输出。这一设计意味着FP16/BF16 输入在 NPU 上实际以float32 中间精度完成 $\max$、$\exp$、$\log$ 等运算只在边界处做一次 Cast可有效降低 16 位浮点逐元素累积带来的精度损失FP32 输入则直接原生计算无精度降级与PrecisionReduceFlag(false)配置相互印证。两份 DAG 的MemCfg均为MemOptCfgMemLevel::LEVEL_2即允许中间数据驻留二级缓存优化访存。6.4 二进制配置与运行期分派平台侧以 JSON 形式维护算子的二进制库分派配置op_host/config/ascend950/sigmoid_cross_entropy_with_logits_binary.json 中为三种数据类型分别登记了二进制文件二进制文件predict dtypetarget dtypeloss dtype格式SigmoidCrossEntropyWithLogits_a1b2c3d4e5f67890float16float16float16NDSigmoidCrossEntropyWithLogits_b2c3d4e5f6789012float32float32float32NDSigmoidCrossEntropyWithLogits_c3d4e5f678901234bfloat16bfloat16bfloat16ND配置中所有输入输出均声明为shape: [-2]动态 Rank 通配paramType: required运行期根据实际输入 dtype 选择对应二进制与算子支持的动态 Shape 能力相配套。七、测试与验证仓库为算子提供了两级验证手段InferShape 单元测试tests/ut/op_host/test_sigmoid_cross_entropy_with_logits_infershape.cpp 使用gert::InferShapeContextFaker构造{96, 256}、{1024}等形状验证算子实现已注册、InferShape 函数非空且推导返回GRAPH_SUCCESSTiling 单元测试tests/ut/op_host/arch35/test_sigmoid_cross_entropy_with_logits_tiling.cpp 覆盖 Tiling 数据生成逻辑端到端 GEIR 示例examples/test_geir_sigmoid_cross_entropy_with_logits.cpp 完整演示从 GE 初始化、构图、建图到RunGraph并在 NPU 上执行、导出输入输出 bin 文件的全流程。对于开发者而言验证算子的标准步骤为使用示例构图代码构造 predict/target 输入 → 运行图得到 loss 输出 → 与公式 $\max(x,0) - x \cdot t \log(1e^{-|x|})$ 的逐元素计算结果比对由于 Kernel 指令序列与公式严格对应逐元素误差应在所选数据类型的精度范围内。八、小结SigmoidCrossEntropyWithLogits 是 ops-nn 中一个典型的小而完整的逐元素损失算子文档层面给出了清晰的功能、公式、参数与约束源码层面则完整覆盖了算子原型注册、Shape 推导、动态 Tiling、AIV 向量 Kernel、二进制分派配置与多级测试。其核心设计要点可归纳为数值稳定公式以 $\max(x,0) \log(1e^{-|x|}) - x \cdot t$ 形式规避 Sigmoid 交叉熵在极端 logits 下的溢出严格输入约束predict/target/loss 三者 dtype 与 shape 必须一致仅支持 FLOAT16/FLOAT/BFLOAT16 与 ND 格式动态能力完备支持动态 Shape、动态 Rank、动态编译可承接 TensorFlow 导出的同名算子16 位精度增强FP16/BF16 输入以 float32 中间精度计算FP32 原生计算不降级完整工程配套图模式示例、InferShape/Tiling 单测与二进制分派配置一应俱全可直接作为 CANN 算子开发与集成的参考模板。【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nn创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
