已关闭
[Bug] rms_norm_regbase_split_d.h: tail==0 时 level1 和 ComputeMultiLevelReduce 无条件执行导致 rstd 计算错误 #3158
raoliang_sac创建于 6月5日关闭于 6月13日
Rraoliang_sac
6月5日 关联了pull request:fix(rms_norm): add tail>0 guard for level1 increment and multi-level reduce in stage2
6月5日 关联了pull request:fix(rms_norm): add tail>0 guard for level1 increment and multi-level reduce in stage2
6月6日 将 raoliang_sac 设为负责人
6月13日 关闭了 issue
6月13日 添加了label:resolved
问题描述
在
norm/rms_norm/op_kernel/arch35/rms_norm_regbase_split_d.h的ComputeFormer函数中,stage2 处理非整尾块逻辑时,level1 += 1和ComputeMultiLevelReduce在 line 182-183 无条件执行,缺少if (tail > 0)保护。触发条件
当
colTail恰好被ubFactor整除时(tail = colTail % ubFactor == 0),stage2 的两个 if/else-if 分支均不进入,ComputeSum不执行,level1Local[level1]无有效数据。影响
level1 += 1无条件执行,向归约树写入一个空槽位(初始化为 0.0)ComputeMultiLevelReduce使用错误的 level1 计数执行,导致 level2/level3 读取位置错位ComputeMultiLevelRstd计算出的 rstd 值不正确对比参考
norm/add_rms_norm_quant/op_kernel/arch35/add_rms_norm_quant_regbase_split_reduce.h的 line 288 已正确使用if (tail > 0)包裹了相同逻辑。修复方案
将 line 182-183 用
if (tail > 0)包裹:if (tail > 0) { level1 += 1; ComputeMultiLevelReduce(level1Local, level2Local, level3Local, level1, level2, level3); }