已关闭
组合离散mask越界修复方案 #377
candyhong创建于  3月20日关闭于  3月27日
candyhong成员
3月20日 创建

一、问题根因分析

1.1 典型场景:Triton Kernel 尾块处理逻辑

Triton Kernel 处理数组 / 矩阵尾块(tail block)时,会组合两类 Mask 实现 “防越界 + 筛数据”,如下:

# 连续边界 mask(防止尾块越界)
boundary_mask = tl.arange(0, BLOCK_N) < (N - pid * BLOCK_N)  # 规律连续 mask
idx = tl.load(idx_ptr + ...)  # 离散行选择 mask(运行时加载的 idx 决定)
discrete_mask = idx < HALF    # 离散 mask

combined_mask = boundary_mask[:] & discrete_mask[:] # 组合 mask

val = tl.load(A_ptr + ..., mask=combined_mask)
tl.store(C_ptr + ..., val, mask=combined_mask)

1.2 现有处理逻辑的缺陷

DiscreteMaskAccessConversionPass 通过 isDiscreteMask() 分析 combined_mask 时:

  • discrete_mask 无法静态解析,导致 parseAnd 整体失败;
  • combined_mask 被判定为 “离散 Mask”,触发 DiscreteMaskLoad/StoreConversion 转换;
  • 转换后丢失边界 Mask 保护,导致尾块越界读取/存储
// 原始:带 combined_mask 的合法 Store
tt.store(%ptr, %val, %combined_mask)

// 转换后:无 Mask 全量 Load 引发越界
%old_val  = tt.load(%ptr)    // ★ 访问 ptr 全部地址(含越界部分)
%selected = arith.select(%combined_mask, %val, %old_val)
tt.store(%ptr, %selected)

1.3 越界根因

  • 保护缺失boundary_mask 本用于限制访存范围(如 [0, N - pid*BLOCK_N)),但转换后被完全移除;
  • 越界行为tt.load(%ptr) 访问超出张量边界的 [N - pid*BLOCK_N, BLOCK_N) 地址;
  • 风险表现:NPU 对 GM(Global Memory)越界读可能读取无效数据,极端场景触发 MTE 越界错误。

二、修复方案设计

2.1 核心思路:拆分 Mask 职责

针对 mask = 连续 Mask(contMask) & 离散 Mask(discMask) 场景,明确两类 Mask 的分工:

  • 连续 Mask:仅负责 “访存范围保护”(防止 GM 越界);
  • 离散 Mask:仅负责 “数据筛选”(select 有效数据)。
操作类型 旧方案(有风险) 新方案(安全)
Load 1. 全量 Load
2. 组合 Mask 筛选
1. contMask 安全 Load(防越界)
2. contMask & discMask 筛选(防未初始化内存泄露)
Store 1. 全量 Load
2. 组合 Mask 筛选
3. 全量 Store
1. contMask 安全 Load(防越界)
2. discMask 筛选
3. contMask 安全 Store(防越界)

2.1.1 Load 操作修复路径

// 原始:带组合 Mask 的 Load
%result = tt.load(%ptr, %combined_mask, %other)
              ↓ (分解 combined_mask = contMask & discMask)
// 转换后:安全 Load + 组合 Mask 筛选
%safe_load       = tt.load(%ptr, %contMask)                        // 仅读合法范围(防 OOB)
%full_select_mask = arith.andi(%contMask, %discMask)               // 重组完整筛选条件
%result          = arith.select(%full_select_mask, %safe_load, %other)  // 筛选有效数据

为何 select 条件是 contMask & discMask 而非仅 discMask

LoadConvertertt.load(ptr, contMask) 降级为 SubView+Copy 模式时,只拷贝
contMask 有效区域(如前 M 行),alloc 中 contMask=false 的 padding 区域未被初始化
若 select 条件仅用 discMask,对于 contMask=false & discMask=true 的位置,select 会
从未初始化内存中取值,导致输出为 NaN 或脏数据(偶发性精度错误)。
将条件改为 contMask & discMask 后,contMask=false 的位置强制走 other 分支,
不再触碰未初始化的 padding。

2.1.2 Store 操作修复路径

// 原始:带组合 Mask 的 Store
tt.store(%ptr, %val, %combined_mask)
              ↓ (分解 combined_mask = contMask & discMask)
// 转换后:安全 Load → 筛选 → 安全 Store
%old_val  = tt.load(%ptr, %contMask)                  // ★ 关键修复:仅读合法范围
%selected = arith.select(%discMask, %val, %old_val)   // 筛选有效数据
tt.store(%ptr, %selected, %contMask)                  // 仅写合法范围

Store 路径 select 为何仍用 discMask

tt.store(ptr, selected, contMask) 自带 contMask 保护,contMask=false 的位置完全不写入,
old_val 中 padding 区域的未初始化值不会被写回内存,因此不存在泄露问题。

2.2 具体实现步骤

Step 1:新增辅助函数 / 结构

新增以下结构和函数:

  • MaskDecomposition 结构体(存储分解后的连续 / 离散 Mask);
  • collectAndLeaves 函数(递归拆解 AND 嵌套);
  • decomposeAndMask 函数(分类 + 合并 Mask)。

Step 2:修改 Load 转换逻辑

DiscreteMaskLoadConversion::matchAndRewrite 中:

  1. 原有 isDiscreteMask 检查通过后,调用 decomposeAndMask(mask, loc, rewriter)
  2. contMask != nullptr && discMask != nullptr,走新路径:
    • safeLoad = tt.load(ptr, contMask)(仅读合法范围)
    • fullSelectMask = arith.andi(contMask, discMask)(防未初始化内存泄露)
    • result = arith.select(fullSelectMask, safeLoad, other)
  3. 否则 fallback 到原有逻辑。

Step 3:修改 Store 转换逻辑

DiscreteMaskStoreConversion::matchAndRewrite 中:

  1. 同 Load 逻辑,先调用 decomposeAndMask 分解 Mask;
  2. 若分解出 contMaskdiscMask,走新路径(安全 Load → select → 安全 Store);
    • select 条件仍用 discMask(store 自带 contMask 保护,无未初始化泄露风险)
  3. 否则 fallback 到原有逻辑。

2.3 关键逻辑:多级 AND 的递归分解

combined_mask 可能是多层 arith.andi 嵌套(如 maskA & maskB & maskC),需先分解再合并:

分解策略
  1. 递归拆解:将嵌套 AND 操作拆分为所有叶子节点 Mask;
  2. 分类筛选
    • contLeaves:可被 MaskState::parse() 静态解析的连续 Mask;
    • discLeaves:无法静态解析的离散 Mask;
  3. 重新合并
    • contMask:所有 contLeaves 做 AND 合并(保留越界保护能力);
    • discMask:所有 discLeaves 做 AND 合并(保留数据筛选能力);
  4. 边界兼容:若 contLeaves 为空,退化为原有逻辑(保证兼容性)。

2.3.1 算法伪代码

// 数据结构:存储分解后的连续/离散 Mask
struct MaskDecomposition {
  Value contMask;  // 可静态解析的连续 Mask(防越界)
  Value discMask;  // 动态离散 Mask(筛数据)
};
/**
 * 递归拆解嵌套的 AND 操作,收集所有叶子节点 Mask
 * @param mask: 待分解的组合 Mask
 * @param leaves: 输出参数,存储所有叶子节点
 */
void collectAndLeaves(Value mask, SmallVectorImpl<Value> &leaves) {
  if (auto andOp = dyn_cast<arith::AndIOp>(mask.getDefiningOp())) {
    collectAndLeaves(andOp.getLhs(), leaves);
    collectAndLeaves(andOp.getRhs(), leaves);
  } else {
    leaves.push_back(mask);  // 非 AND 操作,直接作为叶子节点
  }
}

/**
 * 分解组合 Mask 为连续/离散两部分
 * @param mask: 输入的组合 Mask(可能多层 AND 嵌套)
 * @param loc: 位置信息(MLIR 操作创建用)
 * @param rewriter: MLIR 重写器
 * @return 分解后的连续/离散 Mask
 */
MaskDecomposition decomposeAndMask(Value mask, Location loc,
                                   PatternRewriter &rewriter) {
  // 步骤1:递归拆解所有叶子节点
  SmallVector<Value> leaves;
  collectAndLeaves(mask, leaves);
  // 步骤2:分类筛选连续/离散 Mask
  SmallVector<Value> contLeaves, discLeaves;
  for (Value leaf : leaves) {
    MaskState st;
    // 尝试解析当前叶子节点是否为连续 Mask
    if (succeeded(st.parse(leaf, loc, rewriter)) && st.isMask()) {
      // ★ 关键:清理解析过程中生成的临时操作
      st.eraseInsertedOps(rewriter);
      contLeaves.push_back(leaf);
    } else {
      discLeaves.push_back(leaf);
    }
  }
  // 步骤3:合并连续 Mask(AND 操作)
  Value contMask = nullptr;
  for (Value v : contLeaves) {
    contMask = contMask 
               ? rewriter.create<arith::AndIOp>(loc, contMask, v)
               : v;
  }
  // 步骤4:合并离散 Mask(AND 操作)
  Value discMask = nullptr;
  for (Value v : discLeaves) {
    discMask = discMask 
               ? rewriter.create<arith::AndIOp>(loc, discMask, v)
               : v;
  }
  return {contMask, discMask};
}

三、方案局限性

局限性 具体说明 影响范围 & 应对建议
纯离散 Mask 越界未解决 若 Mask 无连续部分(如 mask = idx < HALF),退化为原有全量 Load 逻辑,仍可能越界 影响小:此类 Kernel 极少(尾块通常需连续 Mask 规约);需上层保证 ptr 访问范围合法
仅支持 AND 组合 仅处理 contMask & discMask 结构,不支持 OR / 异或等复杂组合 影响可控:业务场景中 Mask 组合以 AND 为主

四、 方案验证

4.1 单元测试场景(覆盖核心场景)

Mask 组合类型 Load 测试 Store 测试 Load+Store 测试 测试目的
纯离散 Mask (A) (B) - 验证退化逻辑兼容性
纯连续 Mask (C) (D) - 验证连续 Mask 保护逻辑
连续 + 离散 2 路 AND 组合 (E) (F) (G) 核心场景:防尾块越界
连续 + 离散 4 路 AND 嵌套 - - (H) 验证递归分解逻辑
likedislike
Ccandyhong成员
3月20日 关联了pull request:fix: decompose AND masks to preserve boundary safety in discrete mask conversion
Ccandyhong成员
3月20日 修改标题为 “组合离散mask越界修复方案”,原标题为“离散mask越界修复方案”
Ccandyhong成员
3月20日 修改了issue 的描述
Ccandyhong成员
3月23日 关联了pull request:fix: decompose AND masks to preserve boundary safety in discrete mask conversion
Ccandyhong成员
3月26日 将 candyhong 设为负责人
Ccandyhong成员
3月26日 关联了pull request:feat: promote to be committer of triton-ascend repo
Ccandyhong成员
3月27日 issue状态由 TODO 改变为 DONE
Ccandyhong成员
3月27日 关闭了 issue
candyhong成员
3月27日 评论:

修复已合入,Issue正常关闭

likedislike
ascend-robotascend-robot成员
3月27日 添加了label:resolved