ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

CANN ops-nn MseLossGrad 算子完全指南:均方误差反向传播的数学原理、参数解析与 aclnn 调用实战

CANN ops-nn MseLossGrad 算子完全指南:均方误差反向传播的数学原理、参数解析与 aclnn 调用实战 CANN ops-nn MseLossGrad 算子完全指南均方误差反向传播的数学原理、参数解析与 aclnn 调用实战【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nnMseLossGrad 是 CANN ops-nn 神经网络算子库中均方误差损失函数MSE Loss的反向传播算子它根据前向损失对输入的梯度 dout将梯度继续回传到预测值 predict 与标签 label是训练场景中 MSE 损失反向链路上的关键一环。本文以 loss/mse_loss_grad/README.md 为骨架结合仓库内的算子定义、shape 推导、tiling、kernel 实现与测试用例系统讲解该算子的产品支持范围、数学原理、参数语义、两段式 aclnn 接口调用方法与底层实现机制帮助读者在 Atlas 系列产品上正确使用并深入理解 MseLossGrad。算子定位MSE 损失的梯度回传MseLossGrad 是前向算子 aclnnMseLoss均方误差函数的反向传播实现。前向算子计算预测值与目标值之间逐元素的平方误差并依据 reduction 做缩减而 MseLossGrad 则完成该损失对输入的自变量求导结果与上游梯度的复合运算将dout上游传来的损失梯度乘以2 * (predict - label)MSE 对 predict 的导数从而把梯度正确地回传到网络前层。该算子在 CANN ops-nn 仓库中位于 loss/mse_loss_grad/整体由算子原型定义op_graph、Host 侧实现op_host、Kernel 侧实现op_kernel以及示例与测试examples、tests构成并通过两段式 aclnn 接口aclnnMseLossBackward对外提供调用能力。产品支持情况根据 loss/mse_loss_grad/README.md 的产品支持表MseLossGrad 在如下产品上可用产品是否支持Ascend 950PR/Ascend 950DT√Atlas A3 训练系列产品/Atlas A3 推理系列产品√Atlas A2 训练系列产品/Atlas A2 推理系列产品√Atlas 200I/500 A2 推理产品×Atlas 推理系列产品√Atlas 训练系列产品√需要注意的是产品支持矩阵与算子 kernel 的适配架构相关。从源码看本算子的 AI Core 配置目前注册在ascend950架构上见 mse_loss_grad_def.cpp 中的this-AICore().AddConfig(ascend950, aicoreConfig)kernel 与 tiling 也位于 arch35 目录op_kernel/arch35/这些产品差异决定了算子在具体硬件上的可执行性使用前请先核对目标设备型号。数学原理与计算公式MseLossGrad 的输入输出对应关系为predict对应公式中的 xlabel对应公式中的 ydout对应公式中的 grad上游梯度输出y为回传梯度。当reduction为mean时计算公式为$$ MselossBackward(grad, x, y) grad \times (x - y) \times 2 / x.numel() $$其中x.numel()表示predict中的元素个数。该系数来源于 MSE 损失mean模式下前向为sum((x-y)^2)/N对 x 求导得到2*(x-y)/N再乘以上游梯度grad。当reduction不是mean即none或sum时计算公式为$$ MselossBackward(grad, x, y) grad \times (x - y) \times 2 $$需要特别指出尽管 README 与接口文档中 reduction 名义上支持none/mean/sum三种模式但从实现来看none与sum在反向传播中对应同一计算式系数均为 2只有mean模式额外引入了1/numel的缩放这一点可以从仓库的 tiling 实现中得到印证。从源码验证公式实现上述公式在仓库多处实现中保持一致形成了文档—Host 侧—Kernel 侧—测试的全链路印证Host 侧系数计算在 mse_loss_grad_tiling_arch35.cpp 的CalcReduceMeanCof中reduction 字符串被映射为整数{none: 0, sum: 1, mean: 2}reduceMeanCof初始为2.0f仅当 reduction 为mean时遍历 predict 的 storage shape 累乘得到元素数dimVal再计算reduceMeanCof 2.0 / dimVal并拒绝空 tensor 的 mean 计算Input predict cannot be an empty tensor when the attribute reduction is mean。Kernel 侧计算流水在 mse_loss_grad_dag.h 的算子 DAG 中计算链路清晰可见Sub(predict, label)→Muls(系数)乘上2或2/numel→Mul(dout)→ 类型还原后写出与文档公式逐项对应。测试参考实现executor_aclnnMseLossBackward.py 用 PyTorch 表达为gradOutput * (self - target) * 2mean时再除以self.numel()golden.py 中则先做三输入 broadcast再按cof 2.0或cof 2.0 / numel计算最后还原原始 dtype与公式完全一致。参数说明MseLossGrad 共包含 3 个输入、1 个属性、1 个输出参数语义如下表所示参数名输入/输出/属性描述数据类型数据格式predict输入公式中的输入 x前向预测值BFLOAT16、FLOAT16、FLOATNDlabel输入公式中的输入 y真实标签BFLOAT16、FLOAT16、FLOATNDdout输入公式中的输入 grad上游传入的梯度值BFLOAT16、FLOAT16、FLOATNDreduction属性指定损失函数的计算方式支持 0(none) | 1(mean) | 2(sum)。none 表示不应用缩减mean 表示输出的总和将除以输入中的元素数sum 表示输出将被求和INT64NDy输出公式中的输出 MselossBackward回传梯度BFLOAT16、FLOAT16、FLOATND参数语义的源码级细化reduction 的默认值与形态差异在算子原型 mse_loss_grad_proto.h 中reduction 被定义为String类型且默认值为mean而在 aclnn 接口 aclnnMseLossBackward.md 中它表现为int64_t参数0/1/2两种形态可通过{none:0, sum:1, mean:2}的映射对应。README 参数表中 reduction 的数据类型标注为 INT64即对应 aclnn 接口层的形态。dtype 与 format 约束算子定义 mse_loss_grad_def.cpp 中三个输入与一个输出统一限定为{BF16, FLOAT16, FLOAT}与FORMAT_ND二进制配置文件 mse_loss_grad_binary.json 进一步为 bfloat16、float16、float32 三种 dtype 分别注册了对应的 kernel binaryshape 以-2表示动态维度。输出 shape 与 dtype 推导infershape 实现 mse_loss_grad_infershape.cpp 将 predict、label、dout 三个输入 shape 做 broadcast 作为输出 shape输出数据类型直接继承 predict 的数据类型InferDataType4MseLossGrad。对应的单测见 test_mse_loss_grad_infershape.cpp。约束说明与校验逻辑README 中约束说明一节标注为无但这指的是算子本身没有额外限制条件。实际调用时仍受接口层的参数校验约束从 aclnnMseLossBackward.md 与 tiling 源码可以归纳出以下校验规则数据类型一致性predict、label、dout、输出 y 四者的数据类型必须一致。tiling 的GetShapeAttrsInfo中会逐一校验predict 与 label、predict 与 dout、predict 与 y的 dtype 是否相同不一致时直接返回GRAPH_FAILED并给出明确报错信息。Broadcast 关系gradOutput、self、target 三者 shape 必须满足 broadcast 关系输出 y 的 shape 为三者 broadcast 之后的结果。这也是该算子天然支持预测值、标签、梯度形状不同的输入的前提kernel 侧通过 broadcast_sch.h 的BroadcastSch完成广播调度。维度范围输入 shape 需在 1~8 维范围内超过 8 维报错。reduction 取值必须在 0~2 之间none/mean/sumtiling 的CalcReduceMeanCof会对非法字符串报错reduction ... none, mean or sum。空 tensor 行为gradOutput、self、target 任一为空 tensor 时输出为空 tensor直接返回成功但若 reduction 为 mean 且 predict 为空 tensor则因无法计算元素数而报错。确定性计算aclnnMseLossBackward 默认确定性实现。两段式接口调用aclnnMseLossBackwardMseLossGrad 算子对外提供两段式 aclnn 接口。根据 aclnnMseLossBackward.md必须先调用第一段接口aclnnMseLossBackwardGetWorkspaceSize获取计算所需 workspace 大小与执行器再调用第二段接口aclnnMseLossBackward真正执行计算aclnnStatus aclnnMseLossBackwardGetWorkspaceSize( const aclTensor* gradOutput, const aclTensor* self, const aclTensor* target, int64_t reduction, aclTensor* out, uint64_t* workspaceSize, aclOpExecutor** executor)aclnnStatus aclnnMseLossBackward( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)GetWorkspaceSize 阶段参数说明参数名输入/输出描述使用说明数据类型数据格式维度(shape)非连续TensorgradOutput(aclTensor*)输入梯度反向输入公式中的输入 gradshape 需与 self、target 满足 broadcast 关系与 self 保持一致ND1-8√self(aclTensor*)输入输入张量公式中的输入 xshape 需与 gradOutput、target 满足 broadcast 关系FLOAT、FLOAT16、BFLOAT16ND1-8√target(aclTensor*)输入真实标签公式中的输入 yshape 需与 gradOutput、self 满足 broadcast 关系与 self 保持一致ND1-8√reduction(int64_t)输入指定应用到输出的缩减方式支持 0(none) | 1(mean) | 2(sum)INT64--√out(aclTensor*)输出输出的回传梯度shape 为三者 broadcast 后的结果dtype 为 self 可推导的数据类型FLOAT、FLOAT16、BFLOAT16ND1-8√workspaceSize(uint64_t*)输出返回需在 Device 侧申请的 workspace 大小-----executor(aclOpExecutor**)输出返回 op 执行器包含算子计算流程-----第一段接口的典型错误码第一段接口完成入参校验出现以下场景时报错返回值错误码描述ACLNN_ERR_PARAM_NULLPTR161001传入的 gradOutput、self、target 或 out 是空指针时ACLNN_ERR_PARAM_INVALID161002self 的数据类型不在支持范围内161002gradOutput、target 的数据类型与 self 不同161002gradOutput、self、target 的 shape 无法做 broadcast161002三者 broadcast 后的 shape 与 out 的 shape 不一致161002reduction 值不在 0~2 范围之内161002gradOutput、self 或 target 的 shape 超过 8 维第二段接口参数说明参数名输入/输出描述workspace输入在 Device 侧申请的 workspace 内存地址workspaceSize输入Device 侧 workspace 大小由第一段接口获取executor输入op 执行器包含算子计算流程stream输入指定执行任务的 Stream完整调用示例与运行流程仓库在 examples/test_aclnn_mse_loss_backward.cpp 中提供了完整的 aclnn 调用样例核心流程如下#include acl/acl.h #include aclnnop/aclnn_mse_loss_backward.h int main() { // 1. 固定写法device/stream 初始化 int32_t deviceId 0; aclrtStream stream; auto ret Init(deviceId, stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(Init acl failed. ERROR: %d\n, ret); return ret); // 2. 构造输入与输出 std::vectorint64_t gradOutputShape {2, 2}; std::vectorint64_t selfShape {2, 2}; std::vectorint64_t targetShape {2, 2}; std::vectorint64_t outShape {2, 2}; std::vectorfloat gradOutputHostData {0, 1, 2, 3}; std::vectorfloat selfHostData {0, 1, 2, 3}; std::vectorfloat targetHostData {1, 1, 1, 1}; std::vectorfloat outHostData {0, 0, 0, 0}; // 通过 aclCreateTensor 创建 gradOutput/self/target/out 四个 aclTensorND 格式连续 strides // ... // 3. 创建 reduction1 表示 mean int64_t reduction 1; // 4. 两段式接口调用 uint64_t workspaceSize 0; aclOpExecutor* executor; ret aclnnMseLossBackwardGetWorkspaceSize(gradOutput, self, target, reduction, out, workspaceSize, executor); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(GetWorkspaceSize failed. ERROR: %d\n, ret); return ret); void* workspaceAddr nullptr; if (workspaceSize 0) { ret aclrtMalloc(workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(allocate workspace failed. ERROR: %d\n, ret); return ret); } ret aclnnMseLossBackward(workspaceAddr, workspaceSize, executor, stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclnnMseLossBackward failed. ERROR: %d\n, ret); return ret); // 5. 固定写法同步等待任务执行结束 ret aclrtSynchronizeStream(stream); // 6. 将 device 侧结果拷回 host 并打印 std::vectorfloat resultData(4, 0); ret aclrtMemcpy(resultData.data(), resultData.size() * sizeof(resultData[0]), outDeviceAddr, size * sizeof(resultData[0]), ACL_MEMCPY_DEVICE_TO_HOST); for (int64_t i 0; i size; i) { LOG_PRINT(result[%ld] is: %f\n, i, resultData[i]); } // 7. 释放 aclTensor、device 内存与 stream 资源 aclDestroyTensor(gradOutput); aclDestroyTensor(self); aclDestroyTensor(target); aclDestroyTensor(out); aclrtFree(gradOutputDeviceAddr); aclrtFree(selfDeviceAddr); aclrtFree(targetDeviceAddr); aclrtFree(outDeviceAddr); if (workspaceSize 0) { aclrtFree(workspaceAddr); } aclrtDestroyStream(stream); aclrtResetDevice(deviceId); aclFinalize(); return 0; }以示例数据gradOutput{0,1,2,3}、self{0,1,2,3}、target{1,1,1,1}、reduction1(mean)为例(self - target) {-1, 0, 1, 2}乘 2 得{-2, 0, 2, 4}除以numel4得{-0.5, 0, 0.5, 1}再逐元素乘gradOutput得到最终输出{0, 0, 1, 3}。这与 golden.py 的参考计算逻辑完全一致可用于快速自检调用结果。源码级实现原理Host 侧定义、推导与 tiling算子定义mse_loss_grad_def.cpp通过OpDef注册算子MseLossGrad声明 3 个 REQUIRED 输入、1 个 REQUIRED 输出与可选属性 reduction默认mean并配置 AICore 动态编译、动态 rank、动态 shape 支持能力。Shape/DataType 推导mse_loss_grad_infershape.cpp调用通用BroadcastShape工具对三个输入 shape 求广播结果作为输出 shape输出 dtype 直接继承 predict。Tiling 计算mse_loss_grad_tiling_arch35.cpptiling 阶段完成三件关键工作校验四个张量 dtype 一致判断 dout 是否为标量doutIsScalar从而选择不同的计算 DAG 分支依据 reduction 计算reduceMeanCof标量通过SetScalar注入到 tiling 数据中。 之后按 predict 的 dtypeFLOAT16/BF16/FLOAT分别实例化BroadcastBaseTiling完成广播调度 tiling并生成对应的tilingKey。此外还通过TilingPrepare获取 AI Vector 核数与 UB 内存大小等平台信息。Kernel 侧DAG 算子图与广播调度Kernel 入口mse_loss_grad.cpp根据doutIsScalar选择两类算子图并统一交给Ops::Base::BroadcastSch调度执行张量 dout 分支MseLossGradDagmse_loss_grad_dag.h三个输入分别CopyInBrc带广播拷贝进入随后全部 Cast 到 float 精度参与中间计算Sub(predict, label)→Muls(系数)→Mul(dout)→Cast回原精度 →CopyOut。中间精度统一使用 float可以避免 FP16/BF16 在累乘过程中的精度损失。标量 dout 分支MseLossGradScalarDag差别仅在于 dout 通过Vec::Duplicate将标量广播为向量参与乘运算。两个 DAG 均使用MemOptCfgMemLevel::LEVEL_2配置 L2 级内存优化并由BroadcastSch依据 tiling 阶段生成的调度模式schMode驱动执行这解释了该算子对 broadcast 输入的原生支持能力。测试与验证体系仓库为 MseLossGrad 提供了从算子级到 API 级的完整测试覆盖UT 单测tests/ut/op_host/test_mse_loss_grad_infershape.cpp 验证了 infershape如{182,4}输入广播后输出为 2 维与 infer datatype 的正确性test_mse_loss_grad_tiling.cpp 覆盖 arch35 平台的 tiling 逻辑。ST 用例tests/st/aclnnMseLossBackward/ 提供atk_aclnnMseLossBackward.json测试配置与executor_aclnnMseLossBackward.py参考实现后者以 PyTorch 语法给出期望输出用于端到端比对 aclnn 接口的计算结果。Golden 参考tests/assets/golden.py 定义mse_loss_grad_golden完整模拟了三输入 broadcast → 提升精度计算 → mean 时除以元素数 → 还原 dtype的全过程可作为理解算子语义的权威参考实现。Kernel ST 矩阵tests/st/arch35/ttk_kernel_mse_loss_grad_st.csv 记录了 arch35 上 kernel 级测试的用例矩阵。小结MseLossGrad 是 MSE 损失训练闭环中不可或缺的反向算子文档层面给出了完整的产品支持矩阵、数学公式与参数语义源码层面则揭示了 broadcast 广播计算、float 中间精度、mean 系数预计算、dout 标量/张量双分支等实现细节。开发者在实际使用时只需遵循先 GetWorkspaceSize 后执行的两段式 aclnn 调用模式保证四个张量 dtype 一致、shape 满足 broadcast 关系、reduction 取值合法即可在支持的 Atlas 产品上正确完成 MSE 损失的梯度回传。【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nn创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表