[Spec] Guard async-assert probes against None tensor (#27478)
This commit is contained in:
@@ -5,30 +5,38 @@ When the gate is on, a violation surfaces as an assertion at the next CUDA
|
|||||||
sync point instead of as a silent NaN cascade or illegal-address crash.
|
sync point instead of as a silent NaN cascade or illegal-address crash.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
|
|
||||||
|
|
||||||
def maybe_detect_nan(tensor: torch.Tensor, msg: str = ""):
|
def maybe_detect_nan(tensor: Optional[torch.Tensor], msg: str = ""):
|
||||||
"""Async NaN check — no GPU-CPU sync, error surfaces at next sync point."""
|
"""Async NaN check — no GPU-CPU sync, error surfaces at next sync point."""
|
||||||
if not envs.SGLANG_ENABLE_ASYNC_ASSERT.get():
|
if not envs.SGLANG_ENABLE_ASYNC_ASSERT.get():
|
||||||
return
|
return
|
||||||
|
# A None tensor means there is nothing to probe, e.g. hidden_states on
|
||||||
|
# capture_hidden_mode=NULL paths (STANDALONE speculative decoding).
|
||||||
|
if tensor is None:
|
||||||
|
return
|
||||||
torch._assert_async(~torch.any(torch.isnan(tensor)), f"NaN detected! {msg}")
|
torch._assert_async(~torch.any(torch.isnan(tensor)), f"NaN detected! {msg}")
|
||||||
|
|
||||||
|
|
||||||
def maybe_detect_inf(tensor: torch.Tensor, msg: str = ""):
|
def maybe_detect_inf(tensor: Optional[torch.Tensor], msg: str = ""):
|
||||||
"""Async Inf check — fp16 overflow surfaces as Inf before NaN."""
|
"""Async Inf check — fp16 overflow surfaces as Inf before NaN."""
|
||||||
if not envs.SGLANG_ENABLE_ASYNC_ASSERT.get():
|
if not envs.SGLANG_ENABLE_ASYNC_ASSERT.get():
|
||||||
return
|
return
|
||||||
|
if tensor is None:
|
||||||
|
return
|
||||||
torch._assert_async(~torch.any(torch.isinf(tensor)), f"Inf detected! {msg}")
|
torch._assert_async(~torch.any(torch.isinf(tensor)), f"Inf detected! {msg}")
|
||||||
|
|
||||||
|
|
||||||
def maybe_detect_oob(indices: torch.Tensor, low: int, high: int, msg: str):
|
def maybe_detect_oob(indices: Optional[torch.Tensor], low: int, high: int, msg: str):
|
||||||
"""Async OOB check — no GPU-CPU sync, error surfaces at next sync point."""
|
"""Async OOB check — no GPU-CPU sync, error surfaces at next sync point."""
|
||||||
if not envs.SGLANG_ENABLE_ASYNC_ASSERT.get():
|
if not envs.SGLANG_ENABLE_ASYNC_ASSERT.get():
|
||||||
return
|
return
|
||||||
if indices.numel() == 0:
|
if indices is None or indices.numel() == 0:
|
||||||
return
|
return
|
||||||
torch._assert_async(
|
torch._assert_async(
|
||||||
(indices.min() >= low) & (indices.max() < high),
|
(indices.min() >= low) & (indices.max() < high),
|
||||||
@@ -36,11 +44,13 @@ def maybe_detect_oob(indices: torch.Tensor, low: int, high: int, msg: str):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def maybe_detect_page_aligned(indices: torch.Tensor, page_size: int, msg: str):
|
def maybe_detect_page_aligned(
|
||||||
|
indices: Optional[torch.Tensor], page_size: int, msg: str
|
||||||
|
):
|
||||||
"""Async page-alignment check on slot ids."""
|
"""Async page-alignment check on slot ids."""
|
||||||
if not envs.SGLANG_ENABLE_ASYNC_ASSERT.get():
|
if not envs.SGLANG_ENABLE_ASYNC_ASSERT.get():
|
||||||
return
|
return
|
||||||
if indices.numel() == 0 or page_size <= 1:
|
if indices is None or indices.numel() == 0 or page_size <= 1:
|
||||||
return
|
return
|
||||||
torch._assert_async(
|
torch._assert_async(
|
||||||
(indices % page_size == 0).all(),
|
(indices % page_size == 0).all(),
|
||||||
|
|||||||
Reference in New Issue
Block a user