已合并
fix: 修复 rmsnormquantv2/3 的 GM 内存越界问题 #7338
huanghaihong创建于 7月10日
fix: 修复 rmsnormquantv2/3 的 GM 内存越界问题 #7338
已合并
huanghaihong创建于 7月10日
2 个文件变更+11-4
@@ -352,7 +352,7 @@ private:
352 dataCopyPadExtParamsGamma.paddingValue = 0;352 dataCopyPadExtParamsGamma.paddingValue = 0;
353 DataCopyExtParams copyInParamsGamma;353 DataCopyExtParams copyInParamsGamma;
354 copyInParamsGamma.blockCount = 1;354 copyInParamsGamma.blockCount = 1;
355- copyInParamsGamma.blockLen = xGammaBetaAlign * sizeof(T_X);355+ copyInParamsGamma.blockLen = numR * sizeof(T_X);
356 copyInParamsGamma.srcStride = 0;356 copyInParamsGamma.srcStride = 0;
357 copyInParamsGamma.dstStride = 0;357 copyInParamsGamma.dstStride = 0;
358 DataCopyPad(gammaLocal, gammaGm, copyInParamsGamma, dataCopyPadExtParamsGamma);358 DataCopyPad(gammaLocal, gammaGm, copyInParamsGamma, dataCopyPadExtParamsGamma);
@@ -367,7 +367,11 @@ private:
367 dataCopyPadExtParamsScales.paddingValue = 0;367 dataCopyPadExtParamsScales.paddingValue = 0;
368 DataCopyExtParams copyInParamsScales;368 DataCopyExtParams copyInParamsScales;
369 copyInParamsScales.blockCount = 1;369 copyInParamsScales.blockCount = 1;
370- copyInParamsScales.blockLen = scalesAlign * sizeof(T_SCALES);370+ if (numQ == 1) {
371+ copyInParamsScales.blockLen = sizeof(T_SCALES);
372+ } else {
373+ copyInParamsScales.blockLen = numR * sizeof(T_SCALES);
374+ }
371 copyInParamsScales.srcStride = 0;375 copyInParamsScales.srcStride = 0;
372 copyInParamsScales.dstStride = 0;376 copyInParamsScales.dstStride = 0;
373 DataCopyPad(scales1Local, scales1Gm, copyInParamsScales, dataCopyPadExtParamsScales);377 DataCopyPad(scales1Local, scales1Gm, copyInParamsScales, dataCopyPadExtParamsScales);
@@ -380,7 +384,11 @@ private:
380 dataCopyPadExtParamszeroPoints.paddingValue = 0;384 dataCopyPadExtParamszeroPoints.paddingValue = 0;
381 DataCopyExtParams copyInParamszeroPoints;385 DataCopyExtParams copyInParamszeroPoints;
382 copyInParamszeroPoints.blockCount = 1;386 copyInParamszeroPoints.blockCount = 1;
383- copyInParamszeroPoints.blockLen = zeroPointsAlign * sizeof(T_ZEROPOINTS);387+ if (numQ == 1) {
388+ copyInParamszeroPoints.blockLen = sizeof(T_ZEROPOINTS);
389+ } else {
390+ copyInParamszeroPoints.blockLen = numR * sizeof(T_ZEROPOINTS);
391+ }
atomgit-botatomgit-bot
atomgit-botatomgit-bot7月10日

🔴 Critical

第 385 行声明 DataCopyExtParams copyInParamszeroPoints;,随后第 386 行仅设置了 blockCount = 1blockLen 字段从未被赋值。

由于第 387–391 行的复制粘贴错误将 blockLen 错误地写到了 copyInParamsScales.blockLen(参见上一条 P0 报告),copyInParamszeroPoints.blockLen 实际保持未初始化状态。

第 392–393 行设置了 srcStridedstStride,但 blockLen 仍然是栈上的脏数据。

随后在第 406 行(hasZeroPoints1 路径)和第 413 行(hasZeroPoints2 路径),DataCopyPad 使用这个未初始化的 copyInParamszeroPoints 进行 GM→UB 数据搬移,导致:

  • 拷贝数据量不可预测(栈上的随机值)
  • 可能从 GM 越界读取,触发内存访问错误
  • 可能向 UB 越界写入,破坏其他 LocalTensor 数据

触发条件hasZeroPoints1 == truehasZeroPoints2 == true 时必然触发。

建议:与上一条合并修复:将第 387–391 行的赋值目标从 copyInParamsScales 改为 copyInParamszeroPoints,类型从 T_SCALES 改为 T_ZEROPOINTS。修复后 copyInParamszeroPoints.blockLen 即被正确初始化。

改动建议
391
+ if (numQ == 1) {
392
+ copyInParamszeroPoints.blockLen = sizeof(T_ZEROPOINTS);
393
+ } else {
394
+ copyInParamszeroPoints.blockLen = numR * sizeof(T_ZEROPOINTS);
391
- }
395
+ }
应用建议
likedislike
atomgit-botatomgit-bot7月10日

🔴 Critical

CopyInOhters 函数中,第 387–391 行(zeroPoints 段)的 blockLen 赋值存在明显的复制粘贴错误:

这段代码是从第 370–374 行(scales 段)复制过来但未修改目标变量和类型。它导致了两个独立的缺陷:

  1. copyInParamszeroPoints.blockLen 未初始化:第 385 行声明 DataCopyExtParams copyInParamszeroPoints 后仅设置了 blockCount=1blockLen 从未被赋值,保持栈上的脏数据。第 406 行、413 行调用 DataCopyPad(... copyInParamszeroPoints ...) 时使用未初始化的 blockLen,导致未定义行为——可能拷贝错误数量的数据或越界。

  2. copyInParamsScales.blockLen 被意外覆盖:如果 hasScales2 为 true(第 395 行),第 399 行会复用 copyInParamsScales 做 scales2 拷贝,此时 blockLen 已被覆盖为 zeroPoints 段错误写入的值(sizeof(T_SCALES)numR * sizeof(T_SCALES),但语义上应是 zeroPoints 的 size),导致 scales2 拷贝数据量可能不正确。

触发条件:当 hasZeroPoints1hasZeroPoints2 为 true 时触发未初始化 blockLen 的问题;当 hasScales2 && (hasZeroPoints1 || hasZeroPoints2) 时还会触发 scales2 拷贝错误。

建议:将第 388 行的 copyInParamsScales 改为 copyInParamszeroPoints,将 T_SCALES 改为 T_ZEROPOINTS;第 390 行同理。修正后的代码如上。

改动建议
391
+ if (numQ == 1) {
392
+ copyInParamszeroPoints.blockLen = sizeof(T_ZEROPOINTS);
393
+ } else {
394
+ copyInParamszeroPoints.blockLen = numR * sizeof(T_ZEROPOINTS);
391
- }
395
+ }
应用建议
likedislike
384 copyInParamszeroPoints.srcStride = 0;392 copyInParamszeroPoints.srcStride = 0;
385 copyInParamszeroPoints.dstStride = 0;393 copyInParamszeroPoints.dstStride = 0;
386 394 
@@ -9,7 +9,6 @@
9 */9 */
10 10 
11#include "aclnn/aclnn_base.h"11#include "aclnn/aclnn_base.h"
12-#include "op_api/op_api_def.h"
13 12 
14#include "opdev/common_types.h"13#include "opdev/common_types.h"
15#include "opdev/data_type_utils.h"14#include "opdev/data_type_utils.h"