已关闭
[Requirement|需求建议]: add fused_cross_entropy_loss_with_max_sum for ascend950 #4201
AlfengYuan创建于  7月21日关闭于  7月22日
AlfengYuan
AlfengYuan成员
7月21日 创建

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)的规约,本算子负责:

loss[b]      = log(sum_exp_logits[b]) - predicted_logits[b]
softmax[b,v] = exp(vocab_parallel_logits[b,v] - logits_max[b]) / sum_exp_logits[b]

关键点:行规约结果(max / sum_exp)是输入而不是本算子计算的,因此本算子是纯 elementwise 计算——这是后面 tiling 能按 v 维自由切核的数学基础。

Origin(信息来源)

华为海思昇腾

Benefit / Necessity (价值/作用)

该算子是loss算子中的一种融合算子,910b上有实现, ascend950上暂时没有

Design(设计方案)

两条执行路径

路径 触发 输出 用途
完整路径(tilingKey=0) vocab_parallel_logits 提供 loss + softmax 训练前向需要完整 softmax
省显存路径(tilingKey=1) vocab_parallel_logits 缺省(None) 仅 loss 超大词表下避免 softmax 物化,省显存

接口定义(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]]
  • 标量输入(max/sum_exp/predicted)与两个输出恒为 fp32;dtype 变化仅存在于 vocab

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 的优势:

  • 寄存器级计算,避免 UB 中间往返;
  • 掩码(predicate)控制天然处理非对齐尾块,无需 padding 搬运;
  • fp16/bf16 → fp32 的 DIST_UNPACK_B16 升精度加载 + Cast,一条指令链完成。
2.2 单二进制 + 运行时分发(为什么不用 DTYPE_* 宏多二进制)

常见做法(如 lin_space、rotary 等算子)是每种 dtype 组合出一个二进制,编译期 -DDTYPE_X= 注入,运行时按输入 dtype 的 key 选二进制。该方案要求 dtype 变化的输入全是 REQUIRED,而本算子 vocab 是 OPTIONAL(省显存路径不存在):

  • 若按 vocab dtype 出 3 个二进制,vocab 缺省时运行时 key 中没有 vocab dtype → 无法区分 3 个二进制;
  • 实测佐证:950 binary.json 只声明了一条 bf16 的 bin 记录,但 fp32/fp16 vocab 的 319 条用例全部正常执行——证明 optional 输入 dtype 不参与运行时选择 key。

因此采用单二进制(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_FP32 64 一个向量寄存器容纳的 fp32 数(256B/4B)
A_PER_LOOP 8 每个 UB 行块处理的行数
SCALAR_BUF_BYTES 256 标量行向量缓冲(64×4B)
3.2 完整路径流程(ProcessFull → ProcessRowTile)
对每个 8 行行块:
  1. 搬入 max / sum_exp / predicted 标量行向量(≤8 floats,3 个深度1队列)
  2. ComputeLossInvSum: 单发 MicroAPI(log, sub, div 求 inv_sum),写 loss、inv_sum
     - v切分时仅 vPartIdx==0 的核写 loss GM(其余核自算 inv_sum 自用)
  3. v维循环(本核的 v 分片 [vBegin, vEnd)):
     vocabQue(深度2双缓冲) 搬入 aRows×vCur  →  ComputeSoftmaxTile  →  softmaxQue(深度2) 搬出
     流水:tile k+1 的 MTE2 搬入与 tile k 的 V 计算 / MTE3 搬出重叠
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)

两个演进点(相对最初实现):

  1. 掩码修复(重要 bugfix):初版 UpdateMask(count) 的 count 在内层循环外固定为总长,末尾轮次多余 lane 也被计算;改为按迭代计算 curCount,再演进为现在的满/尾拆分;
  2. 掩码消减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 用乘倒数:kernel 每行先算 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 精度预期;
  • 升精度计算:fp16/bf16 vocab 先经 DIST_UNPACK_B16 + Cast 升 fp32 再计算,全链路 fp32,无半精度累加误差;
  • 数值实测(CPP 零点实验,见 6.4):loss maxAbs 1.5e-7(isclose(1e-3,1e-3) 零违例)。
3.6 对齐契约(正确性依赖)
  • vPerLoopelementsNumbervChunk 均按 64(=VL_FP32)对齐 → 满块迭代的全宽非掩码 load 最多读到行内 padding,绝不越出 UB
  • GM 侧 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 推导
scalarReserve = 5×SCALAR_BUF_BYTES(256) + 2048 = 3328   // 标量队列+预留
完整路径 perColBytes = A_PER_LOOP(8) × DB(2) × (dtypeSize + 4)
  fp32: 128B/列 → vPerLoop = FloorAlign((ubSize−3328)/128, 64) = 1920
  fp16/bf16: 96B/列 → vPerLoop = FloorAlign((ubSize−3328)/96, 64) = 2560
省显存 elementsNumber = FloorAlign((ubSize−2048)/(3队列×2×4B), 64) = 10496

(上表数值对应 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                        // ≤ coreNum

kernel 侧 blockIdx 二维解码:rowBlockIdx = blockIdx / vCoresvPartIdx = blockIdx % vCores,v-loop 只跑 [vPartIdx×vChunk, min(+vChunk, vLen));loss 仅 vPartIdx==0 写出(其余核仍需自算 inv_sum,8 浮点成本可忽略)。

边界与守卫:

  • vCores=1 时(bt ≥ coreNum 或 v ≤ vPerLoop)vChunk = vLen,blockIdx 解码退化为原一维行切,行为与切分前逐字节一致
  • vChunk 对齐取整可能使末尾个别分片核的 vBegin ≥ vLen:此时 v-loop 次数 ≤0 自动跳过(空转核,仅参与标量计算),不影响正确性;
  • 省显存路径只有 bt 一维,恒 vCores=1,完全不受本优化影响;
  • 行维 former/latter 切分与 vCores 正交:SplitRows(bt, rowCores, totalCores) 在 v 切分时 rowCores=bt(每核 1 行),总核数 rowCores×vCores 不超 coreNum。
likedislike
AlfengYuanAlfengYuan成员
7月21日 添加了label:requirement
yuning_chenyuning_chen成员
7月21日 将 alfengyuan 设为负责人
AlfengYuanAlfengYuan成员
7月21日 修改了issue 的描述
AlfengYuanAlfengYuan成员
7月21日 修改了issue 的描述
AlfengYuanAlfengYuan成员
7月21日 修改了issue 的描述
CANN-robotCANN-robot成员
7月22日 关闭了 issue
AlfengYuanAlfengYuan成员
7月22日 issue状态由 进行中 改变为 已完成
CANN-robotCANN-robot成员
7月22日 添加了label:Accepted
CANN-robotCANN-robot成员
7月22日 添加了label:resolved