已关闭
[Requirement|需求建议]: add fused_cross_entropy_loss_with_max_sum for ascend950 #4201
AlfengYuan创建于 7月21日关闭于 7月22日
7月21日 添加了label:requirement
7月21日 将 alfengyuan 设为负责人
7月21日 修改了issue 的描述
7月21日 修改了issue 的描述
7月21日 修改了issue 的描述
7月22日 关闭了 issue
7月22日 issue状态由 进行中 改变为 已完成
7月22日 添加了label:Accepted
7月22日 添加了label:resolved
Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.
Backgroud(背景信息)
add fused_cross_entropy_loss_with_max_sum for ascend950
Megatron 词汇表并行(vocab-parallel)交叉熵被拆成多级融合算子,本算子是最后一段:上游已经完成 logits 的行最大值(
logits_max)与 exp 行和(sum_exp_logits)的规约,本算子负责:关键点:行规约结果(max / sum_exp)是输入而不是本算子计算的,因此本算子是纯 elementwise 计算——这是后面 tiling 能按 v 维自由切核的数学基础。
Origin(信息来源)
华为海思昇腾
Benefit / Necessity (价值/作用)
该算子是loss算子中的一种融合算子,910b上有实现, ascend950上暂时没有
Design(设计方案)
两条执行路径
vocab_parallel_logits提供vocab_parallel_logits缺省(None)接口定义(op_def)
logits_max(fp32,REQUIRED)、sum_exp_logits(fp32,REQUIRED)、predicted_logits(fp32,REQUIRED)、input(fp32,OPTIONAL)、weight(fp32,OPTIONAL)、vocab_parallel_logits(fp32/fp16/bf16,OPTIONAL)loss(fp32,REQUIRED)、softmax_logits(fp32,OPTIONAL)label_smoothing(float,当前仅支持 0,aclnn 层显式拒绝非 0)input/weight为 API 对齐占位参数,kernel 不使用([[maybe_unused]])2. 总体方案
2.1 为什么是 regbase + MicroAPI
arch35(Ascend950)引入 regbase 编程模型:数据从 UB 加载进向量寄存器(VReg,256B = 64 个 fp32),在寄存器上完成计算后写回 UB。对
sub + exp + mul这种 V pipe bound 的 elementwise 算子,MicroAPI(AscendC::MicroAPI::RegTensor+MaskReg)相比传统 LocalTensor API 的优势:DIST_UNPACK_B16升精度加载 +Cast,一条指令链完成。2.2 单二进制 + 运行时分发(为什么不用 DTYPE_* 宏多二进制)
常见做法(如 lin_space、rotary 等算子)是每种 dtype 组合出一个二进制,编译期
-DDTYPE_X=注入,运行时按输入 dtype 的 key 选二进制。该方案要求 dtype 变化的输入全是 REQUIRED,而本算子 vocab 是 OPTIONAL(省显存路径不存在):因此采用单二进制(
FusedCrossEntropyLossWithMaxSum_ALL)+ tiling 字段vocabDtypeId(0=fp32,1=fp16,2=bf16) 运行时分发:if (TILING_KEY_IS(TILINGKEY_FULL)) { // 0:完整路径 if (tilingData.vocabDtypeId == FP32) RegBase<float, true> else if (tilingData.vocabDtypeId == FP16) RegBase<half, true> else RegBase<bfloat16_t, true> } else if (TILING_KEY_IS(TILINGKEY_MEMORY)) { // 1:省显存路径 RegBase<float, false> // 无 vocab,恒 fp32 }tilingKey 只表达"路径"这一个维度,比 910b 实现(tilingKey 0~3 混合路径+dtype)语义更干净。
2.3 soc 分发
op_host/fused_cross_entropy_loss_with_max_sum_tiling.cpp中注册的 tiling 函数首行按Ops::NN::OpTiling::IsRegbaseSocVersion(context)(NpuArch ∈ {3510, 5102})分发:regbase soc → arch35 tiling,否则 → 910b 旧路径。op_kernel/CMakeLists.txt按 soc 选择 kernel 源(910b/910_93 → 旧实现,950 →arch35/)。3. Kernel 实现(
op_kernel/arch35/fused_cross_entropy_loss_with_max_sum_ar.h,424 行)3.1 类与模板
FusedCrossEntropyLossWithMaxSumRegBase<T, fullPath>:T 为 vocab dtype(float/half/bfloat16_t),fullPath区分两条路径,编译期if constexpr消除分支成本。关键常量:
VL_FP32A_PER_LOOPSCALAR_BUF_BYTES3.2 完整路径流程(ProcessFull → ProcessRowTile)
3.3 ComputeSoftmaxTile 的 MicroAPI 核心
for j in rows: // 每行 maxReg = broadcast(maxAddr[j]); // DIST_BRC_B32 广播 invReg = broadcast(invAddr[j]); fullMask = CreateMask<float, ALL>(); // 编译期满掩码,一次生成 for i in fullLoops: // 满块迭代(无尾 lane) load; Sub; Exp; Mul; store (fullMask) // fp16/bf16 经 DIST_UNPACK_B16 + Cast 升 fp32 for i in [fullLoops, totalLoops): // 0/1 次迭代,避免 VF 内 if tailMask = UpdateMask(tailCount); // 仅尾块计算掩码 load; Sub; Exp; Mul; store (tailMask)两个演进点(相对最初实现):
UpdateMask(count)的 count 在内层循环外固定为总长,末尾轮次多余 lane 也被计算;改为按迭代计算curCount,再演进为现在的满/尾拆分;vPerLoop/elementsNumber按 64 对齐,只有最后一个分块的最后迭代存在尾 lane。满块复用CreateMask<ALL>()(编译期生成,零运行时开销),尾块写成 0/1 次 for(避免__VEC_SCOPE__内 if 分支),UpdateMask调用从aRows×innerLoops降到 ≤2 次。3.4 省显存路径(ProcessMemory → ComputeLossChunk)
纯 fp32 一维流:
sumExpQue_/predictedQue_/lossQue_深度 2 双缓冲,按elementsNumber分块log(sum_exp) − predicted。同样的满/尾掩码结构。3.5 精度设计
inv_sum = 1/sum_exp(Div),softmax 用exp(v−max) × inv_sum(Mul),与 golden 的exp(v−max)/sum_exp(除法)存在 ~1ulp 级语义差,实测 softmax maxRel 6.6e-7,符合 fp32 精度预期;DIST_UNPACK_B16+Cast升 fp32 再计算,全链路 fp32,无半精度累加误差;3.6 对齐契约(正确性依赖)
vPerLoop、elementsNumber、vChunk均按 64(=VL_FP32)对齐 → 满块迭代的全宽非掩码 load 最多读到行内 padding,绝不越出 UB;DataCopyPad处理任意长度,CopyInVocabTile/CopyOutSoftmaxTile用 (blockCount=aRows, srcStride=vLen−vCur, dstStride=UB行距差) 的 strided copy 完成行距转换。4. Tiling 实现(
op_host/fused_cross_entropy_loss_with_max_sum_tiling_arch35.cpp,219 行)4.1 UB 预算与 vPerLoop / elementsNumber 推导
(上表数值对应 UB=253952B;其他 arch35 变体由
GetCoreMemSize动态取。)4.2 行切分(SplitRows)
formerCoreNum = bt % rowCores个核各latterRows+1行,其余各latterRows行,±1 均衡。4.3 v 维切核(SplitV,本 PR 的性能主优化)
问题:
bt < coreNum时只有 bt 个核工作。由于 max/sum_exp 是输入,v 维逐列独立、无归约 → 可按 v 切核:// tiling 侧 vCores = min(coreNum / bt, ceil(v / vPerLoop)); // 每核至少一个完整 UB tile 才切 vChunk = CeilAlign(ceil(v / vCores), 64); // 保持对齐契约 rowCores = min(bt, coreNum / vCores); // 行维核数(v切分时退化为 bt,每核 1 行) blockDim = rowCores × vCores // ≤ coreNumkernel 侧
blockIdx二维解码:rowBlockIdx = blockIdx / vCores,vPartIdx = blockIdx % vCores,v-loop 只跑[vPartIdx×vChunk, min(+vChunk, vLen));loss 仅vPartIdx==0写出(其余核仍需自算 inv_sum,8 浮点成本可忽略)。边界与守卫:
vCores=1时(bt ≥ coreNum 或 v ≤ vPerLoop)vChunk = vLen,blockIdx 解码退化为原一维行切,行为与切分前逐字节一致;vBegin ≥ vLen:此时 v-loop 次数 ≤0 自动跳过(空转核,仅参与标量计算),不影响正确性;vCores=1,完全不受本优化影响;SplitRows(bt, rowCores, totalCores)在 v 切分时 rowCores=bt(每核 1 行),总核数rowCores×vCores不超 coreNum。