Gate Mamba slot-donation debug asserts behind SGLANG_MAMBA_DEBUG_ASSERTS (#31982)
This commit is contained in:
@@ -19,6 +19,7 @@ limitations under the License.
|
|||||||
The radix tree data structure for managing the hybrid (full and Mamba) KV cache.
|
The radix tree data structure for managing the hybrid (full and Mamba) KV cache.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import os
|
||||||
from array import array
|
from array import array
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from typing import TYPE_CHECKING, List, Optional, Tuple
|
from typing import TYPE_CHECKING, List, Optional, Tuple
|
||||||
@@ -61,6 +62,12 @@ from sglang.srt.runtime_context import get_parallel
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Debug-only invariant checks in the Mamba slot-donation path call tensor.item(),
|
||||||
|
# which forces a per-request cudaStreamSynchronize on the scheduler thread. Under
|
||||||
|
# load this can serialize/stall the scheduler. Gate them off by default; set
|
||||||
|
# SGLANG_MAMBA_DEBUG_ASSERTS=1 to re-enable for debugging.
|
||||||
|
_MAMBA_DEBUG_ASSERTS = os.environ.get("SGLANG_MAMBA_DEBUG_ASSERTS", "0") == "1"
|
||||||
|
|
||||||
|
|
||||||
class TreeNode:
|
class TreeNode:
|
||||||
|
|
||||||
@@ -592,6 +599,8 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache):
|
|||||||
src_active = req.mamba_ping_pong_track_buffer[
|
src_active = req.mamba_ping_pong_track_buffer[
|
||||||
mamba_ping_pong_track_buffer_to_keep
|
mamba_ping_pong_track_buffer_to_keep
|
||||||
].unsqueeze(-1)
|
].unsqueeze(-1)
|
||||||
|
if _MAMBA_DEBUG_ASSERTS:
|
||||||
|
# .item() forces a cudaStreamSynchronize; only pay it when debugging.
|
||||||
assert src_active.item() != -1, (
|
assert src_active.item() != -1, (
|
||||||
f"Cached mamba slot is -1: keep_idx={mamba_ping_pong_track_buffer_to_keep}, "
|
f"Cached mamba slot is -1: keep_idx={mamba_ping_pong_track_buffer_to_keep}, "
|
||||||
f"buf={req.mamba_ping_pong_track_buffer.tolist()}, "
|
f"buf={req.mamba_ping_pong_track_buffer.tolist()}, "
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ import copy
|
|||||||
import dataclasses
|
import dataclasses
|
||||||
import logging
|
import logging
|
||||||
import math
|
import math
|
||||||
|
import os
|
||||||
from contextlib import contextmanager, nullcontext
|
from contextlib import contextmanager, nullcontext
|
||||||
from dataclasses import dataclass, fields
|
from dataclasses import dataclass, fields
|
||||||
from functools import cached_property
|
from functools import cached_property
|
||||||
@@ -93,6 +94,12 @@ if TYPE_CHECKING:
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# Debug-only invariant in the Mamba slot-donation path calls tensor.item(), which
|
||||||
|
# forces a per-request cudaStreamSynchronize on the scheduler thread and can stall
|
||||||
|
# the scheduler under load. Off by default; set SGLANG_MAMBA_DEBUG_ASSERTS=1 to
|
||||||
|
# re-enable for debugging.
|
||||||
|
_MAMBA_DEBUG_ASSERTS = os.environ.get("SGLANG_MAMBA_DEBUG_ASSERTS", "0") == "1"
|
||||||
|
|
||||||
GB = 1024 * 1024 * 1024
|
GB = 1024 * 1024 * 1024
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
@@ -1404,6 +1411,8 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
|||||||
mamba_value_donated = (
|
mamba_value_donated = (
|
||||||
req.mamba_ping_pong_track_buffer[donate_idx].unsqueeze(-1).clone()
|
req.mamba_ping_pong_track_buffer[donate_idx].unsqueeze(-1).clone()
|
||||||
)
|
)
|
||||||
|
if _MAMBA_DEBUG_ASSERTS:
|
||||||
|
# .item() forces a cudaStreamSynchronize; only pay it when debugging.
|
||||||
assert mamba_value_donated.item() != -1, (
|
assert mamba_value_donated.item() != -1, (
|
||||||
f"Donated mamba slot is -1: donate_idx={donate_idx}, "
|
f"Donated mamba slot is -1: donate_idx={donate_idx}, "
|
||||||
f"buf={req.mamba_ping_pong_track_buffer.tolist()}, "
|
f"buf={req.mamba_ping_pong_track_buffer.tolist()}, "
|
||||||
|
|||||||
Reference in New Issue
Block a user