[Kernel] Add pluggable NaN/Inf tensor checks to NaN detector

Add check_tensor() to NaNDetector for checking arbitrary tensors at
any point in the forward pass. Uses torch.isfinite() — all ops stay
on GPU, CUDA-graph compatible, writes to the same per-token flag
array as the RMSNorm kernel checks.

Any module can register checkpoints via register() and call
check_tensor(tensor, idx) in its forward. update_layer_names()
picks up _nan_detect_indices dicts for readable names.

Wire into DeepseekV2Attention to check attn_output before o_proj,
distinguishing "attention produced NaN" from "o_proj produced NaN".

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
Signed-off-by: Tyler Michael Smith <tlrmchlsmth@gmail.com>
This commit is contained in:
Tyler Michael Smith
2026-03-28 18:58:43 -04:00
co-authored by Claude Opus 4.6
parent 98502f1d64
commit 171eb482e8
2 changed files with 108 additions and 23 deletions
+89 -23
View File
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Zero-overhead NaN/Inf detection via RMSNorm kernel instrumentation."""
"""Zero-overhead NaN/Inf detection via RMSNorm kernel instrumentation
and pluggable tensor checks."""
from __future__ import annotations
@@ -20,15 +21,19 @@ def _as_fp8(data: torch.Tensor) -> torch.Tensor:
class NaNDetector:
"""Manages per-token NaN/Inf detection flags for RMSNorm kernels.
"""Manages per-token NaN/Inf detection flags.
Singleton. Created lazily when VLLM_NAN_DETECT=1.
Singleton. Created lazily when ``VLLM_NAN_DETECT=1``.
The flag array has shape ``int8[num_layers, max_num_tokens]``.
Each RMSNorm CUDA kernel block (one per token) writes
``flag[layer_idx][blockIdx.x] = 1`` when ``isnan(variance) ||
isinf(variance)`` after the CUB reduction — zero cost when the
flag pointer is NULL (the default).
The flag array has shape ``int8[num_checkpoints, max_num_tokens]``.
Checkpoints can be:
* **RMSNorm layers** -- the CUDA kernel writes flags via a pointer
argument (zero-cost when disabled).
* **Arbitrary tensors** -- call :meth:`check_tensor` which uses
``torch.isfinite`` to check for NaN/Inf. CUDA-graph compatible.
Both share the same flag array and reporting path.
"""
_instance: NaNDetector | None = None
@@ -53,18 +58,31 @@ class NaNDetector:
"""Reset singleton (for testing)."""
cls._instance = None
# ------------------------------------------------------------------
# Registration (before finalize)
# ------------------------------------------------------------------
def register(self, name: str) -> int:
"""Called from ``RMSNorm.__init__``. Returns layer index."""
"""Register a checkpoint. Returns its index into the flag array.
Works for both RMSNorm layers (which pass the index to the CUDA
kernel) and arbitrary tensor checks (which pass it to
:meth:`check_tensor`).
"""
assert not self._finalized, (
"Cannot register new layers after NaNDetector.finalize()"
"Cannot register new checkpoints after NaNDetector.finalize()"
)
idx = self._counter
self._layer_names[idx] = name
self._counter += 1
return idx
# ------------------------------------------------------------------
# Properties
# ------------------------------------------------------------------
@property
def num_layers(self) -> int:
def num_checkpoints(self) -> int:
return self._counter
@property
@@ -75,29 +93,37 @@ class NaNDetector:
def max_num_tokens(self) -> int:
return self._max_num_tokens
# ------------------------------------------------------------------
# Lifecycle
# ------------------------------------------------------------------
def finalize(
self,
device: torch.device,
max_num_tokens: int,
kv_caches: list[torch.Tensor] | None = None,
) -> None:
"""Allocate ``int8[num_layers, max_num_tokens]`` flag tensors."""
"""Allocate ``int8[num_checkpoints, max_num_tokens]`` flag tensors."""
if self._finalized:
return
n = self._counter
if n == 0:
logger.warning("NaNDetector.finalize() called but no layers registered")
logger.warning(
"NaNDetector.finalize() called but nothing registered"
)
return
self._max_num_tokens = max_num_tokens
self._nan_flags = torch.zeros(
n, max_num_tokens, dtype=torch.int8, device=device
)
self._host_flags = torch.zeros(n, max_num_tokens, dtype=torch.int8).pin_memory()
self._host_flags = torch.zeros(
n, max_num_tokens, dtype=torch.int8
).pin_memory()
if kv_caches is not None:
self._kv_caches = kv_caches
self._finalized = True
logger.info(
"NaN/Inf detector initialized: %d layers, max %d tokens "
"NaN/Inf detector initialized: %d checkpoints, max %d tokens "
"(%.1f KB flag buffer)",
n,
max_num_tokens,
@@ -109,16 +135,50 @@ class NaNDetector:
from vllm.model_executor.layers.layernorm import RMSNorm
for name, module in model.named_modules():
if isinstance(module, RMSNorm) and hasattr(module, "_nan_detect_layer_idx"):
if isinstance(module, RMSNorm) and hasattr(
module, "_nan_detect_layer_idx"
):
idx = module._nan_detect_layer_idx
if idx in self._layer_names:
self._layer_names[idx] = name
if hasattr(module, "_nan_detect_indices"):
for attr_label, idx in module._nan_detect_indices.items():
if idx in self._layer_names:
self._layer_names[idx] = f"{name}.{attr_label}"
# ------------------------------------------------------------------
# Per-step operations
# ------------------------------------------------------------------
def clear(self) -> None:
"""Zero flags before each forward pass."""
if self._nan_flags is not None:
self._nan_flags.zero_()
def check_tensor(
self, tensor: torch.Tensor, checkpoint_idx: int
) -> None:
"""Check *tensor* for NaN/Inf, writing per-token flags.
Uses ``torch.isfinite`` -- all ops stay on GPU, no D2H sync.
CUDA-graph compatible (fixed output address).
Args:
tensor: 2-D ``[num_tokens, hidden_size]`` tensor to check.
checkpoint_idx: index returned by :meth:`register`.
"""
if self._nan_flags is None:
return
num_tokens = tensor.shape[0]
has_bad = (~torch.isfinite(tensor.view(num_tokens, -1))).any(dim=1)
self._nan_flags[checkpoint_idx, :num_tokens].bitwise_or_(
has_bad.to(torch.int8)
)
# ------------------------------------------------------------------
# Post-forward checking
# ------------------------------------------------------------------
def check(self, num_real_tokens: int) -> None:
"""D2H copy flags, scan, log results.
@@ -140,10 +200,12 @@ class NaNDetector:
token_positions = (
real_flags[layer_idx].nonzero(as_tuple=True)[0].tolist()
)
name = self._layer_names.get(layer_idx, f"layer_{layer_idx}")
name = self._layer_names.get(
layer_idx, f"checkpoint_{layer_idx}"
)
logger.error(
"NaN/Inf detected in real tokens at '%s' "
"(layer %d), token positions: %s",
"(checkpoint %d), token positions: %s",
name,
layer_idx,
token_positions,
@@ -152,17 +214,21 @@ class NaNDetector:
pad_bad = pad_flags.any(dim=1).nonzero(as_tuple=True)[0]
if len(pad_bad) > 0:
logger.debug(
"NaN/Inf in padding tokens at %d layers",
"NaN/Inf in padding tokens at %d checkpoints",
len(pad_bad),
)
if len(real_bad) > 0:
self._check_all_kv_cache()
raise RuntimeError(
f"NaN/Inf detected in {len(real_bad)} layer(s). "
f"NaN/Inf detected at {len(real_bad)} checkpoint(s). "
"See ERROR logs above for details."
)
# ------------------------------------------------------------------
# KV cache checks
# ------------------------------------------------------------------
def _check_all_kv_cache(self) -> None:
"""Scan all KV cache blocks for NaN. Called once before crash."""
if not self._kv_caches:
@@ -171,9 +237,9 @@ class NaNDetector:
if not isinstance(kv_cache, torch.Tensor):
continue
data = _as_fp8(kv_cache)
nan_per_block = (
torch.isnan(data.view(data.shape[0], -1)).any(dim=1)
)
nan_per_block = torch.isnan(
data.view(data.shape[0], -1)
).any(dim=1)
bad_blocks = nan_per_block.nonzero(as_tuple=True)[0]
if len(bad_blocks) > 0:
logger.error(
+19
View File
@@ -33,6 +33,7 @@ from torch import nn
from transformers import DeepseekV2Config, DeepseekV3Config
import vllm._custom_ops as ops
import vllm.envs as envs
from vllm._aiter_ops import rocm_aiter_ops
from vllm.compilation.decorators import support_torch_compile
from vllm.config import CacheConfig, ParallelConfig, VllmConfig, get_current_vllm_config
@@ -534,6 +535,16 @@ class DeepseekV2Attention(nn.Module):
prefix=f"{prefix}.attn",
)
self.prefix = prefix
if envs.VLLM_NAN_DETECT:
from vllm.model_executor.layers.nan_detector import NaNDetector
self._nan_detect_indices = {
"attn_output": NaNDetector.get().register(
f"{prefix}.attn_output"
),
}
def forward(
self,
positions: torch.Tensor,
@@ -577,6 +588,14 @@ class DeepseekV2Attention(nn.Module):
attn_output = attn_output.view(-1, self.num_local_heads, self.qk_head_dim)[
..., : self.v_head_dim
].reshape(-1, self.num_local_heads * self.v_head_dim)
if envs.VLLM_NAN_DETECT and hasattr(self, "_nan_detect_indices"):
from vllm.model_executor.layers.nan_detector import NaNDetector
NaNDetector.get().check_tensor(
attn_output, self._nan_detect_indices["attn_output"]
)
output, _ = self.o_proj(attn_output)
return output