Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.
AcosGradV2 算子 A2(Ascend910B / DAV_2201)AscendC 实现,计算 Acos(反余弦)算子的反向梯度。
算子使用场景:反向训练/推理中,已知前向 Acos 的输入 y 与上游梯度 dy,求对原始输入的梯度 z。 CANN 当前未提供同名官方 aclnnAcosGrad 单算子,用户走 torch_npu 的 torch.acos 反向时会被拆分为多算子链(Acos 前向重算 + Mul×2 + Neg×2 + Rsqrt + Adds)。本算子将其融合为单个 kernel,减少多次基础算子的 HBM 来回搬运与启动开销,[1024,1024] 下相对该反向链平均加速约 3.3x。
aclnn 两段式直调(aclnnAcosGradV2GetWorkspaceSize + aclnnAcosGradV2),registry-invoke 工程。
ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BFLOAT16
shape 推导:输出 z 与输入 y 同 shape 同 dtype;tiling 采用多核 + UB 两级切分(UB 大小由平台信息动态获取,按 dtype 的字节/对齐因子推导 ubFormer,并记录核内 loop/tail);参数校验覆盖空指针、y/dy 同 shape 同 dtype、dtype 合法性。
公式 z = -dy / sqrt(1 - y²)。FP32 直接计算;FP16/BF16 先 Cast 到 FP32 计算再 Cast 回原类型。计算链:Mul(y²) → Muls(-y²) → Adds(1-y²) → Sqrt → Muls(-dy) → Div。
Atlas A2 训练/推理系列产品(Ascend910B / DAV_2201),compute unit = ascend910b,tiling/kernel 目录 = arch32。
💡 备注:关联 PR https://gitcode.com/cann/ops-math/pull/4371 ,用于追踪。
Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.
一、背景信息 (必填)
AcosGradV2 算子 A2(Ascend910B / DAV_2201)AscendC 实现,计算 Acos(反余弦)算子的反向梯度。
二、价值/作用 (必填)
算子使用场景:反向训练/推理中,已知前向 Acos 的输入 y 与上游梯度 dy,求对原始输入的梯度 z。
CANN 当前未提供同名官方 aclnnAcosGrad 单算子,用户走 torch_npu 的 torch.acos 反向时会被拆分为多算子链(Acos 前向重算 + Mul×2 + Neg×2 + Rsqrt + Adds)。本算子将其融合为单个 kernel,减少多次基础算子的 HBM 来回搬运与启动开销,[1024,1024] 下相对该反向链平均加速约 3.3x。
三、设计方案 (必填)
3.1 使能方式
aclnn 两段式直调(aclnnAcosGradV2GetWorkspaceSize + aclnnAcosGradV2),registry-invoke 工程。
3.2 总体设计
3.2.1 算子支持的数据类型
ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BFLOAT16
3.2.2 host 侧设计
shape 推导:输出 z 与输入 y 同 shape 同 dtype;tiling 采用多核 + UB 两级切分(UB 大小由平台信息动态获取,按 dtype 的字节/对齐因子推导 ubFormer,并记录核内 loop/tail);参数校验覆盖空指针、y/dy 同 shape 同 dtype、dtype 合法性。
3.2.3 kernel 侧设计
公式 z = -dy / sqrt(1 - y²)。FP32 直接计算;FP16/BF16 先 Cast 到 FP32 计算再 Cast 回原类型。计算链:Mul(y²) → Muls(-y²) → Adds(1-y²) → Sqrt → Muls(-dy) → Div。
3.3 支持硬件
Atlas A2 训练/推理系列产品(Ascend910B / DAV_2201),compute unit = ascend910b,tiling/kernel 目录 = arch32。
3.4 算子约束限制
💡 备注:关联 PR https://gitcode.com/cann/ops-math/pull/4371 ,用于追踪。