已合并
fix(rms_norm): add tail>0 guard for level1 increment and multi-level reduce in stage2 #5771
raoliang_sac创建于 6月5日
fix(rms_norm): add tail>0 guard for level1 increment and multi-level reduce in stage2 #5771
已合并
共 3 个文件变更+14-8
| @@ -256,9 +256,11 @@ private: | |||
| 256 | offset += tail; | 256 | offset += tail; |
| 257 | workspaceOffset += tail; | 257 | workspaceOffset += tail; |
| 258 | } | 258 | } |
| 259 | - RmsNorm::ComputeSum(level1Local, tempLocal, level1Idx, SUM_COUNT); | 259 | + if (tail > 0) { |
| 260 | - level1Idx += 1; | 260 | + RmsNorm::ComputeSum(level1Local, tempLocal, level1Idx, SUM_COUNT); |
| 261 | - RmsNorm::ComputeMultiLevelReduce(level1Local, level2Local, level3Local, level1Idx, level2Idx, level3Idx); | 261 | + level1Idx += 1; |
| 262 | + RmsNorm::ComputeMultiLevelReduce(level1Local, level2Local, level3Local, level1Idx, level2Idx, level3Idx); | ||
| 263 | + } | ||
| 262 | // Stage3: Cal MasterLoop | 264 | // Stage3: Cal MasterLoop |
| 263 | for (uint32_t repeat = 0; repeat < masterLoop; repeat++) { | 265 | for (uint32_t repeat = 0; repeat < masterLoop; repeat++) { |
| 264 | ComputeFormerHandle(level1Local, offset, level1Idx, workspaceOffset, powerSplit_, powerSplit_); | 266 | ComputeFormerHandle(level1Local, offset, level1Idx, workspaceOffset, powerSplit_, powerSplit_); |
Mnorm/add_rms_norm_dynamic_quant/op_kernel/arch35/add_rms_norm_dynamic_quant_regbase_split_reduce.h+5-3
| @@ -308,9 +308,11 @@ private: | |||
| 308 | offset += tail; | 308 | offset += tail; |
| 309 | workSpaceOffset += tail; | 309 | workSpaceOffset += tail; |
| 310 | } | 310 | } |
| 311 | - RmsNorm::ComputeSum(level1Local, tempLocal, level1, SUM_COUNT); | 311 | + if (tail > 0) { |
| 312 | - level1 += 1; | 312 | + RmsNorm::ComputeSum(level1Local, tempLocal, level1, SUM_COUNT); |
| 313 | - RmsNorm::ComputeMultiLevelReduce(level1Local, level2Local, level3Local, level1, level2, level3); | 313 | + level1 += 1; |
| 314 | + RmsNorm::ComputeMultiLevelReduce(level1Local, level2Local, level3Local, level1, level2, level3); | ||
| 315 | + } | ||
| 314 | // Stage3: Cal MasterLoop | 316 | // Stage3: Cal MasterLoop |
| 315 | for (uint32_t repeat = 0; repeat < masterLoop; repeat++) { | 317 | for (uint32_t repeat = 0; repeat < masterLoop; repeat++) { |
| 316 | ComputeFormerHandle(level1Local, offset, level1, workSpaceOffset, powerSplit_, powerSplit_); | 318 | ComputeFormerHandle(level1Local, offset, level1, workSpaceOffset, powerSplit_, powerSplit_); |
| @@ -179,8 +179,10 @@ private: | |||
| 179 | 179 | ||
| 180 | ComputeSum(level1Local, tempLocal, level1, SUM_COUNT); | 180 | ComputeSum(level1Local, tempLocal, level1, SUM_COUNT); |
| 181 | } | 181 | } |
| 182 | - level1 += 1; | 182 | + if (tail > 0) { |
| 183 | - ComputeMultiLevelReduce(level1Local, level2Local, level3Local, level1, level2, level3); | 183 | + level1 += 1; |
| 184 | + ComputeMultiLevelReduce(level1Local, level2Local, level3Local, level1, level2, level3); | ||
| 185 | + } | ||
| 184 | loop += 1; | 186 | loop += 1; |
| 185 | // stage3: 处理主块逻辑 | 187 | // stage3: 处理主块逻辑 |
| 186 | for (uint32_t repeat = 0; repeat < masterLoop; repeat++) { | 188 | for (uint32_t repeat = 0; repeat < masterLoop; repeat++) { |