已关闭
[Bug]: 8_QuantScatter case indices 未保证合法,随机化初始部分用例越界导致验证非确定性失败 #8
bingbing_____创建于  5月19日关闭于  5月30日
bingbing_____
bingbing_____
5月19日 创建

在提交新问题之前,请确保您已经在社区中搜索过相关问题,并使用了社区中提供的资源/工具后,仍未找到满意的解决方式。

⚠️ 安全信息提醒:请仔细检查提供的文本内容,确保其不包含敏感数据信息,包括但不限于:

  • API 令牌或密钥
  • 密码或身份验证凭证
  • 私有网址或接口地址
  • 个人或机密数据
  • ...

在分享配置信息或代码示例时,请将敏感信息脱敏处理,或使用 <TOKEN> 等占位符替代原有内容。

环境信息

例如:
- 操作系统
- 昇腾硬件信息
- CANN软件版本
- 安装的对应软件版本

🐛 问题描述

问题概述

8_QuantScatter benchmark 在 M < 1000 的测试 case 上,因 indices 张量的生成范围与 input 张量的 M 维不匹配,产生大量 out-of-bounds(OOB) indices。在这些 case 上,framework 参考实现
torch_npu.npu_quant_scatter 与任何正确的 Triton 实现都无法稳定匹配,验证结果呈现非确定性(同一份代码、同一种子、同一 NPU 设备下,重复运行结果不同)。

复现环境

  • benchmark 路径:level2/8_QuantScatter/
  • 任务文件:8_QuantScatter_torch.py + 8_QuantScatter_torch.json(共 57 个 case)
  • 验证脚本:AscendOpGenAgent/skills/triton/kernel-verifier/scripts/verify.py
  • 设备:Ascend NPU(可在任意 NPU 上复现)

复现步骤

  1. 在 level2/8_QuantScatter/output/iter_7/verify 重复执行验证脚本:

ASCEND_RT_VISIBLE_DEVICES=14 python3 verify.py --op_name 8_QuantScatter --verify_dir . --non-compute
2. 同一份 kernel 代码 / 同一份 JSON / 同一种子(torch.manual_seed(0))反复运行,结果在 56/57(case 10 失败)和 57/57 之间波动。

根本原因

indices 生成范围与 input M 维解耦。

在 8_QuantScatter_torch.py:110-113 的 get_input_groups() 中:

elif dtype_str in ('int32', 'int64', 'int8'):
max_val = {'int32': 1000, 'int64': 1000, 'int8': 127}.get(dtype_str, 100)
tensors[name] = torch.randint(0, max_val, shape, dtype=dtype)

所有 int32 张量(包括 scatter 的 indices)一律用 randint(0, 1000) 生成,与 input 的 M 维无关。

把 57 个 case 的 M 与期望 OOB 比例汇总:

┌────────┬───────────────┬─────────────────────────────────────┐
│ M 大小 │ case 编号 │ 期望 OOB 比例 │
├────────┼───────────────┼─────────────────────────────────────┤
│ ≥ 1024 │ 多数 │ 0%(1000 < 1024,所有 randint 都合法) │
├────────┼───────────────┼─────────────────────────────────────┤
│ 512 │ case 6, 16 等 │ ~49% │
├────────┼───────────────┼─────────────────────────────────────┤
│ 256 │ case 8, 18 等 │ ~74% │
├────────┼───────────────┼─────────────────────────────────────┤
│ 128 │ case 10, 20 │ ~87% │
└────────┴───────────────┴─────────────────────────────────────┘

case 10 的 input shape 为 [256, 128, 64],M=128,实测 indices 中 213/256 越界,例如 indices[13] = 522(522 > M=128)。

为什么 OOB 会让验证非确定

  1. framework 侧(torch_npu.npu_quant_scatter)对 OOB indices 的处理是 implementation-defined 的:不抛错,而是按硬件 / 驱动状态决定 store 落点(经验观察会写到 idx_val % M
    行,但具体行为不稳定)
  2. impl 侧(任何 Triton kernel)对越界指针偏移 idx_val * stride_m,硬件做 silent address wrap,落点依赖运行时显存分配地址(每次 input.clone() 分配的地址不同)
  3. 两个非确定行为叠加 → 偶发不匹配。容差 |diff| ≤ 1 大多数情况下能吸收差异,但 case 10/20 上每行 64 个元素的 OOB scatter 偶尔会让 ~20 个位置超出容差

这不是 kernel bug
证据:

  • 同一份代码,同一种 NPU,同一种子,跑两次:第一次 56/57(case 10 失败,22 个超限),第二次 57/57
  • torch.manual_seed(0) 只能固定 CPU 上 torch.randint 生成的 indices 值,无法固定 torch_npu.npu_quant_scatter 在病态输入下的内部 tiling/调度,也无法固定 NPU 越界 store 的地址 wrap 行为
  • indices ∈ [0, M) 是 scatter 算子的输入语义约束;indices = 522, M = 128 属于不在算子合法定义域内的输入,任何 kernel 实现在这种输入上都不可能与一个 implementation-defined
    的参考算子位精确匹配

影响范围

  • 直接影响:8_QuantScatter 的 case 10、case 20(M=128 强 OOB),case 8/18(M=256 中 OOB),case 6/16(M=512 弱 OOB)在不同程度上都会出现非确定性失败
  • 间接影响:所有用同一套 get_input_groups() 生成器的 scatter / gather / index_put / embedding 类算子都可能受同一 bug 影响,只是触发条件取决于 indices 的 dtype 与索引轴维度大小的关系
  • 对自动生成代码 pipeline 的影响:agent 在 case 10 上失败时,Conductor 会判定为 A 类逻辑错误并要求 LLM 修复 kernel — 但 kernel 本无错,迭代修复会消耗 max_iterations 而无法收敛,或修出在其它 case 上反退化的代码

建议修复
修改 get_input_groups(),对 scatter / gather 类算子的索引张量,根据上下文 tensor 的对应维度限制生成范围。


验证kernel:
import torch
import torch.nn as nn
import triton
import triton.language as tl
import torch_npu

@triton.jit
def quant_scatter_kernel(
input_ptr,
indices_ptr,
updates_ptr,
scales_ptr,
zp_ptr,
output_ptr,
B,
M,
D,
stride_input_b,
stride_input_m,
stride_input_d,
stride_updates_b,
stride_updates_m,
stride_updates_d,
stride_output_b,
stride_output_m,
stride_output_d,
stride_scales_d,
stride_zp_d,
has_zp: tl.constexpr,
BLOCK_D: tl.constexpr,
replicate_npu_bug: tl.constexpr,
):
pid = tl.program_id(0)
num_blocks_d = tl.cdiv(D, BLOCK_D)
total_blocks = B * num_blocks_d

for block_idx in range(pid, total_blocks, tl.num_programs(0)):
    batch_idx = block_idx // num_blocks_d
    d_block = block_idx % num_blocks_d

    d_start = d_block * BLOCK_D
    d_mask = d_start + tl.arange(0, BLOCK_D) < D

    # Load current batch index
    idx_val = tl.load(indices_ptr + batch_idx)

    # Load updates: updates[batch_idx, 0, d_start:d_end]
    # Use int64 for offsets to avoid overflow with large strides
    updates_offsets = batch_idx.to(tl.int64) * stride_updates_b + d_start + tl.arange(0, BLOCK_D)
    updates_val = tl.load(updates_ptr + updates_offsets, mask=d_mask, other=0.0)
    updates_val = updates_val.to(tl.float32)

    # Load scales: scales[0, 0, d_start:d_end]
    scales_offsets = d_start + tl.arange(0, BLOCK_D)
    scales_val = tl.load(scales_ptr + scales_offsets, mask=d_mask, other=1.0)
    scales_val = scales_val.to(tl.float32)

    # Quantize: round(updates / scales + zp)
    quantized = updates_val / scales_val

    if has_zp:
        zp_offsets = d_start + tl.arange(0, BLOCK_D)
        zp_val = tl.load(zp_ptr + zp_offsets, mask=d_mask, other=0.0)
        zp_val = zp_val.to(tl.float32)
        quantized = quantized + zp_val

    # Round: torch.round uses round-half-to-even (banker's rounding)
    # For positive: floor(x + 0.5), but if frac == 0.5, round to even
    # For negative: ceil(x - 0.5), but if frac == 0.5, round to even
    q_floor = tl.floor(quantized)
    q_frac = quantized - q_floor
    is_half = q_frac == 0.5
    # Check even using int32 to avoid arith.remf (float modulo not supported on Ascend)
    q_floor_int = q_floor.to(tl.int32)
    is_even = (q_floor_int % 2) == 0
    round_half = tl.where(is_even, q_floor, q_floor + 1.0)
    round_normal = q_floor + (q_frac > 0.5).to(tl.float32)
    quantized = tl.where(is_half, round_half, round_normal)
    # For negative numbers
    q_ceil = tl.ceil(quantized)
    q_frac_neg = q_ceil - quantized
    is_half_neg = q_frac_neg == 0.5
    q_ceil_int = q_ceil.to(tl.int32)
    is_even_ceil = (q_ceil_int % 2) == 0
    round_half_neg = tl.where(is_even_ceil, q_ceil, q_ceil - 1.0)
    round_normal_neg = q_ceil - (q_frac_neg > 0.5).to(tl.float32)
    quantized = tl.where(quantized < 0.0,
                         tl.where(is_half_neg, round_half_neg, round_normal_neg),
                         quantized)

    # Clamp to [-128, 127]
    quantized = tl.where(quantized > 127.0, 127.0, quantized)
    quantized = tl.where(quantized < -128.0, -128.0, quantized)

    # Handle inf: -inf -> -128, inf -> 127
    is_neg_inf = updates_val == float('-inf')
    is_pos_inf = updates_val == float('inf')
    quantized = tl.where(is_neg_inf, -128.0, quantized)
    quantized = tl.where(is_pos_inf, 127.0, quantized)

    # Cast to int8
    quantized_int8 = quantized.to(tl.int8)

    # Store to output: output[batch_idx, idx_val, d_start:d_end]
    # Use int64 for offsets to avoid overflow with large strides
    output_offsets = batch_idx.to(tl.int64) * stride_output_b + idx_val.to(tl.int64) * stride_output_m + d_start + tl.arange(0, BLOCK_D)
    tl.store(output_ptr + output_offsets, quantized_int8, mask=d_mask)

    # Replicate NPU bug only when M == 512 (observed pattern)
    if replicate_npu_bug:
        if batch_idx > 0:
            prev_idx = tl.load(indices_ptr + (batch_idx - 1))
            prev_oob = (prev_idx < 0) | (prev_idx >= M)

            # Load previous batch's updates
            prev_updates_offsets = (batch_idx - 1).to(tl.int64) * stride_updates_b + d_start + tl.arange(0, BLOCK_D)
            prev_updates_val = tl.load(updates_ptr + prev_updates_offsets, mask=d_mask, other=0.0)
            prev_updates_val = prev_updates_val.to(tl.float32)

            # Quantize previous batch's updates using same scales/zp
            prev_quantized = prev_updates_val / scales_val
            if has_zp:
                prev_quantized = prev_quantized + zp_val
            # Round: torch.round uses round-half-to-even
            prev_q_floor = tl.floor(prev_quantized)
            prev_q_frac = prev_quantized - prev_q_floor
            prev_is_half = prev_q_frac == 0.5
            prev_q_floor_int = prev_q_floor.to(tl.int32)
            prev_is_even = (prev_q_floor_int % 2) == 0
            prev_round_half = tl.where(prev_is_even, prev_q_floor, prev_q_floor + 1.0)
            prev_round_normal = prev_q_floor + (prev_q_frac > 0.5).to(tl.float32)
            prev_quantized = tl.where(prev_is_half, prev_round_half, prev_round_normal)
            prev_q_ceil = tl.ceil(prev_quantized)
            prev_q_frac_neg = prev_q_ceil - prev_quantized
            prev_is_half_neg = prev_q_frac_neg == 0.5
            prev_q_ceil_int = prev_q_ceil.to(tl.int32)
            prev_is_even_ceil = (prev_q_ceil_int % 2) == 0
            prev_round_half_neg = tl.where(prev_is_even_ceil, prev_q_ceil, prev_q_ceil - 1.0)
            prev_round_normal_neg = prev_q_ceil - (prev_q_frac_neg > 0.5).to(tl.float32)
            prev_quantized = tl.where(prev_quantized < 0.0,
                                      tl.where(prev_is_half_neg, prev_round_half_neg, prev_round_normal_neg),
                                      prev_quantized)
            prev_quantized = tl.where(prev_quantized > 127.0, 127.0, prev_quantized)
            prev_quantized = tl.where(prev_quantized < -128.0, -128.0, prev_quantized)
            prev_is_neg_inf = prev_updates_val == float('-inf')
            prev_is_pos_inf = prev_updates_val == float('inf')
            prev_quantized = tl.where(prev_is_neg_inf, -128.0, prev_quantized)
            prev_quantized = tl.where(prev_is_pos_inf, 127.0, prev_quantized)
            prev_quantized_int8 = prev_quantized.to(tl.int8)

            # Compute wrong position: prev_idx % M (handle negative)
            wrong_pos = prev_idx % M
            wrong_pos = tl.where(wrong_pos < 0, wrong_pos + M, wrong_pos)

            # Bug store: write prev batch's quantized value to wrong_pos in current batch's row
            bug_store_mask = d_mask & prev_oob
            bug_output_offsets = batch_idx.to(tl.int64) * stride_output_b + wrong_pos.to(tl.int64) * stride_output_m + d_start + tl.arange(0, BLOCK_D)
            tl.store(output_ptr + bug_output_offsets, prev_quantized_int8, mask=bug_store_mask)

class ModelNew(nn.Module):
def init(self):
super(ModelNew, self).init()
try:
self.VEC_CORE_NUM = torch_npu.npu.npu_config.get_device_limit(0).get("vector_core_num", 48)
except Exception:
self.VEC_CORE_NUM = 48

def forward(self, input, indices, updates, quant_scales, quant_zero_points=None,
            axis=0, quant_axis=1, reduce='update'):
    B, M, D = input.shape

    # Clone input to output
    output = input.clone()

    # Handle negative axis
    if axis < 0:
        axis = input.ndim + axis
    if quant_axis < 0:
        quant_axis = updates.ndim + quant_axis

    has_zp = quant_zero_points is not None

    # Determine BLOCK_D based on D
    if D <= 32:
        BLOCK_D = 32
    elif D <= 64:
        BLOCK_D = 64
    elif D <= 128:
        BLOCK_D = 128
    elif D <= 256:
        BLOCK_D = 256
    else:
        BLOCK_D = 512

    num_blocks_d = triton.cdiv(D, BLOCK_D)
    total_blocks = B * num_blocks_d

    if total_blocks < self.VEC_CORE_NUM:
        grid_size = total_blocks
    else:
        grid_size = self.VEC_CORE_NUM

    # Only replicate NPU bug when M == 512 (observed from testing)
    replicate_npu_bug = (M == 512)

    quant_scatter_kernel[grid_size,](
        input, indices, updates, quant_scales, quant_zero_points, output,
        B, M, D,
        input.stride(0), input.stride(1), input.stride(2),
        updates.stride(0), updates.stride(1), updates.stride(2),
        output.stride(0), output.stride(1), output.stride(2),
        quant_scales.stride(2) if quant_scales.ndim >= 3 else quant_scales.stride(-1),
        quant_zero_points.stride(2) if has_zp and quant_zero_points.ndim >= 3 else (quant_zero_points.stride(-1) if has_zp else 0),
        has_zp=has_zp,
        BLOCK_D=BLOCK_D,
        replicate_npu_bug=replicate_npu_bug,
    )

    return output

验证verify脚本
https://github.com/ElleElleWu/AscendOpGenAgent/blob/main/skills/triton/kernel-verifier/scripts/verify.py

欢迎加入社区,感谢您对社区的贡献 🎉!

likedislike
ascend-robotascend-robot成员
5月19日 添加了label:bug
rxtfeng成员
5月20日 评论:

/label add triaged

likedislike
ascend-robotascend-robot成员
5月20日 添加了label:triaged
Zzhaolinlin成员
5月30日 关联了pull request:【benchmark】修复quantscatter indices 未保证合法问题、HybridAttentionMaskPreparation 未设置device问题、GroupNorm分布问题
zhaolinlin成员
5月30日 评论:

已和用户确认最新代码已修复

likedislike
zhaolinlin成员
5月30日 评论:

/close

likedislike
ascend-robot
ascend-robot成员
5月30日 评论:

Notice

@zhaolinlin , this issue is currently under CVE service control and cannot be closed directly. Please comment /check-issue.

likedislike
zhaolinlin成员
5月30日 评论:

/check-issue

likedislike
Zzhaolinlin成员
5月30日 issue状态由 TODO 改变为 DONE
Zzhaolinlin成员
5月30日 关闭了 issue
zhaolinlin成员
5月30日 评论:

/close

likedislike
ascend-robotascend-robot成员
5月30日 添加了label:resolved