aclnnDivs(self 为 tensor、other 为标量的 TrueDiv 路径)在 Ascend950(DAV_3510,RegBase)上,当 out 与 self 同为 bf16/fp16 时,host 侧实现会插入一次冗余的 Cast(self -> fp32),导致算子图上多出 2 个 kernel(正向 Cast 与反向 Cast)。
aclnnDivs
out
self
Cast(self -> fp32)
Aclnn 直调(aclnnDivs / aclnnInplaceDivs)。
aclnnInplaceDivs
self: bf16/fp16/fp32;other: fp16/bf16/fp32/double;out: 与 self 同 dtype(本次优化生效范围)。
在 aclnnDivsGetWorkspaceSize 的倒数乘(Muls)分支中,当 canUseMuls 且 out->GetDataType() == self->GetDataType() 时,跳过 Cast(self -> promoteType),直接以 self 的原 dtype 调用 Muls:
aclnnDivsGetWorkspaceSize
canUseMuls
out->GetDataType() == self->GetDataType()
Cast(self -> promoteType)
const bool canSkipCast = canUseMuls && out->GetDataType() == self->GetDataType(); auto selfCasted = canSkipCast ? selfContiguous : l0op::Cast(selfContiguous, promoteType, uniqueExecutor.get());
依据:Muls 的 bf16/fp16 kernel 内部本身即 bf16/fp16 -> fp32 -> CAST_RINT 回写(见 muls_dag.h 的 MulsOp),与 Cast(self->fp32) + Muls(fp32) 逐位等价。
muls_dag.h
MulsOp
Cast(self->fp32) + Muls(fp32)
约束:out 与 self 不同 dtype 时不可跳过,否则会引入二次舍入,或提前丢范围/精度。
无 kernel 改动,仅 host 侧图优化。
Ascend950(DAV_3510)。
仅当 canUseMuls(RegBase 且 self/other 为支持的浮点类型)且 out 与 self 同 dtype 时生效;其余路径行为不变。
💡 备注(选填)
配套 PR:https://gitcode.com/cann/ops-math/merge_requests/5957
/assign @wangqi_ai
一、背景信息 (必填)
aclnnDivs(self 为 tensor、other 为标量的 TrueDiv 路径)在 Ascend950(DAV_3510,RegBase)上,当out与self同为 bf16/fp16 时,host 侧实现会插入一次冗余的Cast(self -> fp32),导致算子图上多出 2 个 kernel(正向 Cast 与反向 Cast)。二、价值/作用 (必填)
三、设计方案 (必填)
3.1 使能方式(涉及哪些框架:如Aclnn直调、Pytorch训练等)
Aclnn 直调(
aclnnDivs/aclnnInplaceDivs)。3.2 总体设计
3.2.1 算子支持的数据类型
self: bf16/fp16/fp32;other: fp16/bf16/fp32/double;out: 与 self 同 dtype(本次优化生效范围)。
3.2.2 host侧设计
在
aclnnDivsGetWorkspaceSize的倒数乘(Muls)分支中,当canUseMuls且out->GetDataType() == self->GetDataType()时,跳过Cast(self -> promoteType),直接以 self 的原 dtype 调用 Muls:const bool canSkipCast = canUseMuls && out->GetDataType() == self->GetDataType(); auto selfCasted = canSkipCast ? selfContiguous : l0op::Cast(selfContiguous, promoteType, uniqueExecutor.get());依据:Muls 的 bf16/fp16 kernel 内部本身即 bf16/fp16 -> fp32 -> CAST_RINT 回写(见
muls_dag.h的MulsOp),与Cast(self->fp32) + Muls(fp32)逐位等价。约束:out 与 self 不同 dtype 时不可跳过,否则会引入二次舍入,或提前丢范围/精度。
3.2.3 kernel侧设计
无 kernel 改动,仅 host 侧图优化。
3.3 支持硬件
Ascend950(DAV_3510)。
3.4 算子约束限制
仅当
canUseMuls(RegBase 且 self/other 为支持的浮点类型)且 out 与 self 同 dtype 时生效;其余路径行为不变。💡 备注(选填)
配套 PR:https://gitcode.com/cann/ops-math/merge_requests/5957