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)
DiscreteMaskAccessConversionPass 通过 isDiscreteMask() 分析 combined_mask 时:
DiscreteMaskAccessConversionPass
isDiscreteMask()
combined_mask
discrete_mask
parseAnd
DiscreteMaskLoad/StoreConversion
// 原始:带 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)
boundary_mask
[0, N - pid*BLOCK_N)
tt.load(%ptr)
[N - pid*BLOCK_N, BLOCK_N)
针对 mask = 连续 Mask(contMask) & 离散 Mask(discMask) 场景,明确两类 Mask 的分工:
mask = 连续 Mask(contMask) & 离散 Mask(discMask)
// 原始:带组合 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? LoadConverter 将 tt.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。
为何 select 条件是 contMask & discMask 而非仅 discMask?
contMask & discMask
discMask
LoadConverter 将 tt.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。
LoadConverter
tt.load(ptr, contMask)
SubView+Copy
alloc
contMask=false & discMask=true
other
// 原始:带组合 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 区域的未初始化值不会被写回内存,因此不存在泄露问题。
Store 路径 select 为何仍用 discMask?
tt.store(ptr, selected, contMask) 自带 contMask 保护,contMask=false 的位置完全不写入, old_val 中 padding 区域的未初始化值不会被写回内存,因此不存在泄露问题。
tt.store(ptr, selected, contMask)
old_val
新增以下结构和函数:
MaskDecomposition
collectAndLeaves
decomposeAndMask
在 DiscreteMaskLoadConversion::matchAndRewrite 中:
DiscreteMaskLoadConversion::matchAndRewrite
isDiscreteMask
decomposeAndMask(mask, loc, rewriter)
contMask != nullptr && discMask != nullptr
safeLoad = tt.load(ptr, contMask)
fullSelectMask = arith.andi(contMask, discMask)
result = arith.select(fullSelectMask, safeLoad, other)
在 DiscreteMaskStoreConversion::matchAndRewrite 中:
DiscreteMaskStoreConversion::matchAndRewrite
contMask
combined_mask 可能是多层 arith.andi 嵌套(如 maskA & maskB & maskC),需先分解再合并:
arith.andi
maskA & maskB & maskC
contLeaves
MaskState::parse()
discLeaves
// 数据结构:存储分解后的连续/离散 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 = idx < HALF
修复已合入,Issue正常关闭
一、问题根因分析
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转换;1.3 越界根因
boundary_mask本用于限制访存范围(如[0, N - pid*BLOCK_N)),但转换后被完全移除;tt.load(%ptr)访问超出张量边界的[N - pid*BLOCK_N, BLOCK_N)地址;二、修复方案设计
2.1 核心思路:拆分 Mask 职责
针对
mask = 连续 Mask(contMask) & 离散 Mask(discMask)场景,明确两类 Mask 的分工:2. 组合 Mask 筛选
2. contMask & discMask 筛选(防未初始化内存泄露)
2. 组合 Mask 筛选
3. 全量 Store
2. discMask 筛选
3. contMask 安全 Store(防越界)
2.1.1 Load 操作修复路径
2.1.2 Store 操作修复路径
2.2 具体实现步骤
Step 1:新增辅助函数 / 结构
新增以下结构和函数:
MaskDecomposition结构体(存储分解后的连续 / 离散 Mask);collectAndLeaves函数(递归拆解 AND 嵌套);decomposeAndMask函数(分类 + 合并 Mask)。Step 2:修改 Load 转换逻辑
在
DiscreteMaskLoadConversion::matchAndRewrite中:isDiscreteMask检查通过后,调用decomposeAndMask(mask, loc, rewriter);contMask != nullptr && discMask != nullptr,走新路径:safeLoad = tt.load(ptr, contMask)(仅读合法范围)fullSelectMask = arith.andi(contMask, discMask)(防未初始化内存泄露)result = arith.select(fullSelectMask, safeLoad, other)Step 3:修改 Store 转换逻辑
在
DiscreteMaskStoreConversion::matchAndRewrite中:decomposeAndMask分解 Mask;contMask和discMask,走新路径(安全 Load → select → 安全 Store);discMask(store 自带 contMask 保护,无未初始化泄露风险)2.3 关键逻辑:多级 AND 的递归分解
combined_mask可能是多层arith.andi嵌套(如maskA & maskB & maskC),需先分解再合并:分解策略
contLeaves:可被MaskState::parse()静态解析的连续 Mask;discLeaves:无法静态解析的离散 Mask;contMask:所有contLeaves做 AND 合并(保留越界保护能力);discMask:所有discLeaves做 AND 合并(保留数据筛选能力);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 = idx < HALF),退化为原有全量 Load 逻辑,仍可能越界contMask & discMask结构,不支持 OR / 异或等复杂组合四、 方案验证
4.1 单元测试场景(覆盖核心场景)