Prhub

#50323 [CI] Add option to raise an exception when NaNs are detected in logits

原始 PR 作者 tlrmchlsmth 合并时间 2026-08-05 07:02 文件变更 6 提交数 12 评论 8 代码增减 +60 / -3

执行摘要

新增 logits NaN fail-fast 开关,CI 全局检测

PR body 明确指出:logits 中的 NaN 往往与 KV cache 的 NaN 相伴出现,而注意力内核通过乘零进行 mask,一旦 KV cache 被 NaN 污染就会产生灾难性故障,属于推理服务的严重故障模式。关联的 Dao-AILab/flash-attention#1974 给出了具体复现:vLLM V1 中 hybrid 模型(如 attention + mamba2)的 mamba state 与 KV cache 共享底层内存但使用不同 dtype 视图,导致 KV 张量出现 bf16 下的 NaN 位模式,即使 FA3 理论上不应读取这些区域,输出仍会出错。本 PR 的目的是让这类问题在 CI 中立即暴露(fail-fast),而不是让 eval 静默产出错误结果。

值得精读,但精读重点不是代码量(仅 +60/-3),而是其设计演进过程:从 per-eval 指标抓取到全局 fail-fast 环境变量的转变,体现了 CI 诊断能力的基础设施化思路。可以关注 raise_if_nan_logits 的实现细节(成功路径零分配、异常消息携带请求级信息)以及 env 变量间隐式联动的写法;对需要构建类似检测开关的团队有参考价值。

讨论亮点

核心讨论围绕检测方式展开:

  • njhill 在 issue 评论中建议:与其靠抓取 metrics,不如让 VLLM_COMPUTE_NANS_IN_LOGITS 模式在遇到 NaN 时抛异常,然后全局在 CI 设置,doesn't even need to be added per test. Basically it would just blow up the test if encountered。
  • mgoin 附议并强调希望全局开启、让 pytest/script 直接失败。
  • 作者据此将方案从逐个 eval 抓取 vllm:corrupted_requests_total 指标改为新增 VLLM_RAISE_ON_LOGIT_NANS 运行时开关,并提到配套的 ci-infra#444 会在 buildkite 全局启用。
  • AndreasKaratzas 在 deepseek_v2_lite_ep_eplb.sh 的 diff 上询问是否由各硬件后端决定开关,或直接在 AMD 的 run-am-tests.sh 默认开启;最终该文件在方案调整后未保留逐 eval 设置,全局由 ci-infra 控制。
  • depthfirst-app bot 发现中间提交把 os.getenv 默认值写成 1,与类注解 False 矛盾,会让所有生产实例在 NaN 时直接崩溃;作者在收尾提交中恢复默认 0。

实现拆解

  1. 新增环境变量开关:在 vllm/envs.py 中定义 VLLM_RAISE_ON_LOGIT_NANS: bool = False,解析默认值为 0;同时让 VLLM_COMPUTE_NANS_IN_LOGITS 的解析逻辑 OR 上新开关,保证开启 fail-fast 时自动启用 NaN 计算,避免出现想抛异常却没算 NaN 的空转。

  2. 提炼公共判定函数:在 vllm/v1/worker/utils.py 新增 raise_if_nan_logits(num_nans_in_logits),传入请求级 NaN 计数映射;全零则直接返回,否则构造 corrupted_requests 字典并抛出 RuntimeError,消息中直接列出被污染的请求 ID 与 NaN 数量。

  3. 同步执行路径接入:在 vllm/v1/worker/gpu_model_runner.py 的 _get_nans_in_logits 返回统计前,按 envs.VLLM_RAISE_ON_LOGIT_NANS 决定是否调用 raise_if_nan_logits,覆盖 v1/v2 两种模型执行器。

  4. 异步输出路径接入:在 vllm/v1/worker/gpu/async_utils.py 的 AsyncOutput.get_output 中,在组装 num_nans_in_logits 映射后同样按环境变量调用 raise_if_nan_logits,保证 CUDA graph 重叠执行场景下也能及时暴露问题。

  5. 测试与 CI 配套:tests/basic_correctness/test_basic_correctness.py 新增 test_raise_on_logit_nans,通过 monkeypatch 在 compute_logits 中注入单个 NaN,分别以 v1/v2 model runner 跑真实推理并断言 RuntimeError;tests/v1/worker/test_gpu_model_runner.py 的既有单测显式关闭新开关以保持统计逻辑可独立测试。全局开启由配套仓库 vllm-project/ci-infra#444 在 buildkite 上统一设置,本仓库不硬编码。

文件 模块 状态 重要度
vllm/v1/worker/utils.py 诊断工具 modified 6.7
vllm/v1/worker/gpu_model_runner.py 模型执行 modified 5.84
vllm/envs.py 环境配置 modified 5.4
vllm/v1/worker/gpu/async_utils.py 异步输出 modified 5.23
tests/basic_correctness/test_basic_correctness.py 正确性测试 modified 5.67
tests/v1/worker/test_gpu_model_runner.py runner 测试 modified 4.35

关键符号

raise_if_nan_logits _get_nans_in_logits AsyncOutput.get_output test_raise_on_logit_nans compute_nan_logits test_get_nans_in_logits

关键源码片段

vllm/v1/worker/utils.py core-logic

新增 raise_if_nan_logits 工具函数,是整个 fail-fast 机制的判定与异常构造核心,被同步与异步两条路径复用。

def raise_if_nan_logits(num_nans_in_logits: Mapping[str, int]) -> None:
    # 任一请求出现 NaN 即触发异常;全为 0 时直接返回,避免无谓分配。
    if not any(num_nans_in_logits.values()):
        return
​
    # 仅收集真正被污染的请求,便于快速定位故障来源。
    corrupted_requests = {
        req_id: num_nans
        for req_id, num_nans in num_nans_in_logits.items()
        if num_nans > 0
    }
    # RuntimeError 会让上层测试 / 服务直接失败,而不是静默产出乱码。
    raise RuntimeError(f'NaNs detected in logits: {corrupted_requests}')
vllm/v1/worker/gpu_model_runner.py dependency-wiring

同步执行路径的接入点,_get_nans_in_logits 在返回统计前根据 env 决定是否抛异常,影响 v1/v2 model runner。

def _get_nans_in_logits(
    self,
    logits: torch.Tensor | None,
) -> dict[str, int]:
    try:
        if logits is None:
            return {req_id: 0 for req_id in self.input_batch.req_ids}
​
        num_nans_in_logits = {}
        # 按词表维度统计每行 NaN 数量,并同步到 CPU 便于逐请求换算。
        num_nans_for_index = logits.isnan().sum(dim=-1).cpu().numpy()
        for req_id in self.input_batch.req_ids:
            req_index = self.input_batch.req_id_to_index[req_id]
            num_nans_in_logits[req_id] = (
                int(num_nans_for_index[req_index])
                if num_nans_for_index is not None and req_index < logits.shape[0]
                else 0
            )
        # 仅在显式开启 fail-fast 时抛出异常;默认保持向后兼容。
        if envs.VLLM_RAISE_ON_LOGIT_NANS:
            raise_if_nan_logits(num_nans_in_logits)
        return num_nans_in_logits
    except IndexError:
        # 竞态下请求可能已被移除,静默返回空字典。
        return {}
vllm/envs.py configuration

新增 VLLM_RAISE_ON_LOGIT_NANS 环境变量,并让 VLLM_COMPUTE_NANS_IN_LOGITS 隐式随其开启;默认值最终恢复为 0,避免影响生产。

# 开启 logits NaN 检测(可能增加计算开销),主要用于排查底层 bug 或坏硬件。
'VLLM_COMPUTE_NANS_IN_LOGITS': lambda: bool(
    int(os.getenv('VLLM_COMPUTE_NANS_IN_LOGITS', '0'))
    or int(os.getenv('VLLM_RAISE_ON_LOGIT_NANS', '0'))
),
# logits 中出现 NaN 时直接抛异常;开启后隐式启用上方 NaN 计算。
'VLLM_RAISE_ON_LOGIT_NANS': lambda: bool(
    int(os.getenv('VLLM_RAISE_ON_LOGIT_NANS', '0'))
),

评论区精华

全局 fail-fast 替代逐 eval 指标抓取 设计

njhill 建议与其靠抓取 metrics,不如让 NaN 检测模式直接抛异常并全局在 CI 设置;mgoin 附议并希望 pytest/script 直接失败。

结论:作者采纳建议,新增 VLLM_RAISE_ON_LOGIT_NANS 开关,移除逐 eval 的 metrics 抓取与辅助函数,并在 ci-infra#444 中全局启用。 · 已解决

是否由各硬件后端统一默认开启 question

AndreasKaratzas 在 deepseek_v2_lite_ep_eplb.sh 的 diff 上询问是否让各 HW backend 决定开关值,或直接在 AMD run-am-tests.sh 默认开启。

结论:最终该逐 eval 设置文件未保留,vLLM 仓库内保持 opt-in;是否在 AMD 默认开启未明确,全局由 ci-infra 控制。 · unresolved

环境变量默认值一度为 1 的回归风险 正确性

depthfirst-app bot 指出中间提交把 os.getenv 默认值改为 1,与类注解 False 矛盾,会让所有生产实例在 NaN 时直接崩溃。

结论:作者在收尾提交 Fix NaN detection test cleanup 中恢复默认 0,保持 opt-in。 · 已解决

风险与影响

  1. 性能开销:开启后每个前向步都要对 logits 做 isnan().sum(dim=-1) 并执行 .cpu() 同步以逐请求统计,会打断 GPU 流水线;这是设计上的已知取舍,因此默认关闭,仅 CI/调试时开启。
  2. 生产误开风险:中间提交曾把默认值设为 1,若按那个版本发布,任何一次 NaN logits 都会让整个服务抛 RuntimeError 崩溃;最终已恢复 opt-in 默认 0,但后续修改 env 解析时需警惕再次引入同类回归。
  3. 异常语义变化:_get_nans_in_logits 捕获 IndexError 后静默返回空字典,若请求在统计期间被移除,NaN 会被吞掉;fail-fast 模式下该静默路径可能掩盖真实问题。
  4. 测试耦合:端到端测试通过 monkeypatch 注入 NaN 并匹配异常消息 NaNs detected in logits,若未来修改消息文案或 compute_logits 签名,测试需要同步维护。
  1. 用户:默认零影响(新变量默认关闭);显式开启后,logits 出现 NaN 会从静默生成乱码变为明确报错,大幅降低 FA3 类内核问题的定位成本。
  2. 系统:为 CI 提供全局 fail-fast 能力,配合 ci-infra#444 可在 buildkite 全量测试中第一时间捕获 KV cache/注意力核污染,避免 eval 用垃圾数字通过。
  3. 团队:新增一个环境变量和公共工具函数,后续其他执行路径(如 pooling、CPU runner)若要接入只需复用 raise_if_nan_logits;需要同步更新环境变量文档。
默认关闭,生产无影响 开启时有 GPU-CPU 同步开销 异常消息被测试断言耦合 IndexError 静默吞掉 NaN 中间版本曾默认开启

关联 Issue

#1974 FlashAttention3 forward producing NaN output when NaN exist in parts of input data that it should not be reading

完整报告

参与讨论