Prhub

#32410 Fix flaky test_sampling_mask: mask length can legitimately be top_k + 1

原始 PR 作者 merrymercy 合并时间 2026-07-26 07:05 文件变更 1 提交数 1 评论 4 代码增减 +7 / -3

执行摘要

修复不稳定的 testing_mask 测试断言

PR body 指出测试 test_generate_returns_sampling_mask 不稳定,最近 CI 运行出现以下错误:AssertionError: 11 not less than or equal to 10。根因是 Sampler._attach_sampling_mask_to_output 有意在 mask 后追加实际采样的 token(当采样 kernel 选中刚好超出 mask 的 torch.topk 重建范围的 token 时),导致合法长度为 top_k + 1。该测试自引入以来一直存在潜在不稳定性。

值得合并:修复了不稳定测试,无需 runtime 变更。

讨论亮点

该 PR 无 review 评论。作者 merrymercy 自行审查并批准。

实现拆解

该 PR 仅修改测试文件 test/registered/sampling/test_sampling_mask.py,对核心逻辑无影响。变更内容为:

  • 对于 top_p 采样情况:将 assertLessEqual(len(sampling_mask), _TOP_K) 改为 assertLessEqual(len(sampling_mask), _TOP_K + 1)
  • 对于仅 top_k 以及 top_k + top_p=1.0 的情况:将 assertEqual(len(sampling_mask), _TOP_K) 改为 assertIn(len(sampling_mask), (_TOP_K, _TOP_K + 1)),以容忍 +1 情况的同时保留下限保证。
  • 添加注释说明 +1 的原因(浮点累积和偏差)。
文件 模块 状态 重要度
test/registered/sampling/test_sampling_mask.py 采样测试 modified 4.1

关键符号

test_generate_returns_sampling_mask

关键源码片段

test/registered/sampling/test_sampling_mask.py test-coverage

唯一修改的文件:将测试断言从严格的 `assertEqual`/`assertLessEqual` 放宽为允许 mask 长度为 `top_k + 1`,以匹配 `Sampler._attach_sampling_mask_to_output` 的实际行为。

def test_generate_returns_sampling_mask(self):
    # ... setup ...
    # The mask keeps at most top_k tokens, plus possibly the actually
    # sampled token when the sampling kernel picks one just outside the
    # mask's topk reconstruction (fp cumsum divergence); see
    # Sampler._attach_sampling_mask_to_output.
    for sampling_mask in top_p_sampling_masks:
        self.assertLessEqual(len(sampling_mask), _TOP_K + 1)
​
    # ... other cases ...
    for sampling_mask in top_k_sampling_masks:
        self.assertIn(len(sampling_mask), (_TOP_K, _TOP_K + 1))
​
    for sampling_mask in top_k_top_p_one_sampling_masks:
        self.assertIn(len(sampling_mask), (_TOP_K, _TOP_K + 1))

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

风险极低:仅修改测试断言,不改变任何运行时行为。测试仍然验证 sampled token 在 mask 内(assertIn(output_id, sampling_mask))。

影响范围极小:仅影响一个测试文件中三个断言行。消除不确定性,使 CI 更可靠。

测试覆盖调整

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论