Add NaN/Inf detection for NIXL KV cache transfers

Gate behind VLLM_NIXL_NAN_DETECT=1 env var. Checks KV cache blocks
for NaN on the decoder side after recv completes and on the prefiller
side after send is confirmed. Handles uint8-stored fp8 KV caches
(MLA cross-layer) by viewing as float8_e4m3fn before isnan check.

Uses a fast two-pass approach: single torch.isnan().any() across all
layers first, only doing per-layer breakdown if NaN is found.

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 13:12:23 -04:00
co-authored by Claude Opus 4.6
parent ef2c4f778d
commit 7779fccdfd
2 changed files with 106 additions and 0 deletions
@@ -1135,6 +1135,7 @@ class NixlConnectorWorker:
# In progress transfers.
# [req_id -> list[handle]]
self._recving_metadata: dict[ReqId, ReqMeta] = {}
self._sending_metadata: dict[ReqId, ReqMeta] = {}
self._recving_transfers = defaultdict[ReqId, list[TransferHandle]](list)
# Track the expiration time of requests that are waiting to be sent.
self._reqs_to_send: dict[ReqId, float] = {}
@@ -2224,6 +2225,82 @@ class NixlConnectorWorker:
cache, indices, block_size_ratio
)
@staticmethod
def _as_fp8(data: torch.Tensor) -> torch.Tensor:
"""View uint8 KV cache data as fp8 so torch.isnan works."""
if data.dtype == torch.uint8:
return data.view(torch.float8_e4m3fn)
return data
def _check_kv_blocks_for_nan(
self, req_id: str, block_ids: BlockIds, direction: str
):
"""Check KV cache blocks for NaN values after transfer.
Uses a fast two-pass approach: first check all blocks across all
layers with a single torch.isnan, then only do the expensive
per-layer breakdown if something is found.
"""
all_group_blocks = [g for g in block_ids if len(g) > 0]
if not all_group_blocks:
return
# Fast pass: check all layers at once.
has_nan = False
for cache_or_caches in self.device_kv_caches.values():
caches = (
[cache_or_caches]
if isinstance(cache_or_caches, torch.Tensor)
else cache_or_caches
)
for cache in caches:
for group_blocks in all_group_blocks:
indices = torch.tensor(
group_blocks, device=cache.device, dtype=torch.long
)
if torch.isnan(
self._as_fp8(cache[indices])
).any().item():
has_nan = True
break
if has_nan:
break
if has_nan:
break
if not has_nan:
return
# Slow pass: per-layer breakdown for diagnosis.
for layer_name, cache_or_caches in self.device_kv_caches.items():
caches = (
[cache_or_caches]
if isinstance(cache_or_caches, torch.Tensor)
else cache_or_caches
)
for cache in caches:
for group_blocks in all_group_blocks:
indices = torch.tensor(
group_blocks, device=cache.device, dtype=torch.long
)
blocks_data = self._as_fp8(cache[indices])
nan_count = torch.isnan(blocks_data).sum().item()
if nan_count > 0:
total_elements = blocks_data.numel()
logger.error(
"*** NaN DETECTED in KV cache during %s *** "
"req_id=%s, layer=%s, blocks=%s, "
"nan_count=%d, total_elements=%d, "
"nan_pct=%.4f%%",
direction,
req_id,
layer_name,
group_blocks,
nan_count,
total_elements,
100.0 * nan_count / total_elements,
)
def get_finished(self) -> tuple[set[str], set[str]]:
"""
Get requests that are done sending or recving on this specific worker.
@@ -2256,6 +2333,11 @@ class NixlConnectorWorker:
if self.use_host_buffer:
self.sync_recved_kv_to_device(req_id, meta)
if envs.VLLM_NIXL_NAN_DETECT:
self._check_kv_blocks_for_nan(
req_id, meta.local_physical_block_ids, "recv"
)
# post processing for heteroblocksize
block_size_ratio = self.kv_topo.block_size_ratio_from_engine_id(
meta.remote.engine_id
@@ -2291,6 +2373,7 @@ class NixlConnectorWorker:
)
self._reqs_to_process.remove(req_id)
del self._reqs_to_send[req_id]
self._sending_metadata.pop(req_id, None)
done_sending.add(req_id)
return done_sending, done_recving
@@ -2338,6 +2421,15 @@ class NixlConnectorWorker:
del self.consumer_notification_counts_by_req[req_id]
self._reqs_to_process.remove(req_id)
self._reqs_to_send.pop(req_id, None)
if envs.VLLM_NIXL_NAN_DETECT:
send_meta = self._sending_metadata.pop(req_id, None)
if send_meta is not None:
self._check_kv_blocks_for_nan(
req_id,
send_meta.local_physical_block_ids,
"send",
)
return notified_req_ids
def _pop_done_transfers(self, transfers: dict[str, list[int]]) -> set[str]:
@@ -2461,6 +2553,14 @@ class NixlConnectorWorker:
if req_id in self._reqs_to_process:
self._reqs_to_send[req_id] = expiration_time
# Track send-side metadata for NaN detection.
if envs.VLLM_NIXL_NAN_DETECT:
for req_id, meta in metadata.reqs_to_save.items():
meta.local_physical_block_ids = (
self._logical_to_kernel_block_ids(meta.local_block_ids)
)
self._sending_metadata[req_id] = meta
def _read_blocks_for_req(self, req_id: str, meta: ReqMeta):
assert meta.remote is not None and self.kv_topo is not None
remote_ranks = self.kv_topo.get_target_remote_ranks_from_engine_id(
+6
View File
@@ -190,6 +190,7 @@ if TYPE_CHECKING:
VLLM_ROCM_QUICK_REDUCE_CAST_BF16_TO_FP16: bool = True
VLLM_ROCM_QUICK_REDUCE_MAX_SIZE_BYTES_MB: int | None = None
VLLM_NIXL_ABORT_REQUEST_TIMEOUT: int = 480
VLLM_NIXL_NAN_DETECT: bool = False
VLLM_MORIIO_CONNECTOR_READ_MODE: bool = False
VLLM_MORIIO_QP_PER_TRANSFER: int = 1
VLLM_MORIIO_POST_BATCH_SIZE: int = -1
@@ -1385,6 +1386,11 @@ environment_variables: dict[str, Callable[[], Any]] = {
"VLLM_NIXL_ABORT_REQUEST_TIMEOUT": lambda: int(
os.getenv("VLLM_NIXL_ABORT_REQUEST_TIMEOUT", "480")
),
# Enable NaN/Inf detection in KV cache blocks during NIXL transfers.
# Logs errors when NaN/Inf values are found on send or receive side.
"VLLM_NIXL_NAN_DETECT": lambda: bool(
int(os.getenv("VLLM_NIXL_NAN_DETECT", "0"))
),
# Controls the read mode for the Mori-IO connector
"VLLM_MORIIO_CONNECTOR_READ_MODE": lambda: (
os.getenv("VLLM_MORIIO_CONNECTOR_READ_MODE", "False").lower() in ("true", "1")