已合并
fix(rms_norm): add tail>0 guard for level1 increment and multi-level reduce in stage2 #5771
fix(rms_norm): add tail>0 guard for level1 increment and multi-level reduce in stage2 #5771
已合并
raoliang_sac创建于 6月5日
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 MasterLoop264 // 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_);
@@ -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 MasterLoop316 // 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++) {