diff --git a/docs/docs/advanced_features/server_arguments.mdx b/docs/docs/advanced_features/server_arguments.mdx index 7117a7173..38353a70c 100644 --- a/docs/docs/advanced_features/server_arguments.mdx +++ b/docs/docs/advanced_features/server_arguments.mdx @@ -673,6 +673,12 @@ Please consult the documentation below and [server_args.py](https://github.com/s The random seed. None Type: int + + + `--mlx-enable-sampling` + MLX backend only: sample decode tokens (temperature / top-k / top-p / min-p) instead of greedy argmax. Sampling runs inside the lazy MLX graph, so it works with the overlap scheduler; first tokens from prefill/extend are sampled too. Greedy requests keep exact argmax behavior. Also enables on the MLX path: grammar vocab masks and custom logit processors (these break decode chaining per step; custom processors run on pure-decode steps only), logit_bias, output logprobs (sampled token / top-k / token_ids; prompt input logprobs are not computed), NaN sanitization (SGLANG_SANITIZE_NAN_LOGITS), and per-request sampling_seed under --enable-deterministic-inference (deterministic within MLX only). Penalties are not applied. + `False` + bool flag (set to enable) `--constrained-json-whitespace-pattern` diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index 88dcd9788..ca931e80c 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -61,6 +61,7 @@ from sglang.srt.utils.common import ( is_xpu, xpu_has_xmx_support, ) +from sglang.srt.utils.tensor_bridge import use_mlx logger = logging.getLogger(__name__) @@ -933,8 +934,12 @@ def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict: overrides["attention_backend"] = "intel_xpu" elif is_hip(): overrides["attention_backend"] = "aiter" - elif not is_mps(): - # No triton on macOS; MPS keeps the platform default. + elif not (is_mps() and use_mlx()): + # No triton on macOS, but only the MLX runner can actually serve + # gpt-oss there -- it owns attention, so it keeps the platform + # default. macOS *without* MLX must still fall through to triton + # and fail fast below: torch_native has neither sliding-window nor + # attention-sink support, so accepting it would silently mis-serve. overrides["attention_backend"] = "triton" if is_xpu(): # Check for bf16 dtype on Intel XPU. Reads the pristine dtype request, diff --git a/python/sglang/srt/constrained/xgrammar_backend.py b/python/sglang/srt/constrained/xgrammar_backend.py index a0b44930c..e4ddcb55c 100644 --- a/python/sglang/srt/constrained/xgrammar_backend.py +++ b/python/sglang/srt/constrained/xgrammar_backend.py @@ -38,6 +38,7 @@ from sglang.srt.constrained.base_grammar_backend import ( ) from sglang.srt.constrained.utils import is_legacy_structural_tag from sglang.srt.utils import is_hip +from sglang.srt.utils.common import is_pin_memory_available _is_hip = is_hip() @@ -58,14 +59,16 @@ MAX_ROLLBACK_TOKENS = 200 def _allocate_token_bitmask(vocab_size: int, batch_size: int) -> torch.Tensor: - # Always allocate a pinned bitmask so the later H2D to the device can be a - # genuine non_blocking copy (a pageable source silently downgrades it to a - # blocking copy). + # Allocate a pinned bitmask where pinning exists so the later H2D to the + # device can be a genuine non_blocking copy (a pageable source silently + # downgrades it to a blocking copy). MPS torch has no pin-memory kernel + # and asserts on pin_memory=True; the MLX path consumes the mask on the + # CPU anyway. return torch.full( get_bitmask_shape(batch_size, vocab_size), -1, dtype=bitmask_dtype, - pin_memory=True, + pin_memory=is_pin_memory_available(), ) @@ -132,6 +135,12 @@ class XGrammarGrammar(BaseGrammarObject): import sgl_kernel_npu # noqa: F401 torch.ops.npu.apply_token_bitmask(logits, vocab_mask) + elif logits.device.type == "cpu": + # Used by the MLX backend, which builds its additive mask rows + # on the CPU before inserting them into the lazy graph. + from xgrammar import apply_token_bitmask_inplace + + apply_token_bitmask_inplace(logits, vocab_mask, backend="cpu") else: raise RuntimeError(f"Unsupported device: {logits.device.type}") diff --git a/python/sglang/srt/hardware_backend/mlx/aot.py b/python/sglang/srt/hardware_backend/mlx/aot.py index af9784e99..890ad49eb 100644 --- a/python/sglang/srt/hardware_backend/mlx/aot.py +++ b/python/sglang/srt/hardware_backend/mlx/aot.py @@ -4,7 +4,7 @@ from __future__ import annotations import logging from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any, Callable, Optional +from typing import Any, Callable, Optional import mlx.core as mx @@ -12,11 +12,6 @@ from sglang.srt.environ import envs logger = logging.getLogger(__name__) -if TYPE_CHECKING: - from sglang.srt.hardware_backend.mlx.kv_cache.attention_kv_cache import ( - ContiguousAttentionKVCache, - ) - def _load_metal_rope_pool_fused(): try: @@ -219,7 +214,10 @@ class MlxAOTKernelContext: req_ids: list[str], req_pool_idx: dict[str, int], req_to_token_pool: Any | None, - layer_caches: list[list[ContiguousAttentionKVCache]], + # Only .offset is read (absolute on every cache kind) and the slot + # lookup is layer-agnostic; the wrapper gates the fused pool scatter + # to full-attention layers. + layer_caches: list[list[Any]], ) -> MlxAOTKernelContext: """Build optional AOT context for one batched decode step.""" if not aot_kernels.rope.enabled or kv_pool is None: diff --git a/python/sglang/srt/hardware_backend/mlx/kv_cache/__init__.py b/python/sglang/srt/hardware_backend/mlx/kv_cache/__init__.py index 985db1e7b..339573dc7 100644 --- a/python/sglang/srt/hardware_backend/mlx/kv_cache/__init__.py +++ b/python/sglang/srt/hardware_backend/mlx/kv_cache/__init__.py @@ -14,6 +14,7 @@ from sglang.srt.hardware_backend.mlx.kv_cache.attention_kv_cache import ( AttentionOffsetCache, ContiguousAttentionKVCache, PoolBackedAttentionKVCache, + WindowedAttentionKVCache, make_attention_mask, ) from sglang.srt.hardware_backend.mlx.kv_cache.attention_kv_pool import ( @@ -64,4 +65,5 @@ __all__ = [ "PoolBackedAttentionKVCache", "set_context", "uses_sliding_window_attention", + "WindowedAttentionKVCache", ] diff --git a/python/sglang/srt/hardware_backend/mlx/kv_cache/attention_contract.py b/python/sglang/srt/hardware_backend/mlx/kv_cache/attention_contract.py index d96b790c6..19414dc0b 100644 --- a/python/sglang/srt/hardware_backend/mlx/kv_cache/attention_contract.py +++ b/python/sglang/srt/hardware_backend/mlx/kv_cache/attention_contract.py @@ -68,11 +68,15 @@ def is_attention_module(module: Any) -> bool: ) +def language_model_container(model: Any) -> Any: + """The mlx-lm container carrying the layer list and model-level config.""" + root = getattr(model, "language_model", model) + return getattr(root, "model", root) + + def get_container_window_size(model: Any) -> int | None: """The container-level scalar sliding window, if the model declares one.""" - root = getattr(model, "language_model", model) - container = getattr(root, "model", root) - return first_present_attr(container, WINDOW_SIZE_ATTRS) + return first_present_attr(language_model_container(model), WINDOW_SIZE_ATTRS) def get_layer_window_sizes(model: Any) -> dict[int, int | None]: @@ -84,10 +88,9 @@ def get_layer_window_sizes(model: Any) -> dict[int, int | None]: ``{layer_idx: window or None}``, or ``{}`` when the model does not follow the convention. """ - root = getattr(model, "language_model", model) - container = getattr(root, "model", root) + container = language_model_container(model) layer_types = getattr(container, "layer_types", None) - window_size = get_container_window_size(model) + window_size = first_present_attr(container, WINDOW_SIZE_ATTRS) if not layer_types or window_size is None: return {} return { diff --git a/python/sglang/srt/hardware_backend/mlx/kv_cache/attention_kv_cache.py b/python/sglang/srt/hardware_backend/mlx/kv_cache/attention_kv_cache.py index 0317c85ab..6f00e46a1 100644 --- a/python/sglang/srt/hardware_backend/mlx/kv_cache/attention_kv_cache.py +++ b/python/sglang/srt/hardware_backend/mlx/kv_cache/attention_kv_cache.py @@ -21,8 +21,19 @@ def make_attention_mask(N, offset, return_array=False, window_size=None): layers pass it, including for N == 1) or windowed models silently fall back to full attention. """ - if window_size is not None: + if window_size is not None and offset + N > window_size: return create_causal_mask(N, offset, window_size=window_size) + # Either no window, or a window that cannot bind. The lowest query + # position is ``offset``, so ``offset + N <= window_size`` means every + # causally visible key is inside the band and the banded mask is + # elementwise identical to a plain causal one -- the shortcut mlx_lm's own + # RotatingKVCache.make_mask takes, and these shims stand in for exactly + # that cache on sliding layers. Worth the branch because a materialised + # mask forces mx.fast.scaled_dot_product_attention off its fused causal + # path: ~2x slower per layer, plus an N x (offset + N) allocation. + # It also survives a window-bounded store: offset + N <= window_size + # implies N + window_size - 1 >= offset + N, so no keys are dropped and + # the mask width still matches the returned key length. if N == 1: return None if return_array: @@ -141,9 +152,130 @@ class ContiguousAttentionKVCache: self.values[:, :, self.offset : end, :] = v self.offset = end - def get_kv(self) -> tuple[mx.array, mx.array]: - """Return valid K/V: (1, n_kv_heads, offset, head_dim).""" - return self.keys[:, :, : self.offset, :], self.values[:, :, : self.offset, :] + def get_kv(self, window: int | None = None) -> tuple[mx.array, mx.array]: + """Return valid K/V: (1, n_kv_heads, min(offset, window), head_dim). + + ``window`` keeps only the trailing window a sliding-window layer can + attend to. Slicing here rather than slicing the full history and then + slicing again costs one op instead of two per request per layer. + """ + start = 0 if window is None else max(0, self.offset - window) + return ( + self.keys[:, :, start : self.offset, :], + self.values[:, :, start : self.offset, :], + ) + + def reset(self) -> None: + """Reset for reuse, keeping allocated buffers.""" + self.offset = 0 + + +class WindowedAttentionKVCache: + """Sliding-window attention KV buffer for one request and one layer. + + Holds the trailing ``window`` tokens plus the in-flight chunk, in + temporal order, instead of the full sequence. ``offset`` stays + absolute (RoPE positions, decode bookkeeping); the dropped prefix + shows up only in the shorter arrays returned by + ``update_and_fetch``/``get_kv`` and in the mask offset ``make_mask`` + clamps to, so mask width always equals returned key length. + """ + + __slots__ = ("keys", "values", "offset", "window", "_local") + + def __init__(self, window: int): + self.window = window + self.keys: mx.array | None = None + self.values: mx.array | None = None + self.offset = 0 # absolute: every token ever written + self._local = 0 # tokens currently in the buffer + + @property + def state(self): + """Arrays for ``mx.eval`` unpacking.""" + if self.keys is None: + return () + return (self.keys, self.values) + + def reset(self) -> None: + """Reset for reuse, keeping allocated buffers.""" + self.offset = 0 + self._local = 0 + + def make_mask(self, N, return_array=False, window_size=None, **kwargs): + kept = min(self._local, self.window) + if window_size is None and self.offset > kept: + raise RuntimeError( + "WindowedAttentionKVCache holds only the trailing window and " + "cannot serve a full-context attention mask" + ) + # No N == 1 shortcut here, tempting as it looks: mlx_lm's banded mask is + # ``linds < rinds + window_size`` (strict), so a window of W admits + # exactly W keys. Once ``kept == window`` this buffer returns W + 1 of + # them -- the trailing window plus the token just written -- and the + # oldest must still be masked out. + return make_attention_mask( + N, kept, return_array=return_array, window_size=window_size + ) + + def _append(self, keys: mx.array, values: mx.array) -> tuple[int, int]: + """Append a chunk in place; return the (start, end) span it serves. + + Split out from ``update_and_fetch`` so the decode path can skip + building the two return slices, which its caller discards in + favour of ``get_kv``. + """ + S = keys.shape[2] + kept = min(self._local, self.window) + capacity = self.window + max(S, self.window) + held = self.keys.shape[2] if self.keys is not None else 0 + if self._local + S > held or held > capacity: + # Compact the trailing window into a right-sized buffer: this + # allocates on the first write, drops history when the buffer + # fills (amortised O(1) per decode token), and shrinks back to + # 2 * window once an oversized prefill chunk is behind us. + B, n_kv_heads, _, head_dim = keys.shape + new_k = mx.zeros((B, n_kv_heads, capacity, head_dim), dtype=keys.dtype) + new_v = mx.zeros((B, n_kv_heads, capacity, head_dim), dtype=keys.dtype) + if kept: + src = slice(self._local - kept, self._local) + new_k[:, :, :kept, :] = self.keys[:, :, src, :] + new_v[:, :, :kept, :] = self.values[:, :, src, :] + self.keys, self.values, self._local = new_k, new_v, kept + start, end = self._local - kept, self._local + S + self.keys[:, :, self._local : end, :] = keys + self.values[:, :, self._local : end, :] = values + self._local = end + self.offset += S + return start, end + + def update_and_fetch( + self, keys: mx.array, values: mx.array + ) -> tuple[mx.array, mx.array]: + """Append a chunk and return the kept trailing window plus the chunk. + + The kept prefix is ``min(local, window)``, matching what + ``make_mask`` clamps to earlier in the same forward pass. + """ + start, end = self._append(keys, values) + return self.keys[:, :, start:end, :], self.values[:, :, start:end, :] + + def write_token(self, k: mx.array, v: mx.array) -> None: + """Write one token. k, v shape: (1, n_kv_heads, 1, head_dim).""" + self._append(k, v) + + def get_kv(self, window: int | None = None) -> tuple[mx.array, mx.array]: + """Return buffered trailing K/V: (1, n_kv_heads, kept, head_dim). + + ``window`` mirrors :meth:`ContiguousAttentionKVCache.get_kv`, but the + slice is buffer-relative: this buffer holds at most ``2 * window`` + tokens, so the trailing window starts from ``_local``, not ``offset``. + """ + start = 0 if window is None else max(0, self._local - window) + return ( + self.keys[:, :, start : self._local, :], + self.values[:, :, start : self._local, :], + ) class PoolBackedAttentionKVCache: diff --git a/python/sglang/srt/hardware_backend/mlx/kv_cache/attention_kv_pool.py b/python/sglang/srt/hardware_backend/mlx/kv_cache/attention_kv_pool.py index 8761438dc..6706965c5 100644 --- a/python/sglang/srt/hardware_backend/mlx/kv_cache/attention_kv_pool.py +++ b/python/sglang/srt/hardware_backend/mlx/kv_cache/attention_kv_pool.py @@ -1,10 +1,10 @@ """Flat attention KV pool for the MLX backend. Each layer buffer has shape ``(pool_size, n_kv_heads, head_dim)``. -This v1 pool is intentionally uniform: every wrapped softmax-attention -layer must share the same KV shape and full-context KV semantics. -Heterogeneous KV shapes and sliding-window KV need per-layer/window-aware -pools before they can use MLX radix reuse. +The pool stores full-attention layers only and stays intentionally +uniform: every stored layer must share the same KV shape and +full-context KV semantics. Sliding-window layers keep window-bounded +per-request caches instead; heterogeneous KV shapes need per-layer pools. Slot 0 is reserved as padding (1-based indexing). """ diff --git a/python/sglang/srt/hardware_backend/mlx/kv_cache/attention_wrapper.py b/python/sglang/srt/hardware_backend/mlx/kv_cache/attention_wrapper.py index 38dcdb2b8..89b134651 100644 --- a/python/sglang/srt/hardware_backend/mlx/kv_cache/attention_wrapper.py +++ b/python/sglang/srt/hardware_backend/mlx/kv_cache/attention_wrapper.py @@ -22,6 +22,7 @@ from sglang.srt.hardware_backend.mlx.kv_cache.attention_contract import ( ) from sglang.srt.hardware_backend.mlx.kv_cache.attention_kv_cache import ( ContiguousAttentionKVCache, + WindowedAttentionKVCache, ) _thread_local = threading.local() @@ -34,9 +35,15 @@ class BatchedDecodeContext: batch_size: int seq_lens: list[int] # per-request token count before the new token - # attention_layer_caches[attention_pool_idx][req_idx] = ContiguousAttentionKVCache - attention_layer_caches: list[list[ContiguousAttentionKVCache]] + # attention_layer_caches[cache_idx][req_idx], dense over attention layers. + # Windowed caches hold only trailing-window KV, so read them through + # write_token/get_kv, never by slicing .keys at an absolute offset. + attention_layer_caches: list[ + list[ContiguousAttentionKVCache | WindowedAttentionKVCache] + ] attention_pool_index_by_layer: dict[int, int] = field(default_factory=dict) + # Dense index into the shared pool's buffers; full-attention layers only. + full_kv_pool_index_by_layer: dict[int, int] = field(default_factory=dict) # Optional AOT kernel state. Keep kernel-specific fields out of the regular # MLX decode path so future AOT kernels can be added without growing this @@ -50,6 +57,11 @@ class BatchedDecodeContext: needs_padding: bool = field(init=False) pad_sizes: list[int] = field(init=False) positions: Optional[mx.array] = field(init=False) + # Padding metadata memo, keyed by window size. It depends only on + # ``seq_lens`` and the window, so every layer sharing a window reuses one + # entry instead of rebuilding it (gpt-oss decodes 24 attention layers per + # step, in two window classes). + _padding_by_window: dict = field(init=False, default_factory=dict) def __post_init__(self) -> None: seq_lens = self.seq_lens @@ -64,6 +76,56 @@ class BatchedDecodeContext: self.attention_pool_index_by_layer = { idx: idx for idx in range(len(self.attention_layer_caches)) } + if self.aot.rope is not None and not self.full_kv_pool_index_by_layer: + # No silent default here: the fused scatter addresses pool buffers + # by full-attention index, so falling back to the cache index would + # write the wrong buffer whenever sliding-window layers are + # interleaved. A model with a pool always has full layers to index. + raise ValueError( + "BatchedDecodeContext requires full_kv_pool_index_by_layer " + "when the fused AOT RoPE + pool-scatter kernel is active" + ) + + def decode_padding( + self, window: int | None + ) -> tuple[list[int], Optional[mx.array]]: + """Right-pad sizes and the keep-mask for one decode step. + + Requests are padded to a common KV width so they can be batched into + one SDPA call. Without a window that width is ``max_len``; a + sliding-window layer only reads the trailing ``window`` keys, so its + width is ``max(min(seq_len + 1, window))`` instead -- which is why the + context's full-length metadata cannot be reused for it. + + The mask is boolean (``True`` keeps the key), broadcast-shaped + ``(B, 1, 1, width)``, and ``None`` when no request needs padding. + Cached per window: all layers in the step share one build. + """ + cached = self._padding_by_window.get(window, None) + if cached is not None: + return cached + + if window is None: + pad_sizes = self.pad_sizes + keep = ( + self.positions[None, :] < self.valid_lens[:, None] + if self.needs_padding + else None + ) + else: + eff_lens = [min(n + 1, window) for n in self.seq_lens] + max_eff = max(eff_lens) + pad_sizes = [max_eff - n for n in eff_lens] + keep = ( + mx.arange(max_eff)[None, :] + < mx.array(eff_lens, dtype=mx.int32)[:, None] + if min(eff_lens) < max_eff + else None + ) + + result = (pad_sizes, None if keep is None else keep[:, None, None, :]) + self._padding_by_window[window] = result + return result @classmethod def from_decode( @@ -77,10 +139,15 @@ class BatchedDecodeContext: req_to_token_pool: Any | None, attention_layer_indices: list[int] | None = None, attention_pool_index_by_layer: dict[int, int] | None = None, + full_kv_pool_index_by_layer: dict[int, int] | None = None, ) -> BatchedDecodeContext: batch_size = len(req_ids) if attention_layer_indices is None: attention_layer_indices = list(range(len(caches[0]))) + # One arbitrary attention layer speaks for the whole step: every + # attention cache's ``offset`` is the ABSOLUTE sequence position, so + # they all agree even though a windowed cache stores far fewer + # tokens than that (see the class docstring's read-through rule). seq_lens = [ caches[i][attention_layer_indices[0]].offset for i in range(batch_size) ] @@ -93,6 +160,7 @@ class BatchedDecodeContext: seq_lens=seq_lens, attention_layer_caches=attention_layer_caches, attention_pool_index_by_layer=attention_pool_index_by_layer or {}, + full_kv_pool_index_by_layer=full_kv_pool_index_by_layer or {}, aot=MlxAOTKernelContext.from_decode( aot_kernels=aot_kernels, kv_pool=kv_pool, @@ -122,9 +190,11 @@ class MLXAttentionWrapper(nn.Module): When ``BatchedDecodeContext`` is set, performs per-request RoPE, cache writes, and batched SDPA. Otherwise delegates to inner module. - ``window_size`` marks a sliding-window layer: the pool keeps the full - KV history and the wrapper attends to the trailing window only, which - is numerically identical to a rotating cache. + ``window_size`` marks a sliding-window layer: the wrapper attends to + the trailing window of the cached keys only, which is numerically + identical to a rotating cache. Both cache kinds keep KV in temporal + order and report absolute offsets, so the same trailing-window slice + works whether the cache holds full history or only the window. """ def __init__( @@ -142,8 +212,27 @@ class MLXAttentionWrapper(nn.Module): raise RuntimeError( f"Cannot determine attention scale for {type(inner).__name__}" ) + n_heads = get_num_heads(inner) + n_kv_heads = get_num_kv_heads(inner) + if n_heads is None or n_kv_heads is None: + raise RuntimeError( + f"Cannot determine attention head counts for {type(inner).__name__}" + ) object.__setattr__(self, "_scale", scale) - object.__setattr__(self, "_sinks", getattr(inner, "sinks", None)) + object.__setattr__(self, "_n_heads", n_heads) + object.__setattr__(self, "_n_kv_heads", n_kv_heads) + # None for modules that expose head_dim only through a projection + # shape; _batched_decode falls back to the runtime K shape. + object.__setattr__(self, "_head_dim", get_head_dim(inner)) + object.__setattr__(self, "_has_q_norm", hasattr(inner, "q_norm")) + object.__setattr__(self, "_has_k_norm", hasattr(inner, "k_norm")) + # Only pass sinks when the module has them: the kwarg requires a + # recent mlx and must not constrain models without sinks. + sinks = getattr(inner, "sinks", None) + object.__setattr__(self, "_sinks", sinks) + object.__setattr__( + self, "_sink_kwargs", {} if sinks is None else {"sinks": sinks} + ) def __call__(self, x: mx.array, mask: Any = None, cache: Any = None) -> mx.array: ctx = get_context() @@ -155,18 +244,14 @@ class MLXAttentionWrapper(nn.Module): inner = self._inner layer_idx = self._layer_idx B = ctx.batch_size - n_heads = get_num_heads(inner) - n_kv_heads = get_num_kv_heads(inner) - if n_heads is None or n_kv_heads is None: - raise RuntimeError( - f"Cannot determine attention head counts for {type(inner).__name__}" - ) + n_heads = self._n_heads + n_kv_heads = self._n_kv_heads q_proj_output = inner.q_proj(x) keys = inner.k_proj(x) values = inner.v_proj(x) - head_dim = get_head_dim(inner) + head_dim = self._head_dim if head_dim is None: head_dim = keys.shape[-1] // n_kv_heads @@ -188,9 +273,9 @@ class MLXAttentionWrapper(nn.Module): keys = keys.reshape(B, 1, n_kv_heads, head_dim) values = values.reshape(B, 1, n_kv_heads, head_dim) - if hasattr(inner, "q_norm"): + if self._has_q_norm: queries = inner.q_norm(queries) - if hasattr(inner, "k_norm"): + if self._has_k_norm: keys = inner.k_norm(keys) queries = queries.transpose(0, 2, 1, 3) @@ -199,36 +284,32 @@ class MLXAttentionWrapper(nn.Module): # Vectorized RoPE with per-batch offsets (cached on the context). offsets = ctx.offsets - attention_pool_idx = ctx.attention_pool_index_by_layer[layer_idx] + cache_idx = ctx.attention_pool_index_by_layer[layer_idx] + window = self._window_size - if ctx.aot.rope is not None: - # AOT path: real .metallib RoPE + fused KV pool scatter. + if ctx.aot.rope is not None and window is None: + # AOT path: real .metallib RoPE + fused scatter into this layer's + # pool buffer. queries, keys = self._rope_custom_aot( queries, keys, values, offsets, - attention_pool_idx, + ctx.full_kv_pool_index_by_layer[layer_idx], ctx.aot.rope, ) else: - # Fallback: MLX's built-in mx.fast.rope (used when the AOT kernel - # isn't built or the model uses an unsupported RoPE variant). + # Fallback: MLX's built-in mx.fast.rope. Used when the AOT kernel + # isn't built, the model uses an unsupported RoPE variant, or the + # layer is sliding-window (windowed KV never enters the pool). queries = inner.rope(queries, offset=offsets) keys = inner.rope(keys, offset=offsets) - layer_caches = ctx.attention_layer_caches[attention_pool_idx] - window = self._window_size - if window is None: - pad_sizes = ctx.pad_sizes - else: - # Sliding-window layer: the cache keeps the full history but the - # newest token only attends to the trailing ``window`` keys. The - # padding metadata shared on the context is full-length, so it is - # rebuilt locally for the windowed lengths. - eff_lens = [min(n + 1, window) for n in ctx.seq_lens] - max_eff = max(eff_lens) - pad_sizes = [max_eff - n for n in eff_lens] + layer_caches = ctx.attention_layer_caches[cache_idx] + # A sliding-window layer reads only the trailing ``window`` keys, so its + # padded width differs from the unwindowed one. Both are memoised on + # the context and shared by every layer of their kind. + pad_sizes, attn_mask = ctx.decode_padding(window) # TODO: replace per-request loop with native batched/ragged # attention once mx.fast.scaled_dot_product_attention supports @@ -239,10 +320,7 @@ class MLXAttentionWrapper(nn.Module): for i in range(B): layer_caches[i].write_token(keys[i : i + 1], values[i : i + 1]) - k_all, v_all = layer_caches[i].get_kv() - if window is not None and k_all.shape[2] > window: - k_all = k_all[:, :, -window:, :] - v_all = v_all[:, :, -window:, :] + k_all, v_all = layer_caches[i].get_kv(window) pad = pad_sizes[i] if pad > 0: @@ -257,34 +335,13 @@ class MLXAttentionWrapper(nn.Module): keys_b = mx.concatenate(all_k, axis=0) values_b = mx.concatenate(all_v, axis=0) - pad_mask = None - if window is None: - if ctx.needs_padding: - pad_mask = ctx.positions[None, :] >= ctx.valid_lens[:, None] - elif max(pad_sizes) > 0: - eff = mx.array(eff_lens, dtype=mx.int32) - pad_mask = mx.arange(max_eff)[None, :] >= eff[:, None] - - attn_mask = None - if pad_mask is not None: - attn_mask = mx.where( - pad_mask[:, None, None, :], - mx.array(mx.finfo(queries.dtype).min, dtype=queries.dtype), - mx.array(0.0, dtype=queries.dtype), - ) - - # Only pass sinks when the module has them: the kwarg requires a - # recent mlx and must not constrain models without sinks. - sink_kwargs = {} - if self._sinks is not None: - sink_kwargs["sinks"] = self._sinks output = mx.fast.scaled_dot_product_attention( queries, keys_b, values_b, scale=self._scale, mask=attn_mask, - **sink_kwargs, + **self._sink_kwargs, ) output = output.transpose(0, 2, 1, 3).reshape(B, 1, -1) @@ -298,14 +355,14 @@ class MLXAttentionWrapper(nn.Module): keys: mx.array, values: mx.array, positions: mx.array, - attention_pool_idx: int, + full_pool_idx: int, rope_ctx: MlxAOTRoPEContext, ) -> tuple[mx.array, mx.array]: """AOT path: rotate Q/K and scatter K/V into the shared pool. The kernel call does RoPE on Q/K and scatters - rotated K + (untouched) V into ``kv_pool`` at ``new_token_slots`` - for ``layer_idx``. + rotated K + (untouched) V into ``kv_pool`` buffer ``full_pool_idx`` + at ``new_token_slots``. If ``new_token_slots`` is None, slot=-1 sentinel is used (no pool write, RoPE-only mode). Returns rotated (queries, keys) in the @@ -322,8 +379,8 @@ class MLXAttentionWrapper(nn.Module): else: slots = rope_ctx.new_token_slots.astype(mx.int32) - k_pool = rope_ctx.kv_pool.k_buffer[attention_pool_idx] - v_pool = rope_ctx.kv_pool.v_buffer[attention_pool_idx] + k_pool = rope_ctx.kv_pool.k_buffer[full_pool_idx] + v_pool = rope_ctx.kv_pool.v_buffer[full_pool_idx] q_rot, k_rot, k_pool_new, v_pool_new = rope_ctx.kernel.rope_pool_fused( q_flat, @@ -339,8 +396,8 @@ class MLXAttentionWrapper(nn.Module): rope_base=rope_ctx.kernel.base, ) # Rebind pool buffers (zero-copy donation result). - rope_ctx.kv_pool.k_buffer[attention_pool_idx] = k_pool_new - rope_ctx.kv_pool.v_buffer[attention_pool_idx] = v_pool_new + rope_ctx.kv_pool.k_buffer[full_pool_idx] = k_pool_new + rope_ctx.kv_pool.v_buffer[full_pool_idx] = v_pool_new # (B, n_heads, head_dim) -> (B, n_heads, 1, head_dim) for SDPA path return q_rot[:, :, None, :], k_rot[:, :, None, :] diff --git a/python/sglang/srt/hardware_backend/mlx/kv_cache/layout.py b/python/sglang/srt/hardware_backend/mlx/kv_cache/layout.py index 5483a4f48..d354db9d5 100644 --- a/python/sglang/srt/hardware_backend/mlx/kv_cache/layout.py +++ b/python/sglang/srt/hardware_backend/mlx/kv_cache/layout.py @@ -2,7 +2,7 @@ from __future__ import annotations -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import Any, Sequence @@ -10,9 +10,13 @@ from typing import Any, Sequence class MlxModelCacheLayout: """Map model layers to MLX cache storage components. - Attention layers store softmax-attention KV in the MLX attention KV pool. - Auxiliary layers keep native ``mlx-lm`` cache state and are snapshotted by - the MLX auxiliary-state component. + Full-attention layers store softmax-attention KV in the shared MLX + attention KV pool; sliding-window layers keep per-request windowed KV + only. Auxiliary layers keep native ``mlx-lm`` cache state. + + ``attention_pool_index_by_layer`` indexes per-request cache arrays + (dense over all attention layers); ``full_kv_pool_index_by_layer`` + indexes the shared pool's buffers (dense over full-attention layers). """ layers: tuple[Any, ...] @@ -20,12 +24,37 @@ class MlxModelCacheLayout: attention_layer_indices: tuple[int, ...] auxiliary_layer_indices: tuple[int, ...] attention_pool_index_by_layer: dict[int, int] + # Per-layer sliding window (None or absent = full attention). + layer_window_sizes: dict[int, int | None] = field(default_factory=dict) + full_attention_layer_indices: tuple[int, ...] = field(init=False) + swa_attention_layer_indices: tuple[int, ...] = field(init=False) + full_kv_pool_index_by_layer: dict[int, int] = field(init=False) + + def __post_init__(self) -> None: + full_indices = tuple( + idx + for idx in self.attention_layer_indices + if self.layer_window_sizes.get(idx) is None + ) + swa_indices = tuple( + idx + for idx in self.attention_layer_indices + if self.layer_window_sizes.get(idx) is not None + ) + object.__setattr__(self, "full_attention_layer_indices", full_indices) + object.__setattr__(self, "swa_attention_layer_indices", swa_indices) + object.__setattr__( + self, + "full_kv_pool_index_by_layer", + {layer_idx: pool_idx for pool_idx, layer_idx in enumerate(full_indices)}, + ) @classmethod def from_attention_discovery( cls, layers: Sequence[Any], attention_attrs: Sequence[str | None], + layer_window_sizes: dict[int, int | None] | None = None, ) -> MlxModelCacheLayout: if len(layers) != len(attention_attrs): raise ValueError( @@ -50,6 +79,7 @@ class MlxModelCacheLayout: attention_layer_indices=attention_layer_indices, auxiliary_layer_indices=auxiliary_layer_indices, attention_pool_index_by_layer=attention_pool_index_by_layer, + layer_window_sizes=dict(layer_window_sizes or {}), ) @property @@ -60,6 +90,14 @@ class MlxModelCacheLayout: def num_attention_layers(self) -> int: return len(self.attention_layer_indices) + @property + def num_full_attention_layers(self) -> int: + return len(self.full_attention_layer_indices) + + @property + def has_sliding_window_layers(self) -> bool: + return bool(self.swa_attention_layer_indices) + @property def has_auxiliary_state(self) -> bool: return bool(self.auxiliary_layer_indices) @@ -70,12 +108,22 @@ class MlxModelCacheLayout: raise RuntimeError("MLX model has no supported attention layers") return self.attention_layer_indices[0] + def window_size(self, layer_idx: int) -> int | None: + """Sliding window of *layer_idx*, or None for full attention.""" + return self.layer_window_sizes.get(layer_idx) + def attention_pool_index(self, layer_idx: int) -> int: try: return self.attention_pool_index_by_layer[layer_idx] except KeyError as exc: raise KeyError(f"Layer {layer_idx} is not an attention layer") from exc + def full_kv_pool_index(self, layer_idx: int) -> int: + try: + return self.full_kv_pool_index_by_layer[layer_idx] + except KeyError as exc: + raise KeyError(f"Layer {layer_idx} is not a full-attention layer") from exc + def attention_attr(self, layer_idx: int) -> str: attr = self.attention_attrs[layer_idx] if attr is None: diff --git a/python/sglang/srt/hardware_backend/mlx/model_runner.py b/python/sglang/srt/hardware_backend/mlx/model_runner.py index 57e396afa..d2bb58382 100644 --- a/python/sglang/srt/hardware_backend/mlx/model_runner.py +++ b/python/sglang/srt/hardware_backend/mlx/model_runner.py @@ -5,8 +5,15 @@ scheduler (``TokenToKVPoolAllocator`` / ``RadixCache``). This runner reads cached attention KV from ``MlxAttentionKVPool``, restores any native auxiliary layer state, runs the forward pass, and writes the new cache state back. Each request keeps model-shaped cache entries: -attention layers use ``ContiguousAttentionKVCache`` and auxiliary layers -use native ``mlx-lm`` cache objects. +full-attention layers use ``ContiguousAttentionKVCache``, sliding-window +layers use a fixed-size ``WindowedAttentionKVCache`` on both KV paths, +and auxiliary layers use native ``mlx-lm`` cache objects. + +The shared pool stores full-attention layers only, so no cross-request +SWA prefix KV exists: a radix prefix hit on a sliding-window model +recomputes the whole prefix (a trailing-band rebuild is inexact because +window receptive fields chain backwards through layers). The +scheduler's slot bookkeeping is untouched. The module also exposes a lazy-eval (`*_start` / `*_finalize`) surface used by the MLX overlap scheduler to pipeline CPU bookkeeping with @@ -22,6 +29,7 @@ from dataclasses import dataclass from typing import Any import mlx.core as mx +import numpy as np import psutil from mlx.utils import tree_flatten from mlx_lm import load as mlx_lm_load @@ -40,14 +48,29 @@ from sglang.srt.hardware_backend.mlx.kv_cache import ( MLXAttentionWrapper, MlxModelCacheLayout, PoolBackedAttentionKVCache, + WindowedAttentionKVCache, clear_context, find_attention_layers, get_head_dim, + get_layer_window_sizes, get_num_kv_heads, patch_model_attention, set_context, uses_sliding_window_attention, ) +from sglang.srt.hardware_backend.mlx.sampling import ( + GREEDY_PARAMS, + MlxLazyLogprobs, + MlxLogprobSpec, + MlxSamplingParams, + MlxStepLogprobs, + all_greedy, + compute_logprobs, + lazy_logprob_arrays, + sample_tokens, + sanitize_logits, + scale_by_temperature, +) from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.runtime_context import ( mamba_cache_chunk_size, @@ -72,6 +95,7 @@ class MlxPendingPrefill: full_token_ids: list[int] req_pool_idx: int synced_offset: int + lazy_logprobs: MlxLazyLogprobs | None = None @dataclass @@ -85,9 +109,11 @@ class MlxPendingExtend: """ lazy_token: mx.array + cache: list[Any] req_id: str new_token_ids: list[int] new_synced_offset: int + lazy_logprobs: MlxLazyLogprobs | None = None @dataclass @@ -104,6 +130,13 @@ class MlxPendingDecode: lazy_tokens: mx.array req_ids: list[str] caches: list[list[Any]] + lazy_logprobs: MlxLazyLogprobs | None = None + # Carried so a chained step recomputes the same logprob request. + logprob_spec: MlxLogprobSpec | None = None + # Carried so chained steps keep applying static logit_bias rows. + # Never holds a grammar mask on the chained path: grammar batches are + # not chain_safe, so their pendings never become a chain root. + edit_rows: mx.array | None = None _MLX_QUANTIZATION_PRESETS: dict[str, tuple[int, int]] = { @@ -117,6 +150,14 @@ _MLX_KV_FLOAT_DTYPES = {mx.float16, mx.bfloat16, mx.float32} class MlxModelRunner: """MLX model runner with radix-cache prefix sharing.""" + # Class defaults cover unit tests that build runners via object.__new__ + # without running __init__/_load_model, which set the real values. + # ``_trunk`` is the headless trunk, resolved in _load_model. + _trunk = None + _enable_sampling = False + _sanitize_nan = False + _deterministic_seeding = False + def __init__( self, model_path: str, @@ -125,12 +166,27 @@ class MlxModelRunner: pool_size: int | None = None, mem_fraction_static: float = 0.8, quantization: str | None = None, + enable_sampling: bool = False, + sampling_rng_seed: int = 0, + deterministic_seeding: bool = False, ): self.model_path = model_path self.trust_remote_code = trust_remote_code self.model = None self.disable_radix_cache = disable_radix_cache self._mem_fraction_static = mem_fraction_static + self._enable_sampling = enable_sampling + # --enable-deterministic-inference parity: seed every row (the + # sampling module's DEFAULT_SAMPLING_SEED when the request has no + # sampling_seed), like the pytorch backend. + self._deterministic_seeding = deterministic_seeding + self._sanitize_nan = envs.SGLANG_SANITIZE_NAN_LOGITS.get() + # RNG state for unseeded sampled rows; split at every sampling + # graph build, so runs are reproducible given the same seed and + # request schedule. + self._rng_key: mx.array | None = ( + mx.random.key(sampling_rng_seed) if enable_sampling else None + ) # Counter used to trigger periodic mx.clear_cache() calls. self._decode_step_ct: int = 0 self._clear_steps = envs.SGLANG_MLX_CLEAR_CACHE_STEPS.get() @@ -155,6 +211,8 @@ class MlxModelRunner: self._cache_layout = MlxModelCacheLayout.from_attention_discovery( layer_list, attn_attrs, + # Per-layer sliding windows (container convention, e.g. gpt-oss). + layer_window_sizes=get_layer_window_sizes(self.model), ) if self._cache_layout.num_attention_layers == 0: raise RuntimeError("MLX model has no supported attention layers") @@ -164,6 +222,16 @@ class MlxModelRunner: raise RuntimeError( "MLX models with auxiliary cache state require model.make_cache()." ) + if ( + self._cache_layout.has_auxiliary_state + and self._cache_layout.has_sliding_window_layers + ): + # Auxiliary-state restore assumes a prefix hit runs only the new + # tokens; an SWA prefix hit recomputes the prefix on a fresh cache. + raise NotImplementedError( + "MLX runner does not support models with both auxiliary " + "cache state and sliding-window attention layers." + ) if self._cache_layout.has_auxiliary_state: self._model_embed, self._model_norm, self._model_lm_head = ( self._extract_model_components() @@ -172,7 +240,9 @@ class MlxModelRunner: self._req_caches: dict[str, list[Any]] = {} self._req_token_ids: dict[str, list[int]] = {} - self._cache_pool: list[list[Any]] = [] # reusable full-attention caches + self._req_sampling: dict[str, MlxSamplingParams] = {} + # Reusable cache lists, for models without auxiliary layer state. + self._cache_pool: list[list[Any]] = [] self._attention_kv_pool: MlxAttentionKVPool | None = None self._req_to_token_pool: ReqToTokenPool | None = None @@ -206,7 +276,12 @@ class MlxModelRunner: """Create a model-shaped cache list with attention KV adapters.""" cache = self._new_cache_skeleton() for layer_idx in self._cache_layout.attention_layer_indices: - cache[layer_idx] = ContiguousAttentionKVCache(max_seq_len=self._max_seq_len) + window = self._cache_layout.window_size(layer_idx) + cache[layer_idx] = ( + WindowedAttentionKVCache(window) + if window is not None + else ContiguousAttentionKVCache(max_seq_len=self._max_seq_len) + ) return cache def _acquire_cache(self) -> list[Any]: @@ -214,7 +289,7 @@ class MlxModelRunner: if not self._cache_layout.has_auxiliary_state and self._cache_pool: cache = self._cache_pool.pop() for c in cache: - c.offset = 0 + c.reset() return cache return self._new_native_cache() @@ -335,13 +410,18 @@ class MlxModelRunner: def _cache_with_pool_backed_attention( self, prefix_slot_ids: list[int], prefix_len: int ) -> list[Any]: + """Build a prefill cache list gathering *prefix_len* pool tokens. + + Only reachable without sliding-window layers (SWA prefix hits + recompute instead), so every attention layer here is pool-backed. + """ assert self._attention_kv_pool is not None slot_ids_mx = mx.array(prefix_slot_ids, dtype=mx.int32) cache = self._new_cache_skeleton() for layer_idx in self._cache_layout.attention_layer_indices: cache[layer_idx] = PoolBackedAttentionKVCache( self._attention_kv_pool, - self._cache_layout.attention_pool_index(layer_idx), + self._cache_layout.full_kv_pool_index(layer_idx), slot_ids_mx, prefix_len, ) @@ -379,22 +459,15 @@ class MlxModelRunner: return arrays @staticmethod - def _eval_with_cache(token_result: mx.array, cache: list[Any]) -> None: - """Evaluate token result and all cache buffers in one mx.eval call.""" - mx.eval( - token_result, - *[s for c in cache for s in MlxModelRunner._cache_arrays(c)], - ) + def cache_state_arrays(caches: list[list[Any]]) -> list[mx.array]: + """Flatten per-request cache lists (``caches[req][layer]``) to arrays. - @staticmethod - def _cache_state_arrays(pending_caches: list[list[Any]]) -> list[mx.array]: - """Flatten pending decode cache state list into an array list. - - Safe to hand to ``mx.async_eval``. + Pass ``[cache]`` for a single request. Safe to hand to + ``mx.eval`` / ``mx.async_eval``. """ return [ s - for cache_list in pending_caches + for cache_list in caches for cache in cache_list for s in MlxModelRunner._cache_arrays(cache) ] @@ -465,6 +538,17 @@ class MlxModelRunner: load_time = time.time() - start_time logger.info(f"MLX model loaded in {load_time:.2f}s") + # mlx-lm models expose the headless trunk as ``Model.model``; without + # it, non-final chunked-prefill chunks cannot skip the logit head. + trunk = getattr(self.model, "model", None) + self._trunk = trunk if callable(trunk) else None + if self._trunk is None: + logger.info( + "Model %s exposes no headless trunk (`.model`); non-final " + "chunked-prefill chunks will compute full vocab logits.", + type(self.model).__name__, + ) + # Optional: Path B fusion — keep up_proj/gate_proj weights separate # (no matmul-kernel tile regression) but fuse the swiglu activation # into the gate matmul via a custom Metal kernel. Activated by @@ -494,11 +578,13 @@ class MlxModelRunner: ) -> tuple[int, int, mx.Dtype]: layer = self._cache_layout.layers[layer_idx] sample_attn = self._attention_module_for_layer(layer_idx) - if uses_sliding_window_attention(layer, sample_attn): + unsized_window = self._cache_layout.window_size(layer_idx) is None + if unsized_window and uses_sliding_window_attention(layer, sample_attn): raise NotImplementedError( - "MLX radix attention KV pool does not support sliding-window " - f"attention yet at layer {layer_idx}. Sliding-window KV needs " - "per-layer/window-aware pools." + f"Attention layer {layer_idx} declares sliding-window " + "attention but the model exposes no per-layer window map " + "(container `layer_types` plus a scalar window), so the MLX " + "KV cache cannot bound its sliding-window KV." ) n_kv_heads = get_num_kv_heads(sample_attn) if n_kv_heads is None: @@ -528,7 +614,11 @@ class MlxModelRunner: return n_kv_heads, head_dim, dtype def _get_attn_config(self) -> tuple[int, int, mx.Dtype]: - """Return the uniform attention KV config used by the shared MLX pool.""" + """Return the uniform KV config shared by every attention layer. + + Sizes the shared pool and the AOT kernels; sliding-window layers + must match the same shape because they share the decode kernels. + """ if self._cache_layout.num_attention_layers == 0: raise RuntimeError( "Cannot determine attention config: no attention module found" @@ -539,12 +629,11 @@ class MlxModelRunner: config = self._attention_kv_config_for_layer(layer_idx) if config != first_config: raise NotImplementedError( - "MLX radix attention KV pool requires uniform softmax-attention " + "MLX attention KV caching requires uniform softmax-attention " "KV shape across layers. " f"Layer {first_layer_idx} has {first_config}, " f"but layer {layer_idx} has {config}. " - "Heterogeneous attention KV or sliding-window KV needs " - "per-layer pools." + "Heterogeneous attention KV needs per-layer pools." ) return first_config @@ -553,7 +642,13 @@ class MlxModelRunner: if explicit_size is not None: return explicit_size n_kv_heads, head_dim, dtype = self._get_attn_config() - num_layers = self._cache_layout.num_attention_layers + # Only full-attention layers occupy pool slots. All-SWA models have no + # pool at all and fall back to the all-layer formula purely to keep the + # scheduler's token budget finite. + num_layers = ( + self._cache_layout.num_full_attention_layers + or self._cache_layout.num_attention_layers + ) sys_available = psutil.virtual_memory().available mlx_limit = mx.device_info().get( "max_recommended_working_set_size", @@ -598,15 +693,44 @@ class MlxModelRunner: ) def init_cache_pools(self, req_to_token_pool: ReqToTokenPool | None) -> None: - """Create attention KV pool (+1 for padding slot 0).""" + """Create the full-attention KV pool (+1 for padding slot 0).""" self._req_to_token_pool = req_to_token_pool if self.disable_radix_cache: return + num_pool_layers = self._cache_layout.num_full_attention_layers + if self._cache_layout.has_sliding_window_layers: + # The pool exists to serve radix prefix hits, and an SWA prefix hit + # recomputes the prefix instead of gathering it (see prefill_start), + # so on any SWA model the pool has no reader: its sole consumer is + # PoolBackedAttentionKVCache, reachable only when + # trusted_prefix_len > 0, which requires no SWA layers. Allocating + # it anyway would burn the whole auto-sized KV budget + # (_compute_pool_size fills mem_fraction_static) on a buffer that is + # only ever written. Skipping it also disables the fused AOT + # RoPE + pool-scatter kernel for the full layers, whose scatter half + # is dead work here; that kernel is opt-in + # (SGLANG_MLX_USE_CUSTOM_ROPE, default off), so the default path + # loses nothing. + # + # Un-gate this together with the window-aware shared SWA pool that + # restores fast prefix hits: the seams it needs (the layout + # partition, the full-pool index, layer-type dispatch) are already + # in place. + logger.info( + "Model has %d sliding-window attention layers; skipping the " + "shared attention KV pool (an SWA prefix hit recomputes the " + "prefix, so the pool would never be read). Per-request " + "windowed caches only.", + len(self._cache_layout.swa_attention_layer_indices), + ) + return + if num_pool_layers == 0: + return n_kv_heads, head_dim, dtype = self._get_attn_config() # +1 for padding slot 0 self._attention_kv_pool = MlxAttentionKVPool( pool_size=self._pool_size + 1, - num_layers=self._cache_layout.num_attention_layers, + num_layers=num_pool_layers, n_kv_heads=n_kv_heads, head_dim=head_dim, dtype=dtype, @@ -614,7 +738,9 @@ class MlxModelRunner: logger.info( f"Attention KV pool initialized: pool_size={self._pool_size} " f"(buffer size {self._pool_size + 1} incl. padding slot 0), " - f"{self._cache_layout.num_attention_layers} attention layers, " + f"{num_pool_layers} full-attention layers " + f"({len(self._cache_layout.swa_attention_layer_indices)} " + "sliding-window layers stay per-request), " f"{n_kv_heads} kv_heads, {head_dim} head_dim" ) @@ -627,8 +753,14 @@ class MlxModelRunner: new_slot_ids: list[int], req_pool_idx: int, req: Any | None = None, + needs_logits: bool = True, ) -> int: - """Prefill a request. Returns next_token_id.""" + """Prefill a request. Returns next_token_id. + + One-shot convenience wrapper around ``prefill_start`` / + ``prefill_finalize``; logit edits and logprobs are only available + through that lazy surface. + """ pending = self.prefill_start( req_id=req_id, new_token_ids=new_token_ids, @@ -637,8 +769,9 @@ class MlxModelRunner: new_slot_ids=new_slot_ids, req_pool_idx=req_pool_idx, req=req, + needs_logits=needs_logits, ) - self._eval_with_cache(pending.lazy_token, pending.cache) + self.eval_pending(pending) return self.prefill_finalize(pending) def extend( @@ -646,10 +779,16 @@ class MlxModelRunner: req_id: str, new_token_ids: list[int], new_slot_ids: list[int], + needs_logits: bool = True, ) -> int: - """Continue prefill for a chunked request. Returns next_token_id.""" - pending = self.extend_start(req_id, new_token_ids, new_slot_ids) - self._eval_with_cache(pending.lazy_token, self._req_caches[req_id]) + """Continue prefill for a chunked request. Returns next_token_id. + + One-shot convenience wrapper; see :meth:`prefill`. + """ + pending = self.extend_start( + req_id, new_token_ids, new_slot_ids, needs_logits=needs_logits + ) + self.eval_pending(pending) return self.extend_finalize(pending) def _sync_new_kv_to_pool( @@ -658,9 +797,15 @@ class MlxModelRunner: cache_start: int, slot_ids: list[int], ) -> None: - """Sync attention KV from contiguous cache to pool at the given slots.""" + """Sync full-attention KV from contiguous caches to the pool slots. + + Sliding-window layers are skipped: they keep no pool KV, and their + buffers are window-local so the absolute slicing below would not + apply to them anyway. + """ if not slot_ids or self._attention_kv_pool is None: return + full_layer_indices = self._cache_layout.full_attention_layer_indices end = cache_start + len(slot_ids) slot_ids_mx = mx.array(slot_ids, dtype=mx.int32) # TODO: Standardize ContiguousAttentionKVCache size to avoid transpose @@ -668,13 +813,13 @@ class MlxModelRunner: k_all = mx.stack( [ cache[layer_idx].keys[0, :, cache_start:end, :].transpose(1, 0, 2) - for layer_idx in self._cache_layout.attention_layer_indices + for layer_idx in full_layer_indices ] ) v_all = mx.stack( [ cache[layer_idx].values[0, :, cache_start:end, :].transpose(1, 0, 2) - for layer_idx in self._cache_layout.attention_layer_indices + for layer_idx in full_layer_indices ] ) self._attention_kv_pool.set_kv_all_layers(slot_ids_mx, k_all, v_all) @@ -711,17 +856,13 @@ class MlxModelRunner: for req_id in list(self._req_caches.keys()): self._sync_decode_kv_to_pool(req_id) - def decode_batch( - self, - req_ids: list[str], - ) -> list[int]: - """Decode one token per request.""" + def decode_batch(self, req_ids: list[str]) -> list[int]: + """Decode one token per request. + + One-shot convenience wrapper; see :meth:`prefill`. + """ pending = self.decode_batch_start(req_ids) - # Evaluate lazy_tokens together with every affected cache buffer so - # the attention write-then-read ordering is materialised in one - # kernel submission. - cache_arrays = self._cache_state_arrays(pending.caches) - mx.eval(pending.lazy_tokens, *cache_arrays) + self.eval_pending(pending) return self.decode_batch_finalize(pending) def prefill_start( @@ -733,6 +874,9 @@ class MlxModelRunner: new_slot_ids: list[int], req_pool_idx: int, req: Any | None = None, + needs_logits: bool = True, + logit_edit_row: mx.array | None = None, + logprob_spec: MlxLogprobSpec | None = None, ) -> MlxPendingPrefill: """Queue a prefill forward pass without evaluating. @@ -740,17 +884,28 @@ class MlxModelRunner: next-token ``mx.array`` plus everything needed to commit the request in :meth:`prefill_finalize`. The caller drives the GPU by handing ``lazy_token`` (and cache state) to ``mx.async_eval``. + + ``needs_logits=False`` marks the first chunk of a chunked prompt + (its next-token output is discarded); see :meth:`extend_start`. """ prefix_len = len(prefix_slot_ids) if req is not None: req.mamba_last_track_seqlen = None + if self._enable_sampling: + self._req_sampling[req_id] = ( + MlxSamplingParams.from_req( + req, deterministic_seeding=self._deterministic_seeding + ) + if req is not None + else GREEDY_PARAMS + ) if self.disable_radix_cache: cache = self._acquire_cache() input_ids = mx.array([new_token_ids], dtype=mx.int32) - model_output = self.model(input_ids, cache=cache) - logits = self._extract_logits(model_output) - lazy_token = mx.argmax(logits[:, -1, :], axis=-1) + lazy_token, lazy_logprobs = self._forward_lazy_token( + input_ids, cache, needs_logits, req_id, logit_edit_row, logprob_spec + ) return MlxPendingPrefill( lazy_token=lazy_token, cache=cache, @@ -758,9 +913,17 @@ class MlxModelRunner: full_token_ids=list(full_token_ids), req_pool_idx=req_pool_idx, synced_offset=0, + lazy_logprobs=lazy_logprobs, ) - assert self._attention_kv_pool is not None + # A pool is required only where one can actually be read: a model with + # full-attention layers and no sliding-window layers. init_cache_pools + # skips it otherwise, and the gather path below is unreachable then. + assert ( + self._attention_kv_pool is not None + or self._cache_layout.num_full_attention_layers == 0 + or self._cache_layout.has_sliding_window_layers + ) new_token_count = len(new_token_ids) track_len = self._select_auxiliary_state_track_len( @@ -770,8 +933,22 @@ class MlxModelRunner: req=req, ) - if prefix_len > 0: - cache = self._cache_with_pool_backed_attention(prefix_slot_ids, prefix_len) + # Sliding-window layers keep no pool KV, so a prefix hit has nothing to + # gather and re-runs the prefix. A trailing-band rebuild would not be + # exact: each rebuilt position needs its own window of exact hidden + # states, and that dependency chains back through every layer. Only the + # run is clamped -- slot ids and synced offsets stay unclamped. + if prefix_len > 0 and self._cache_layout.has_sliding_window_layers: + trusted_prefix_len = 0 + run_token_ids = list(full_token_ids[:prefix_len]) + new_token_ids + else: + trusted_prefix_len = prefix_len + run_token_ids = new_token_ids + + if trusted_prefix_len > 0: + cache = self._cache_with_pool_backed_attention( + prefix_slot_ids[:trusted_prefix_len], trusted_prefix_len + ) pool_backed_attention = True restored_auxiliary_state = ( not self._cache_layout.has_auxiliary_state @@ -787,9 +964,9 @@ class MlxModelRunner: # allocated attention KV below. cache = self._acquire_cache() input_ids = mx.array([full_token_ids or new_token_ids], dtype=mx.int32) - model_output = self.model(input_ids, cache=cache) - logits = self._extract_logits(model_output) - lazy_token = mx.argmax(logits[:, -1, :], axis=-1) + lazy_token, lazy_logprobs = self._forward_lazy_token( + input_ids, cache, needs_logits, req_id, logit_edit_row, logprob_spec + ) if new_slot_ids: self._sync_new_kv_to_pool(cache, prefix_len, new_slot_ids) return MlxPendingPrefill( @@ -799,39 +976,45 @@ class MlxModelRunner: full_token_ids=list(full_token_ids), req_pool_idx=req_pool_idx, synced_offset=prefix_len + len(new_slot_ids), + lazy_logprobs=lazy_logprobs, ) else: cache = self._acquire_cache() pool_backed_attention = False - if new_token_count > 0: + if run_token_ids: track_new_count = track_len - prefix_len if track_len is not None else None if track_new_count is not None and 0 < track_new_count < new_token_count: + # aux + SWA is rejected at init, so run_token_ids is + # new_token_ids on this branch. input_ids = mx.array([new_token_ids[:track_new_count]], dtype=mx.int32) - self.model(input_ids, cache=cache) + # Cache side effects only — this intermediate forward's + # output is never read, so skip the head when possible. + if self._trunk_forward(input_ids, cache) is None: + self.model(input_ids, cache=cache) self._store_tracked_auxiliary_state(req, cache, track_len) if pool_backed_attention: cache = self._materialize_pool_backed_attention(cache) pool_backed_attention = False extend_tokens = new_token_ids[track_new_count:] else: - extend_tokens = new_token_ids + extend_tokens = run_token_ids else: - # Full cache hit - rerun last token to get next-token logits + # Full cache hit - rerun last token to get next-token logits. + # Unreachable with SWA layers: a prefix rebuild always leaves run + # tokens whose final logits already predict the next token. extend_tokens = full_token_ids[-1:] for c in cache: c.offset = max(c.offset - 1, 0) input_ids = mx.array([extend_tokens], dtype=mx.int32) - model_output = self.model(input_ids, cache=cache) - logits = self._extract_logits(model_output) + lazy_token, lazy_logprobs = self._forward_lazy_token( + input_ids, cache, needs_logits, req_id, logit_edit_row, logprob_spec + ) if track_len is not None and track_len == prefix_len + new_token_count: self._store_tracked_auxiliary_state(req, cache, track_len) - last_logits = logits[:, -1, :] - lazy_token = mx.argmax(last_logits, axis=-1) - # Convert pool-backed attention KV to contiguous attention KV for decode. # This appends a lazy slice-assign onto the forward graph; the # arrays get materialised when the caller evaluates lazy_token. @@ -848,6 +1031,7 @@ class MlxModelRunner: full_token_ids=list(full_token_ids), req_pool_idx=req_pool_idx, synced_offset=prefix_len + len(new_slot_ids), + lazy_logprobs=lazy_logprobs, ) def prefill_finalize(self, pending: MlxPendingPrefill) -> int: @@ -872,8 +1056,16 @@ class MlxModelRunner: req_id: str, new_token_ids: list[int], new_slot_ids: list[int], + needs_logits: bool = True, + logit_edit_row: mx.array | None = None, + logprob_spec: MlxLogprobSpec | None = None, ) -> MlxPendingExtend: - """Queue chunked-prefill continuation without evaluating.""" + """Queue chunked-prefill continuation without evaluating. + + ``needs_logits=False`` marks a non-final chunk whose next-token + output the scheduler discards; the logit head is skipped when the + model exposes a headless trunk. + """ assert ( req_id in self._req_caches ), f"extend_start called for unknown request {req_id}" @@ -881,9 +1073,9 @@ class MlxModelRunner: cache = self._req_caches[req_id] input_ids = mx.array([new_token_ids], dtype=mx.int32) - model_output = self.model(input_ids, cache=cache) - logits = self._extract_logits(model_output) - lazy_token = mx.argmax(logits[:, -1, :], axis=-1) + lazy_token, lazy_logprobs = self._forward_lazy_token( + input_ids, cache, needs_logits, req_id, logit_edit_row, logprob_spec + ) if not self.disable_radix_cache and new_slot_ids: synced = self._req_synced_offset[req_id] @@ -894,9 +1086,11 @@ class MlxModelRunner: return MlxPendingExtend( lazy_token=lazy_token, + cache=cache, req_id=req_id, new_token_ids=list(new_token_ids), new_synced_offset=new_synced_offset, + lazy_logprobs=lazy_logprobs, ) def extend_finalize(self, pending: MlxPendingExtend) -> int: @@ -916,6 +1110,201 @@ class MlxModelRunner: ) return next_token + def _trunk_forward(self, input_ids: mx.array, cache: list[Any]) -> mx.array | None: + """Run the model WITHOUT its logit head, for cache side effects only. + + Non-final chunked-prefill chunks discard their next-token output + (``extend_finalize`` pops it), yet the full model call still computes + vocab-sized float32 logits for every chunk position — for a 200K-vocab + model that is ~100x the useful head work and the largest transient + allocation in the process. ``self._trunk`` is resolved once at load; + returns None when the model exposes no headless trunk (caller falls + back to the full forward). + """ + if self._trunk is None: + return None + return self._trunk(input_ids, cache=cache) + + def _forward_lazy_token( + self, + input_ids: mx.array, + cache: list[Any], + needs_logits: bool, + req_id: str, + logit_edit_row: mx.array | None = None, + logprob_spec: MlxLogprobSpec | None = None, + ) -> tuple[mx.array, MlxLazyLogprobs | None]: + """Forward one chunk, returning (lazy next-token, lazy logprobs). + + Skips the logit head for discarded-output chunks when possible. + """ + if not needs_logits: + hidden = self._trunk_forward(input_ids, cache) + if hidden is not None: + return self._dummy_next_token(hidden), None + # Headless trunk unavailable: run the full model, but keep the + # discarded token on argmax — sampling here would consume RNG + # state and make final output depend on prefill chunking. + model_output = self.model(input_ids, cache=cache) + logits = self._extract_logits(model_output) + return mx.argmax(logits[:, -1, :], axis=-1), None + model_output = self.model(input_ids, cache=cache) + logits = self._extract_logits(model_output) + edits = logit_edit_row[None, :] if logit_edit_row is not None else None + return self._select_tokens_with_logprobs( + logits[:, -1, :], [req_id], [cache], edits, logprob_spec + ) + + def _select_tokens_with_logprobs( + self, + last_logits: mx.array, + req_ids: list[str], + caches: list[list[Any]], + edit_rows: mx.array | None = None, + logprob_spec: MlxLogprobSpec | None = None, + ) -> tuple[mx.array, MlxLazyLogprobs | None]: + """Pick one token per row of ``last_logits`` — lazily, inside the graph. + + Greedy behavior (sampling disabled, or every row greedy with no + logit edits) is exactly the pre-sampling ``mx.argmax`` and consumes + no RNG state. ``edit_rows`` is the worker's pre-combined additive + [B, vocab] array (grammar mask + logit_bias), applied before token + selection and logprobs, mirroring the CUDA + ``ModelRunner._preprocess_logits`` order. Positions for seeded rows + come from the attention cache offsets, which the just-built forward + has already advanced past the token being sampled — the same + ``seq_len - 1`` the pytorch path feeds its sampler. They are + build-time Python ints, so this is chained-decode safe. + """ + if not self._enable_sampling: + return mx.argmax(last_logits, axis=-1), None + params = [self._req_sampling[rid] for rid in req_ids] + edited = self._edited_logits(last_logits, edit_rows) + greedy = all_greedy(params) + # Built once and shared: sampling and logprobs both start from + # logits/temperature, and MLX does not CSE the two identical graphs. + # Stays None when neither needs it (the greedy, no-logprob path). + scaled = ( + scale_by_temperature(edited, params) + if not greedy or logprob_spec is not None + else None + ) + if greedy: + tokens = mx.argmax(edited, axis=-1) + else: + positions = [self._first_attention_cache(c).offset - 1 for c in caches] + self._rng_key, key = mx.random.split(self._rng_key) + tokens = sample_tokens( + last_logits=edited, + params=params, + positions=positions, + key=key, + scaled=scaled, + ) + lazy_logprobs = ( + compute_logprobs( + last_logits=edited, + params=params, + tokens=tokens, + spec=logprob_spec, + scaled=scaled, + ) + if logprob_spec is not None + else None + ) + return tokens, lazy_logprobs + + def _edited_logits( + self, last_logits: mx.array, edit_rows: mx.array | None + ) -> mx.array: + """Apply the additive logit edits and env-gated NaN sanitization.""" + edited = last_logits + if edit_rows is not None: + # The edit rows are sized from SamplingBatchInfo.vocab_size while + # these logits come from the model's lm_head. A model whose head + # is padded past the tokenizer vocabulary would otherwise fail as + # an opaque broadcast error deep in the lazy graph. + if edit_rows.shape[-1] != last_logits.shape[-1]: + raise RuntimeError( + "Logit edit rows do not match the model's vocabulary: " + f"sampling_info.vocab_size={edit_rows.shape[-1]} vs " + f"lm_head width {last_logits.shape[-1]}" + ) + edited = edited.astype(mx.float32) + edit_rows + if self._sanitize_nan: + edited = sanitize_logits(edited.astype(mx.float32)) + return edited + + def _run_logits_hook(self, last_logits: mx.array, logits_hook) -> mx.array: + """Materialize logits and let the worker edit them on the CPU. + + Used for custom logit processors (arbitrary torch callables) — the + one edit that cannot be expressed lazily. Synchronizes the graph; + callers gate this to fresh, pure-decode launches, so the chained + overlap pipeline never pays for it. + """ + logits32 = last_logits.astype(mx.float32) + mx.eval(logits32) + return mx.array(logits_hook(np.array(logits32))) + + def collect_logprobs( + self, lazy_logprobs: MlxLazyLogprobs | None + ) -> MlxStepLogprobs | None: + """Materialize one step's lazy logprob arrays into Python lists.""" + if lazy_logprobs is None: + return None + mx.eval(*lazy_logprob_arrays(lazy_logprobs)) + spec = lazy_logprobs.spec + if lazy_logprobs.top_val is not None: + all_val = lazy_logprobs.top_val.tolist() + all_idx = lazy_logprobs.top_idx.tolist() + top_val = [all_val[i][:k] for i, k in enumerate(spec.top_ks)] + top_idx = [all_idx[i][:k] for i, k in enumerate(spec.top_ks)] + else: + top_val = [[] for _ in spec.top_ks] + top_idx = [[] for _ in spec.top_ks] + return MlxStepLogprobs( + chosen=lazy_logprobs.chosen.tolist(), + top_val=top_val, + top_idx=top_idx, + token_ids_val=[ + a.tolist() if a is not None else [] for a in lazy_logprobs.token_ids_val + ], + token_ids_idx=[list(ids) if ids else [] for ids in spec.token_ids], + ) + + def eval_pending( + self, pending: MlxPendingPrefill | MlxPendingExtend | MlxPendingDecode + ) -> None: + """Materialize a queued forward: token(s), cache writes and logprobs. + + One ``mx.eval`` for the whole pending, so the attention + write-then-read ordering is materialised in a single kernel + submission. Prefill and extend carry one request's per-layer + cache; a decode carries one cache list per request. + """ + if isinstance(pending, MlxPendingDecode): + tokens, caches = pending.lazy_tokens, pending.caches + else: + tokens, caches = pending.lazy_token, [pending.cache] + mx.eval( + tokens, + *self.cache_state_arrays(caches), + *lazy_logprob_arrays(pending.lazy_logprobs), + ) + + @staticmethod + def _dummy_next_token(hidden: mx.array) -> mx.array: + """Graph-connected placeholder token for a skipped-head chunk. + + Value is always 0 (a valid vocab id); it is appended and then popped + as the "stale intermediate token" by the next chunk's finalize. + Deriving it from ``hidden`` keeps the trunk in the lazy graph handed + to ``mx.eval``/``mx.async_eval`` (cache arrays are also evaluated + explicitly by both call paths). + """ + return (hidden[:, -1, 0] * 0).astype(mx.int32) + def _extract_model_components(self): """Cache embedding, norm, and lm_head for layer-by-layer hybrid forward.""" root = getattr(self.model, "language_model", self.model) @@ -936,7 +1325,7 @@ class MlxModelRunner: batched_input: mx.array, req_ids: list[str], ) -> mx.array: - """Layer-by-layer hybrid decode for attention plus auxiliary state. + """Layer-by-layer hybrid decode; returns [B, vocab] last-token logits. Attention layers run with batched hidden states via ``BatchedDecodeContext``. Auxiliary layers run batched when their @@ -971,7 +1360,7 @@ class MlxModelRunner: hidden_states = self._model_norm(hidden_states) logits = self._extract_logits(self._model_lm_head(hidden_states)) - return mx.argmax(logits[:, -1, :], axis=-1) + return logits[:, -1, :] def _decode_auxiliary_layer( self, @@ -1107,15 +1496,15 @@ class MlxModelRunner: caches: list[list[Any]], input_ids_by_request: list[mx.array], ) -> mx.array: - lazy_token_list = [] + lazy_logits_list = [] for input_ids, cache in zip(input_ids_by_request, caches): model_output = self.model(input_ids, cache=cache) logits = self._extract_logits(model_output) - lazy_token_list.append(mx.argmax(logits[:, -1, :], axis=-1)) + lazy_logits_list.append(logits[:, -1, :]) return ( - lazy_token_list[0] - if len(lazy_token_list) == 1 - else mx.concatenate(lazy_token_list, axis=0) + lazy_logits_list[0] + if len(lazy_logits_list) == 1 + else mx.concatenate(lazy_logits_list, axis=0) ) def _decode_with_batched_attention( @@ -1135,7 +1524,7 @@ class MlxModelRunner: ] model_output = self.model(batched_input, cache=shim_cache) logits = self._extract_logits(model_output) - return mx.argmax(logits[:, -1, :], axis=-1) + return logits[:, -1, :] finally: clear_context() @@ -1156,9 +1545,16 @@ class MlxModelRunner: attention_pool_index_by_layer=( self._cache_layout.attention_pool_index_by_layer ), + full_kv_pool_index_by_layer=self._cache_layout.full_kv_pool_index_by_layer, ) - def decode_batch_start(self, req_ids: list[str]) -> MlxPendingDecode: + def decode_batch_start( + self, + req_ids: list[str], + edit_rows: mx.array | None = None, + logprob_spec: MlxLogprobSpec | None = None, + logits_hook=None, + ) -> MlxPendingDecode: """Queue a decode forward pass without evaluating. The caller is responsible for calling ``mx.async_eval`` on the @@ -1170,18 +1566,31 @@ class MlxModelRunner: batched_input = mx.array(last_tokens, dtype=mx.int32)[:, None] if self._cache_layout.has_auxiliary_state: - lazy_tokens = self._decode_with_hybrid_batching( + last_logits = self._decode_with_hybrid_batching( caches, batched_input, list(req_ids) ) else: - lazy_tokens = self._decode_with_batched_attention( + last_logits = self._decode_with_batched_attention( caches, batched_input, list(req_ids) ) + if logits_hook is not None: + # CUDA edit order: grammar mask + logit_bias first, custom + # processors second, sanitization last (inside selection). + if edit_rows is not None: + last_logits = last_logits.astype(mx.float32) + edit_rows + edit_rows = None + last_logits = self._run_logits_hook(last_logits, logits_hook) + lazy_tokens, lazy_logprobs = self._select_tokens_with_logprobs( + last_logits, list(req_ids), caches, edit_rows, logprob_spec + ) return MlxPendingDecode( lazy_tokens=lazy_tokens, req_ids=list(req_ids), caches=caches, + lazy_logprobs=lazy_logprobs, + logprob_spec=logprob_spec, + edit_rows=edit_rows, ) def decode_batch_start_chained( @@ -1215,18 +1624,24 @@ class MlxModelRunner: # be written at in step N+1 (and equivalently the RoPE offset). batched_input = prev.lazy_tokens[:, None] if self._cache_layout.has_auxiliary_state: - lazy_tokens = self._decode_with_hybrid_batching( + last_logits = self._decode_with_hybrid_batching( caches, batched_input, prev.req_ids ) else: - lazy_tokens = self._decode_with_batched_attention( + last_logits = self._decode_with_batched_attention( caches, batched_input, prev.req_ids ) + lazy_tokens, lazy_logprobs = self._select_tokens_with_logprobs( + last_logits, prev.req_ids, caches, prev.edit_rows, prev.logprob_spec + ) return MlxPendingDecode( lazy_tokens=lazy_tokens, req_ids=prev.req_ids, caches=caches, + lazy_logprobs=lazy_logprobs, + logprob_spec=prev.logprob_spec, + edit_rows=prev.edit_rows, ) def decode_batch_finalize( @@ -1266,6 +1681,7 @@ class MlxModelRunner: self._sync_decode_kv_to_pool(req_id) self._req_token_ids.pop(req_id, None) + self._req_sampling.pop(req_id, None) cache = self._req_caches.pop(req_id, None) if cache is not None: self._release_cache(cache) @@ -1275,6 +1691,7 @@ class MlxModelRunner: def clear(self): """Clear all request states.""" self._req_token_ids.clear() + self._req_sampling.clear() for cache in self._req_caches.values(): self._release_cache(cache) self._req_caches.clear() diff --git a/python/sglang/srt/hardware_backend/mlx/sampling.py b/python/sglang/srt/hardware_backend/mlx/sampling.py new file mode 100644 index 000000000..fd40e2234 --- /dev/null +++ b/python/sglang/srt/hardware_backend/mlx/sampling.py @@ -0,0 +1,491 @@ +"""MLX-native in-graph sampling for the MLX backend. + +Token selection (temperature / top-k / top-p / min-p / per-request seed) +built entirely from ``mx`` ops, so it lives inside the same lazy graph as +the forward pass. This is what lets sampling coexist with the overlap +scheduler: ``decode_batch_start_chained`` feeds step N's still-unevaluated +sampled tokens as step N+1's input ids, exactly as it does for greedy +argmax, and the GPU runs both steps back-to-back with no host sync. + +An earlier proposal (#25804) bridged MLX logits to the CPU pytorch +``Sampler`` instead. That design forces a host sync in the middle of the +graph-build window, which is precisely what the overlap scheduler exists +to avoid, so token selection is rebuilt from ``mx`` ops here. + +Semantics mirror the sglang pytorch sampling backend +(``top_k_top_p_min_p_sampling_from_probs_torch`` / +``multinomial_with_seed`` in ``sglang/srt/layers/sampler.py``): + +* ``probs = softmax(logits / temperature)`` per row. +* Descending sort, then zero out rank >= top_k, cumulative-prob mass + beyond top_p (the top token is always kept), and probs below + ``max_prob * min_p``. +* Multinomial sampling via the Gumbel-max identity: + ``argmax(log(weights) + gumbel_noise)`` over the masked, unnormalized + weights is distributed identically to ``torch.multinomial(weights)`` + (normalization only shifts ``log`` by a per-row constant). +* When every row asks for a small enough ``top_k`` + (:data:`MAX_BOUNDED_TOP_K`), everything after the sort runs on the + ``[B, K]`` candidates instead of ``[B, vocab]``. Weights are zero + outside those K, so their ``log`` is ``-inf`` and the full-vocab + argmax could never have picked them — same token, a fraction of the + work. Any row without a bounded ``top_k`` sends the batch back to the + full-vocab chain. +* Rows with ``sampling_seed`` set use deterministic Gumbel noise derived + from the same MurmurHash3 formula as the CUDA kernel + (``sglang/kernels/ops/sampling/murmur_hash.py``): hash(seed, position, + token_id) -> uniform -> ``-log(-log(u))``. Seeded noise is keyed on + the token id in every branch (the full-vocab chain scatters the masked + weights back through the sort; the bounded chain hashes the candidate + ids), so a seeded row's token never depends on whether a batchmate + triggered top-k/top-p/min-p filtering or on which chain ran. +* Greedy rows (``top_k == 1`` after sglang normalization, which rewrites + ``temperature < eps`` to ``temperature=1, top_k=1``) short-circuit to + ``argmax`` and consume no randomness. + +Seeds follow the same gate as every other backend: ``sampling_seed`` is +consumed only under ``--enable-deterministic-inference``, which then +seeds every row (:data:`DEFAULT_SAMPLING_SEED` for requests that did not +ask for one). See :meth:`MlxSamplingParams.from_req`. + +Known deviations from the pytorch backend (not bugs): + +* Seeded determinism is MLX-local: noise math runs in float32 (Metal has + no float64) and tie order follows MLX's sort, so the same seed on a + CUDA backend may pick a different token from the same distribution. +* Unseeded rows draw their Gumbel noise in whichever space the chain is + running (candidate or vocab), so the bounded top-K path consumes the + RNG differently from the full-vocab one. Seeded rows are unaffected — + they hash the token id — so ``--enable-deterministic-inference`` is + bit-for-bit identical either way. +* Penalties (frequency/presence/repetition) are not applied on the MLX + path (warned once per process). #25804 skips them as well. +* Custom logit processors run on pure-decode steps only: the first + generated token and decode steps mixed into an extend batch are not + processed (``apply_custom_logit_processor`` requires logits rows to + match the full ``sampling_info``). Same scope as #25804, which only + hooked the pure-decode path at all. +* Logprob output covers the sampled token, top-k, and requested token + ids for every generated token; prompt/input logprobs + (``logprob_start_len``) are not computed. + +Logit edits (grammar vocab masks, ``logit_bias``) arrive as a +pre-combined additive [B, vocab] array built by the worker at graph +launch — grammar FSM state is always current at a fresh launch because +the previous token was finalized before scheduling, so the mask is known +at build time and the graph stays lazy. Grammar/custom-processor +batches must not CHAIN (the mask for step N+1 needs token N +materialized); the scheduler breaks the chain for them. NaN/inf +sanitization mirrors ``sanitize_nan_logits`` and is gated on the same +``SGLANG_SANITIZE_NAN_LOGITS`` env var. +""" + +from __future__ import annotations + +import logging +from dataclasses import dataclass +from typing import Any + +import mlx.core as mx + +logger = logging.getLogger(__name__) + +# Seed given to rows without an explicit ``sampling_seed`` when +# --enable-deterministic-inference is on. Mirrors the literal in +# ``SamplingBatchInfo.from_schedule_batch``. +DEFAULT_SAMPLING_SEED = 42 + +# Largest ``top_k`` that still takes the bounded candidate path in +# :func:`sample_tokens`. Past this the [B, K] chain stops being +# meaningfully cheaper than the [B, vocab] one it replaces. +MAX_BOUNDED_TOP_K = 1024 + +_warned_ignored_penalties = False + + +@dataclass(frozen=True) +class MlxSamplingParams: + """Per-request sampling parameters, frozen at prefill registration.""" + + temperature: float + top_k: int + top_p: float + min_p: float + seed: int | None + + @classmethod + def from_req( + cls, req: Any, deterministic_seeding: bool = False + ) -> MlxSamplingParams: + sp = req.sampling_params + global _warned_ignored_penalties + if not _warned_ignored_penalties and ( + sp.frequency_penalty != 0.0 + or sp.presence_penalty != 0.0 + or sp.repetition_penalty != 1.0 + ): + _warned_ignored_penalties = True + logger.warning( + "MLX sampling ignores frequency/presence/repetition penalties; " + "a request specified them. (Warning logged once.)" + ) + # Seed contract, identical to every other backend: SamplingBatchInfo + # populates sampling_seed only under --enable-deterministic-inference, + # and then seeds every row (default DEFAULT_SAMPLING_SEED). Outside + # that flag a per-request seed is ignored, so it is ignored here too. + seed = None + if deterministic_seeding: + seed = ( + sp.sampling_seed + if sp.sampling_seed is not None + else DEFAULT_SAMPLING_SEED + ) + return cls( + temperature=sp.temperature, + top_k=sp.top_k, + top_p=sp.top_p, + min_p=sp.min_p, + seed=seed, + ) + + @property + def is_greedy(self) -> bool: + return self.top_k == 1 + + +GREEDY_PARAMS = MlxSamplingParams( + temperature=1.0, top_k=1, top_p=1.0, min_p=0.0, seed=None +) + + +@dataclass(frozen=True) +class MlxLogprobSpec: + """Per-row logprob request for one step. + + Mirrors the CUDA ``OutputLogprobProcessor`` inputs: ``top_ks[i]`` is + the row's ``top_logprobs_num`` (0 = none) and ``token_ids[i]`` the + row's requested token ids (None = none). + """ + + top_ks: tuple[int, ...] + token_ids: tuple[tuple[int, ...] | None, ...] + + +@dataclass +class MlxLazyLogprobs: + """Lazy logprob arrays for one step; materialized at finalize.""" + + chosen: mx.array # [B] + top_val: mx.array | None # [B, max_k] + top_idx: mx.array | None # [B, max_k] + token_ids_val: list[mx.array | None] # per row + spec: MlxLogprobSpec + + +@dataclass +class MlxStepLogprobs: + """Materialized per-step logprobs, per row, cut to the request shape.""" + + chosen: list[float] + top_val: list[list[float]] + top_idx: list[list[int]] + token_ids_val: list[list[float]] + token_ids_idx: list[list[int]] + + +def lazy_logprob_arrays(lazy_logprobs: MlxLazyLogprobs | None) -> list[mx.array]: + """The mx arrays of a lazy logprob bundle, for eval/async_eval calls.""" + if lazy_logprobs is None: + return [] + arrays = [lazy_logprobs.chosen] + if lazy_logprobs.top_val is not None: + arrays += [lazy_logprobs.top_val, lazy_logprobs.top_idx] + arrays += [a for a in lazy_logprobs.token_ids_val if a is not None] + return arrays + + +def all_greedy(params: list[MlxSamplingParams]) -> bool: + return all(p.is_greedy for p in params) + + +def sanitize_logits(logits: mx.array) -> mx.array: + """Lazy analogue of sglang's ``sanitize_nan_logits``: NaN -> -1e30, + +-inf -> +-1e30 (not dtype extremes — temperature division would + overflow those back to inf).""" + return mx.clip(mx.where(mx.isnan(logits), -1e30, logits), -1e30, 1e30) + + +def scale_by_temperature( + last_logits: mx.array, params: list[MlxSamplingParams] +) -> mx.array: + """``logits / temperature`` per row, in float32. + + Both :func:`sample_tokens` and :func:`compute_logprobs` start here. MLX + builds eager graphs and does not eliminate common subexpressions, so a + step that samples *and* reports logprobs would otherwise pay two + full-vocab divisions; callers needing both compute this once and pass it + to each. + """ + temps = mx.array([p.temperature for p in params], dtype=mx.float32)[:, None] + return last_logits.astype(mx.float32) / temps + + +def compute_logprobs( + last_logits: mx.array, + params: list[MlxSamplingParams], + tokens: mx.array, + spec: MlxLogprobSpec, + scaled: mx.array | None = None, +) -> MlxLazyLogprobs: + """Lazy log-probabilities of this step's distribution. + + Matches the pytorch sampler: the distribution is + ``log_softmax(edited_logits / temperature)`` — after grammar-mask / + logit_bias / sanitization, before top-k/top-p/min-p filtering (the + filters affect which token is drawn, not the reported logprobs). + + ``scaled`` optionally supplies :func:`scale_by_temperature`'s result when + the caller already built it for :func:`sample_tokens`. + """ + if scaled is None: + scaled = scale_by_temperature(last_logits, params) + logp = scaled - mx.logsumexp(scaled, axis=-1, keepdims=True) + + chosen = mx.take_along_axis(logp, tokens[:, None], axis=-1).squeeze(-1) + + max_k = max(spec.top_ks) if spec.top_ks else 0 + if max_k > 0: + top_idx = mx.argsort(-logp, axis=-1)[:, :max_k] + top_val = mx.take_along_axis(logp, top_idx, axis=-1) + else: + top_idx = None + top_val = None + + token_ids_val: list[mx.array | None] = [ + logp[row, mx.array(ids)] if ids else None + for row, ids in enumerate(spec.token_ids) + ] + return MlxLazyLogprobs( + chosen=chosen, + top_val=top_val, + top_idx=top_idx, + token_ids_val=token_ids_val, + spec=spec, + ) + + +def sample_tokens( + last_logits: mx.array, + params: list[MlxSamplingParams], + positions: list[int], + key: mx.array, + scaled: mx.array | None = None, +) -> mx.array: + """Select one token per row of ``last_logits`` ([B, vocab], lazy ok). + + Pure ``mx`` ops — the result stays inside the lazy graph. ``positions`` + are the absolute sequence positions of the tokens being sampled (only + consumed for seeded rows). Callers should shortcut to ``mx.argmax`` + when ``all_greedy(params)`` — this function assumes at least one row + samples. ``scaled`` optionally supplies + :func:`scale_by_temperature`'s result when the caller already built it + for :func:`compute_logprobs`. + """ + batch_size, vocab_size = last_logits.shape + logits32 = last_logits.astype(mx.float32) + if scaled is None: + scaled = scale_by_temperature(logits32, params) + + filtering = any( + not p.is_greedy and (p.top_k < vocab_size or p.top_p < 1.0 or p.min_p > 0.0) + for p in params + ) + # Token ids the Gumbel-max runs over: None means "the whole vocabulary", + # otherwise a [B, K] candidate array that the argmax result indexes into. + candidates: mx.array | None = None + if filtering: + probs = mx.softmax(scaled, axis=-1) + width = _candidate_width(params, vocab_size) + sorted_idx = mx.argsort(-probs, axis=-1)[:, :width] + p_sort = mx.take_along_axis(probs, sorted_idx, axis=-1) + ranks = mx.arange(width, dtype=mx.int32)[None, :] + # SamplingParams normalizes top_k=-1 to TOP_K_ALL and temperature + # below eps to (temperature=1, top_k=1), so min(top_k, width) + # is always >= 1: rank 0 survives and log(weights) is never all -inf. + top_ks = mx.array([min(p.top_k, width) for p in params], dtype=mx.int32)[ + :, None + ] + top_ps = mx.array([p.top_p for p in params], dtype=mx.float32)[:, None] + min_ps = mx.array([p.min_p for p in params], dtype=mx.float32)[:, None] + cum = mx.cumsum(p_sort, axis=-1) + masked_out = ( + (ranks >= top_ks) + | ((cum - p_sort) > top_ps) + | (p_sort < p_sort[:, :1] * min_ps) + ) + w_sort = mx.where(masked_out, 0.0, p_sort) + if width < vocab_size: + # Bounded top-K: every surviving rank is inside the K + # candidates, so run the rest of the chain in candidate space + # and map the winning rank back to its vocab id at the end. + log_weights = mx.log(w_sort) + candidates = sorted_idx + else: + # Scatter the masked weights back to vocab order: noise must be + # applied in vocab-id space in every branch, or a seeded row's + # token would change when a batchmate happens to need filtering. + weights = mx.put_along_axis( + mx.zeros_like(probs), sorted_idx, w_sort, axis=-1 + ) + log_weights = mx.log(weights) + else: + # Nothing is masked, so the weights are the plain softmax and + # ``log(softmax(scaled)) == scaled - logsumexp(scaled)``: a per-row + # constant offset, which the argmax below is invariant to. Feeding + # the scaled logits straight to the Gumbel-max drops two full-vocab + # passes (softmax, log) on the common temperature-only batch. + log_weights = scaled + + noise = _gumbel_noise( + params=params, + positions=positions, + shape=log_weights.shape, + key=key, + columns=candidates, + ) + # Gumbel-max over the UNNORMALIZED masked weights: normalization would + # only shift log(w) by a per-row constant, which argmax is invariant to. + # That is also why seed + min_p is well-defined here, and why the + # pytorch backend's `assert sampling_seed is None` under min-p (and its + # TODO at layers/sampler.py "probs_sort should be re-normalized for the + # use of multinomial_with_seed") has no analogue on this path. + sampled = mx.argmax(log_weights + noise, axis=-1) + if candidates is not None: + sampled = mx.take_along_axis(candidates, sampled[:, None], axis=-1).squeeze(-1) + + greedy = [p.is_greedy for p in params] + if not any(greedy): + return sampled + # A batch that mixes greedy rows in still runs them through the sampled + # path above (the row exists either way); overwrite those rows with the + # unnoised argmax, which is what makes greedy rows consume no randomness. + return mx.where(mx.array(greedy), mx.argmax(logits32, axis=-1), sampled) + + +def _candidate_width(params: list[MlxSamplingParams], vocab_size: int) -> int: + """Rank cut-off the filtered chain can run on, or ``vocab_size``. + + Every rank at or beyond a row's ``top_k`` is masked to weight 0, whose + ``log`` is ``-inf``, so the Gumbel-max can never select it. When every + row's ``top_k`` is small, the whole chain after the sort — gather, + cumsum, mask, log, noise, argmax — can therefore run on ``[B, K]`` + instead of ``[B, vocab]`` and still pick exactly the same token. + + Falls back to the full vocabulary as soon as one row wants more + candidates than :data:`MAX_BOUNDED_TOP_K` (or no top-k at all, which + ``SamplingParams`` spells as ``top_k = TOP_K_ALL``). + """ + largest_top_k = max(p.top_k for p in params) + if largest_top_k > MAX_BOUNDED_TOP_K or largest_top_k >= vocab_size: + return vocab_size + return largest_top_k + + +def _gumbel_noise( + params: list[MlxSamplingParams], + positions: list[int], + shape: tuple[int, int], + key: mx.array, + columns: mx.array | None = None, +) -> mx.array: + """Gumbel noise shaped like the weights: hashed if seeded, RNG otherwise. + + ``columns`` is the [B, K] token ids each weight column stands for on the + bounded top-K path; ``None`` means column j is token id j. Seeded rows + hash the token id, so their noise — and therefore their token — is the + same either way. + """ + seeded_rows = [p.seed is not None for p in params] + if not any(seeded_rows): + return mx.random.gumbel(shape=shape, key=key) + + hashed = _murmur_hash32( + seeds=[p.seed if p.seed is not None else 0 for p in params], + positions=positions, + vocab_size=shape[1], + columns=columns, + ) + u = hashed.astype(mx.float32) / float(0xFFFFFFFF) + # REQUIRED, not cosmetic: uint32(0xFFFFFFFF) rounds UP to 2**32 in + # float32, so the quotient can land just above 1.0 and make + # log(-log(u)) NaN; and an exact 1.0 gives -log(-log(1)) = +inf, which + # would deterministically force that token. Clamp both ends to the + # nearest representable interior values. + u = mx.clip(u, 2.0**-32, 1.0 - 2.0**-24) + hash_noise = -mx.log(-mx.log(u)) + + if all(seeded_rows): + return hash_noise + random_noise = mx.random.gumbel(shape=shape, key=key) + return mx.where(mx.array(seeded_rows)[:, None], hash_noise, random_noise) + + +def _murmur3_mix_py(h: int, k: int) -> int: + """One MurmurHash3 block mix on Python ints (exact 32-bit semantics).""" + k = (k * 0xCC9E2D51) & 0xFFFFFFFF + k = ((k << 15) | (k >> 17)) & 0xFFFFFFFF + k = (k * 0x1B873593) & 0xFFFFFFFF + h ^= k + h = ((h << 13) | (h >> 19)) & 0xFFFFFFFF + h = (h * 5 + 0xE6546B64) & 0xFFFFFFFF + return h + + +def _murmur_hash32( + seeds: list[int], + positions: list[int], + vocab_size: int, + columns: mx.array | None = None, +) -> mx.array: + """Port of ``murmur_hash32`` (Triton) to mx ops: [B, V] uint32. + + Blocks mixed in kernel order: seed_low, seed_high, position, column. + The first three are per-row scalars, so they are folded exactly on the + CPU with Python ints; only the column block and finalization run as + vectorized uint32 ops (verified to wrap like the Triton kernel). + + ``columns`` hashes an explicit [B, K] set of token ids instead of every + id in ``[0, vocab_size)`` — same value per (row, token id) either way, + which is what keeps a seeded row's token identical on the bounded + top-K path. + """ + row_states = [] + for seed, pos in zip(seeds, positions): + seed &= 0xFFFFFFFFFFFFFFFF + h = _murmur3_mix_py(0, seed & 0xFFFFFFFF) + h = _murmur3_mix_py(h, (seed >> 32) & 0xFFFFFFFF) + h = _murmur3_mix_py(h, pos & 0xFFFFFFFF) + row_states.append(h) + + h = mx.array(row_states, dtype=mx.uint32)[:, None] + if columns is None: + k = mx.arange(vocab_size, dtype=mx.uint32)[None, :] + else: + k = columns.astype(mx.uint32) + + # murmur3_mix on [B, V] + k = k * mx.array(0xCC9E2D51, dtype=mx.uint32) + k = (k << 15) | (k >> 17) + k = k * mx.array(0x1B873593, dtype=mx.uint32) + h = h ^ k + h = (h << 13) | (h >> 19) + h = h * mx.array(5, dtype=mx.uint32) + mx.array(0xE6546B64, dtype=mx.uint32) + + # finalize: len = 16 bytes (seed + pos + col), then fmix32 + h = h ^ mx.array(16, dtype=mx.uint32) + h = h ^ (h >> 16) + h = h * mx.array(0x85EBCA6B, dtype=mx.uint32) + h = h ^ (h >> 13) + h = h * mx.array(0xC2B2AE35, dtype=mx.uint32) + h = h ^ (h >> 16) + return h diff --git a/python/sglang/srt/hardware_backend/mlx/scheduler_mixin.py b/python/sglang/srt/hardware_backend/mlx/scheduler_mixin.py index c462c0342..aecf0ba8b 100644 --- a/python/sglang/srt/hardware_backend/mlx/scheduler_mixin.py +++ b/python/sglang/srt/hardware_backend/mlx/scheduler_mixin.py @@ -17,23 +17,20 @@ from __future__ import annotations import logging import time -from dataclasses import dataclass +from dataclasses import dataclass, replace from typing import TYPE_CHECKING, List, Optional import mlx.core as mx from sglang.srt.environ import envs from sglang.srt.managers.overlap_utils import resolve_forward_inputs +from sglang.srt.runtime_context import get_device from sglang.srt.utils import DynamicGradMode logger = logging.getLogger(__name__) if TYPE_CHECKING: - from sglang.srt.hardware_backend.mlx.model_runner import ( - MlxPendingDecode, - MlxPendingExtend, - MlxPendingPrefill, - ) + from sglang.srt.hardware_backend.mlx.tp_worker import MlxLaunch from sglang.srt.managers.schedule_batch import Req, ScheduleBatch from sglang.srt.managers.scheduler import Scheduler @@ -43,23 +40,10 @@ class MlxPendingJob: """Unfinished MLX work and graphs queued on the GPU. Attributes: - lazy_tokens: Lazily evaluated token IDs produced by the forward - pass. Unevaluated; calling ``.tolist()`` / ``.item()`` / - ``mx.eval`` on it will block until the Metal kernel finishes. - ``None`` for idle batches. - prefills: MLX prefill state returned by the model worker — one - entry per new request in an extend batch. Used by - ``finalize_mlx_result`` to commit per-request caches. Empty - list for pure-decode steps. - extends: Chunked-prefill-continuation state, one entry per - already-active request whose extend seq_len > 1. Also empty - for pure-decode steps. - decode: Decode state covering full-decode mode AND mixed - single-token decodes inside an extend batch. Used as the - chaining root by :meth:`async_chained_decode_mlx`. - mode: One of ``"decode"``, ``"extend"``, ``"idle"`` describing - which forward pass produced this job. Drives finalise - dispatch and whether chaining is safe. + launch: The :class:`MlxLaunch` this job is waiting on — the lazy + token handle plus the prefill / extend / decode pendings the + forward produced, and the mode that drives finalise dispatch + and whether chaining is safe. batch_copy: Snapshot of the :class:`ScheduleBatch` at launch time. Decoupled from the live batch so ``process_batch_result`` can update request state without @@ -73,19 +57,41 @@ class MlxPendingJob: mutable batch object. """ - lazy_tokens: Optional[mx.array] - prefills: list[MlxPendingPrefill] - extends: list[MlxPendingExtend] - decode: Optional[MlxPendingDecode] - mode: str + launch: MlxLaunch batch_copy: ScheduleBatch schedule_batch: ScheduleBatch reqs: List[Req] + # False when the batch needs per-step CPU logit state (grammar vocab + # masks / custom logit processors under --mlx-enable-sampling): step + # N+1's mask needs token N materialized, so such batches must launch + # fresh every step instead of chaining. + chain_safe: bool = True + # Captured at launch when batch.return_logprob, exactly like + # Scheduler.run_batch does for the CUDA paths (the live values mutate + # before output processing). + extend_input_len_per_req: Optional[List[int]] = None + extend_logprob_start_len_per_req: Optional[List[int]] = None class SchedulerMlxOverlapMixin: """Mixin that adds MLX overlap scheduling to :class:`Scheduler`.""" + def _mlx_batch_chain_safe(self: Scheduler, batch: ScheduleBatch) -> bool: + """False when per-step CPU logit state forbids chained decode. + + Grammar vocab masks and custom logit processors depend on the + previous token being materialized; a chained step is built before + that, so those batches launch fresh every step. + """ + if not get_device().mlx_enable_sampling: + return True + sampling_info = batch.sampling_info + if sampling_info is None: + return True + # batch.has_grammar, not sampling_info.grammars: the latter is only + # populated at forward launch (see _build_logit_edit_rows). + return not (batch.has_grammar or sampling_info.has_custom_logit_processor) + def _prepare_mlx_launch(self: Scheduler, batch: ScheduleBatch): """Stamp scheduler bookkeeping before an MLX forward is launched.""" # Match run_batch's launch boundary. In particular, the profiler @@ -98,12 +104,10 @@ class SchedulerMlxOverlapMixin: self.profiler_manager._profile_batch_predicate(batch) def _finalize_mlx_pending_job(self: Scheduler, pending: MlxPendingJob): - result = self.tp_worker.finalize_mlx_result( - pending.prefills, - pending.extends, - pending.decode, - pending.mode, - pending.reqs, + result = self.tp_worker.finalize_mlx_result(pending.launch, pending.reqs) + result.extend_input_len_per_req = pending.extend_input_len_per_req + result.extend_logprob_start_len_per_req = ( + pending.extend_logprob_start_len_per_req ) if result.next_token_ids is not None: pending.batch_copy.input_ids = result.next_token_ids @@ -163,29 +167,31 @@ class SchedulerMlxOverlapMixin: # loop must do it too, otherwise async_forward_batch_generation_mlx # dereferences a None input_ids. resolve_forward_inputs(batch, self.future_map) - # run_batch stamps launch_ts on every scheduler-built forward; the - # MLX overlap loop bypasses run_batch, and process_batch_result -> - # _record_step_counters subtracts launch_ts unconditionally for - # prefill/decode batches. ScheduleBatch.copy() below carries the - # stamp to process_batch_result. - lazy_tokens, prefills, extends, decode, mode = ( - self.tp_worker.async_forward_batch_generation_mlx(batch) - ) + launch = self.tp_worker.async_forward_batch_generation_mlx(batch) + extend_input_len_per_req = None + extend_logprob_start_len_per_req = None + if batch.return_logprob: + # Mirror Scheduler.run_batch's launch-time copy. + extend_input_len_per_req = [ + req.extend_range.length if req.extend_range is not None else 0 + for req in batch.reqs + ] + extend_logprob_start_len_per_req = batch.extend_logprob_start_lens return MlxPendingJob( - lazy_tokens=lazy_tokens, - prefills=prefills, - extends=extends, - decode=decode, - mode=mode, + launch=launch, batch_copy=batch.copy(), schedule_batch=batch, reqs=list(batch.reqs), + chain_safe=self._mlx_batch_chain_safe(batch), + extend_input_len_per_req=extend_input_len_per_req, + extend_logprob_start_len_per_req=extend_logprob_start_len_per_req, ) def _launch_chained(prev: MlxPendingJob) -> MlxPendingJob: - assert prev.decode is not None - # Composition is identical to prev: reuse a fresh batch copy - # of the same underlying ScheduleBatch so process_batch_result + assert prev.launch.decode is not None + # Composition is identical to prev: every scheduler-side field + # carries over, and only a fresh batch copy of the same + # underlying ScheduleBatch is needed so process_batch_result # updates the same req objects with the new token. batch_copy = prev.batch_copy.copy() self._prepare_mlx_launch(batch_copy) @@ -193,18 +199,10 @@ class SchedulerMlxOverlapMixin: # chain breaks, prepare_for_decode() may run SWA maintenance # before the next fresh launch gets a chance to re-stamp it. prev.schedule_batch.forward_iter = batch_copy.forward_iter - lazy_tokens, prefills, extends, decode, mode = ( - self.tp_worker.async_chained_decode_mlx(prev.decode) - ) - return MlxPendingJob( - lazy_tokens=lazy_tokens, - prefills=prefills, - extends=extends, - decode=decode, - mode=mode, + return replace( + prev, + launch=self.tp_worker.async_chained_decode_mlx(prev.launch.decode), batch_copy=batch_copy, - schedule_batch=prev.schedule_batch, - reqs=prev.reqs, ) while True: @@ -224,8 +222,9 @@ class SchedulerMlxOverlapMixin: # build pending_next on top of it NOW — before we block on curr. can_chain = ( pending_curr is not None - and pending_curr.mode == "decode" - and pending_curr.decode is not None + and pending_curr.launch.mode == "decode" + and pending_curr.launch.decode is not None + and pending_curr.chain_safe and not self.waiting_queue ) if can_chain and pending_next is None: diff --git a/python/sglang/srt/hardware_backend/mlx/tp_worker.py b/python/sglang/srt/hardware_backend/mlx/tp_worker.py index d79ac6784..438936aa7 100644 --- a/python/sglang/srt/hardware_backend/mlx/tp_worker.py +++ b/python/sglang/srt/hardware_backend/mlx/tp_worker.py @@ -13,7 +13,8 @@ normal ``GenerationBatchResult``. """ import logging -from typing import Optional, Union +from dataclasses import dataclass +from typing import Optional import mlx.core as mx import torch @@ -23,6 +24,11 @@ from sglang.srt.hardware_backend.mlx.model_runner import ( MlxPendingExtend, MlxPendingPrefill, ) +from sglang.srt.hardware_backend.mlx.sampling import ( + MlxLogprobSpec, + MlxStepLogprobs, + lazy_logprob_arrays, +) from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.managers.utils import GenerationBatchResult @@ -31,11 +37,45 @@ from sglang.srt.model_executor.forward_batch_info import ( ForwardBatch, PPProxyTensors, ) -from sglang.srt.runtime_context import get_memory, get_model, get_schedule +from sglang.srt.runtime_context import ( + get_device, + get_exec, + get_memory, + get_model, + get_schedule, +) logger = logging.getLogger(__name__) +@dataclass +class MlxLaunch: + """One lazily launched MLX forward pass: its handle and its pending work. + + Produced by :meth:`MlxTpModelWorker.async_forward_batch_generation_mlx` + and :meth:`MlxTpModelWorker.async_chained_decode_mlx`, consumed by + :meth:`MlxTpModelWorker.finalize_mlx_result`. + + Attributes: + lazy_tokens: an ``mx.array`` that, when evaluated, forces + materialisation of the whole batch's outputs. ``None`` for + idle batches. + prefills: one :class:`MlxPendingPrefill` per new request in an + extend batch; empty for pure-decode steps. + extends: one :class:`MlxPendingExtend` per chunked-prefill + continuation; also empty for pure-decode steps. + decode: the :class:`MlxPendingDecode` covering full decode mode + AND mixed single-token decodes inside an extend batch. + mode: one of ``"idle"``, ``"decode"``, ``"extend"``. + """ + + lazy_tokens: Optional[mx.array] + prefills: list[MlxPendingPrefill] + extends: list[MlxPendingExtend] + decode: Optional[MlxPendingDecode] + mode: str + + class MlxTpModelWorker(TpModelWorker): """A tensor parallel model worker that routes inference through MLX. @@ -59,6 +99,11 @@ class MlxTpModelWorker(TpModelWorker): disable_radix_cache=get_memory().disable_radix_cache, mem_fraction_static=get_schedule().mem_fraction_static, quantization=get_model().quantization, + enable_sampling=get_device().mlx_enable_sampling, + sampling_rng_seed=get_device().random_seed, + deterministic_seeding=( + get_exec().deterministic.enable_deterministic_inference + ), ) if get_schedule().max_total_tokens is not None: init_kwargs["pool_size"] = get_schedule().max_total_tokens @@ -137,8 +182,8 @@ class MlxTpModelWorker(TpModelWorker): def _route_extend_request(self, rid: str, decoding_rids: set[str]) -> str: """Classify a request within an extend / mixed batch. - Shared by the sync (:meth:`_forward_batch_generation_mlx`) and async - (:meth:`_async_extend_batch`) paths so both route identically. + Called once per request from :meth:`_async_extend_batch`, which both + the overlap loop and the synchronous entry point launch through. Returns one of: @@ -157,134 +202,220 @@ class MlxTpModelWorker(TpModelWorker): return "decode" return "continuation" + @staticmethod + def _chunk_needs_logits(req) -> bool: + """False iff this extend chunk is a non-final chunked-prefill chunk. + + The scheduler truncates a chunked request's extend range below the + tokens it already knows about; such a chunk's next-token output is + discarded (the runner pops it as the stale intermediate token), so + the runner may skip the logit head for it. + """ + if req.extend_range is None: + return True + return req.extend_range.end >= len(req.full_untruncated_fill_ids) + + @staticmethod + def _sampling_active(batch: ScheduleBatch) -> bool: + return get_device().mlx_enable_sampling and batch.sampling_info is not None + + def _build_logit_edit_rows( + self, batch: ScheduleBatch + ) -> dict[str, mx.array] | None: + """Pre-combine grammar vocab masks and logit_bias into one additive + [vocab] float32 row per request, ready to enter the lazy graph. + + Grammar FSM state is current at every fresh launch — the previous + token was finalized before this batch was scheduled — so the mask + is knowable at graph-build time with no device sync. The + scheduler never chains grammar batches + (:attr:`MlxPendingJob.chain_safe`), so a chained step never needs + a stale mask. Mask application reuses the grammar backend's own + ``apply_vocab_mask`` on a zeros tensor, which keeps this + backend-agnostic (xgrammar / llguidance / outlines). + """ + if not self._sampling_active(batch): + return None + sinfo = batch.sampling_info + # Mirror ForwardBatch.init_new's grammars population — the MLX paths + # never build a ForwardBatch, so without this the list stays None + # even when requests carry live grammar objects. + sinfo.grammars = ( + [req.grammar for req in batch.reqs] if batch.has_grammar else None + ) + has_grammar = bool(sinfo.grammars) + if not has_grammar and sinfo.logit_bias is None: + return None + if not has_grammar: + # logit_bias alone is already the dense [B, vocab] additive row we + # want; converting it directly skips a second [B, vocab] float32 + # allocation and an add on every step (~6 MB of churn per step at + # vocab 200k, batch 8). Not mutated below, so no clone is needed. + combined = sinfo.logit_bias.to(device="cpu", dtype=torch.float32) + else: + combined = torch.zeros( + len(batch.reqs), sinfo.vocab_size, dtype=torch.float32 + ) + sinfo.update_regex_vocab_mask() + if sinfo.grammar_mask is not None: + grammar_mask = sinfo.grammar_mask + grammar_mask.grammar.apply_vocab_mask( + logits=combined, + vocab_mask=grammar_mask.vocab_mask.to("cpu"), + ) + # Release promptly; mirrors the VRAM-leak note in the CUDA + # ModelRunner._preprocess_logits. + sinfo.grammar_mask = None + if sinfo.logit_bias is not None: + combined += sinfo.logit_bias.to("cpu") + rows = mx.array(combined.numpy()) + return {req.rid: rows[i] for i, req in enumerate(batch.reqs)} + + def _logprob_rows( + self, batch: ScheduleBatch + ) -> dict[str, tuple[int, tuple[int, ...] | None]] | None: + """Per-request (top_logprobs_num, token_ids) for logprob output.""" + if not self._sampling_active(batch) or not batch.return_logprob: + return None + tops = batch.top_logprobs_nums or [0] * len(batch.reqs) + tids = batch.token_ids_logprobs or [None] * len(batch.reqs) + rows = {} + for req, top_k, token_ids in zip(batch.reqs, tops, tids): + if req.return_logprob: + rows[req.rid] = ( + int(top_k or 0), + tuple(token_ids) if token_ids else None, + ) + return rows or None + + @staticmethod + def _logprob_spec_for( + rows: dict[str, tuple[int, tuple[int, ...] | None]] | None, + rids: list[str], + ) -> MlxLogprobSpec | None: + if rows is None or not any(rid in rows for rid in rids): + return None + return MlxLogprobSpec( + top_ks=tuple(rows.get(rid, (0, None))[0] for rid in rids), + token_ids=tuple(rows.get(rid, (0, None))[1] for rid in rids), + ) + + def _custom_logits_hook(self, batch: ScheduleBatch): + """CPU edit hook for custom logit processors, or None. + + Only built for fresh pure-decode launches; the runner materializes + the logits for the hook, so these batches never chain. + """ + if not ( + self._sampling_active(batch) + and batch.sampling_info.has_custom_logit_processor + ): + return None + sinfo = batch.sampling_info + + def hook(logits_np): + from sglang.srt.layers.sampler import apply_custom_logit_processor + + # torch.from_numpy shares memory with logits_np, so the + # processors' in-place edits land in the returned array. + logits_t = torch.from_numpy(logits_np) + apply_custom_logit_processor(logits_t, sinfo) + return logits_np + + return hook + + @staticmethod + def _assemble_logprob_output(step_rows: dict[str, tuple], reqs: list): + """Batch-ordered LogitsProcessorOutput from per-request logprob rows. + + Field shapes follow what ``move_logprobs_to_cpu`` and + ``add_logprob_return_values`` consume: tensors for values the + scheduler ``.tolist()``s, plain lists for token-id indices. + """ + from sglang.srt.layers.logits_processor import LogitsProcessorOutput + + chosen, top_val, top_idx, tid_val, tid_idx = [], [], [], [], [] + for req in reqs: + row = step_rows.get(req.rid) + if row is None: + row = (0.0, [], [], [], []) + chosen.append(row[0]) + top_val.append(torch.tensor(row[1], dtype=torch.float32)) + top_idx.append(torch.tensor(row[2], dtype=torch.long)) + tid_val.append(torch.tensor(row[3], dtype=torch.float32)) + tid_idx.append(list(row[4])) + return LogitsProcessorOutput( + next_token_logits=None, + next_token_logprobs=torch.tensor(chosen, dtype=torch.float32), + next_token_top_logprobs_val=top_val, + next_token_top_logprobs_idx=top_idx, + next_token_token_ids_logprobs_val=tid_val, + next_token_token_ids_logprobs_idx=tid_idx, + ) + + @staticmethod + def _step_logprob_rows( + step: Optional[MlxStepLogprobs], rids: list[str] + ) -> dict[str, tuple]: + """Split a step's batch logprobs into per-request rows.""" + if step is None: + return {} + return { + rid: ( + step.chosen[i], + step.top_val[i], + step.top_idx[i], + step.token_ids_val[i], + step.token_ids_idx[i], + ) + for i, rid in enumerate(rids) + } + + def _collect_step_logprobs( + self, + step_rows: dict[str, tuple], + lazy_logprobs, + rids: list[str], + ) -> None: + """Materialize one pending's lazy logprobs into ``step_rows``.""" + step = self._mlx_runner.collect_logprobs(lazy_logprobs) + step_rows.update(self._step_logprob_rows(step, rids)) + def _forward_batch_generation_mlx( self, batch: ScheduleBatch ) -> GenerationBatchResult: - """Run forward pass through the MLX model runner (greedy only).""" - from sglang.srt.layers.logits_processor import LogitsProcessorOutput + """Run one forward pass through the MLX model runner, synchronously. - forward_mode = batch.forward_mode - reqs = batch.reqs + Reachable only under ``--disable-overlap-schedule``: the default MLX + loop drives :meth:`async_forward_batch_generation_mlx` / + :meth:`finalize_mlx_result` directly and never calls ``run_batch``. + Launching and finalising back-to-back IS the synchronous path — the + lazy graph is built exactly the same way, then blocked on + immediately — so routing, logit edits, logprob collection and + chunk-head skipping have one implementation rather than two that + must be kept in step. It is also strictly cheaper than evaluating + each request as it is queued: one ``mx.async_eval`` covers the whole + batch. + """ + launch = self.async_forward_batch_generation_mlx(batch) + return self.finalize_mlx_result(launch, batch.reqs) - if forward_mode.is_idle(): - return GenerationBatchResult( - logits_output=LogitsProcessorOutput(next_token_logits=None), - can_run_cuda_graph=False, - ) + @staticmethod + def _stacked_edit_rows( + edit_rows: dict[str, mx.array] | None, req_ids: list[str] + ) -> Optional[mx.array]: + """Stack the per-request additive edit rows for a decode sub-batch.""" + if not edit_rows: + return None + return mx.stack([edit_rows[rid] for rid in req_ids]) - self._cleanup_stale_rids(forward_mode, {req.rid for req in reqs}) - - next_token_ids_list: list[int] = [] - - if forward_mode.is_extend(): - # Ensure pool is up-to-date before pool-backed attention reads it - # for prefix-cached prefills. Only runs on extend batches. - self._mlx_runner.flush_all_decode_kv() - input_ids_cpu = batch.input_ids.cpu().tolist() - out_cache_loc_cpu = batch.out_cache_loc.cpu().tolist() - extend_seq_lens = batch.extend_lens - - offset = 0 # into input_ids_cpu - slot_offset = 0 # into out_cache_loc_cpu - prefill_rids: list[tuple[str, int]] = [] - extend_rids: list[tuple[str, int]] = [] - decode_rids: list[str] = [] - # Genuine decode steps mixed into this extend batch; see - # _route_extend_request. - decoding_rids = {r.rid for r in (batch.decoding_reqs or [])} - - for i, req in enumerate(reqs): - seq_len = extend_seq_lens[i] - req_token_ids = input_ids_cpu[offset : offset + seq_len] - req_new_slots = out_cache_loc_cpu[slot_offset : slot_offset + seq_len] - offset += seq_len - slot_offset += seq_len - - route = self._route_extend_request(req.rid, decoding_rids) - if route == "continuation": - next_token = self._mlx_runner.extend( - req.rid, req_token_ids, req_new_slots - ) - extend_rids.append((req.rid, next_token)) - elif route == "decode": - decode_rids.append(req.rid) - else: # "prefill" - prefix_slot_ids = req.prefix_indices.tolist() - full_token_ids = list(req.get_fill_ids()) - next_token = self._mlx_runner.prefill( - req_id=req.rid, - new_token_ids=req_token_ids, - full_token_ids=full_token_ids, - prefix_slot_ids=prefix_slot_ids, - new_slot_ids=req_new_slots, - req_pool_idx=req.req_pool_idx, - req=req, - ) - prefill_rids.append((req.rid, next_token)) - - # Batch decode all existing requests at once - if decode_rids: - decode_results = self._mlx_runner.decode_batch(decode_rids) - decode_map = dict(zip(decode_rids, decode_results)) - else: - decode_map = {} - - prefill_map = dict(prefill_rids) - extend_map = dict(extend_rids) - - for req in reqs: - if req.rid in decode_map: - next_token_ids_list.append(decode_map[req.rid]) - elif req.rid in extend_map: - next_token_ids_list.append(extend_map[req.rid]) - else: - next_token_ids_list.append(prefill_map[req.rid]) - - elif forward_mode.is_decode(): - req_ids = [req.rid for req in reqs] - next_token_ids_list = self._mlx_runner.decode_batch(req_ids) - - else: - raise ValueError( - f"MLX runner does not support forward mode: {forward_mode}" - ) - - next_token_ids = torch.tensor( - next_token_ids_list, dtype=torch.long, device="cpu" - ) - - return GenerationBatchResult( - logits_output=LogitsProcessorOutput(next_token_logits=None), - next_token_ids=next_token_ids, - can_run_cuda_graph=False, - ) - - def async_forward_batch_generation_mlx(self, batch: ScheduleBatch) -> tuple[ - Union[mx.array, None], - list[MlxPendingPrefill], - list[MlxPendingExtend], - Optional[MlxPendingDecode], - str, - ]: + def async_forward_batch_generation_mlx(self, batch: ScheduleBatch) -> MlxLaunch: """Start an async (lazy) forward pass through the MLX model runner. - Returns ``(lazy_result, prefills, extends, decode, mode)``: - - * ``lazy_result`` — an ``mx.array`` that, when evaluated, forces - materialisation of the whole batch's outputs. ``None`` for - idle batches. - * ``prefills`` — list of :class:`MlxPendingPrefill` for new - requests in an extend batch. - * ``extends`` — list of :class:`MlxPendingExtend` for chunked - prefill continuations in an extend batch. - * ``decode`` — :class:`MlxPendingDecode` for the decode - sub-batch (covers full decode mode AND mixed decodes inside - an extend batch). - * ``mode`` — one of ``"idle"``, ``"decode"``, ``"extend"``. - - The caller must make sure the returned pendings are fed into a - subsequent ``mx.async_eval`` or ``.item()`` / ``.tolist()`` call - — :meth:`finalize_mlx_result` does that. + See :class:`MlxLaunch` for the returned fields. The caller must + make sure the launch's pendings are fed into a subsequent + ``mx.async_eval`` or ``.item()`` / ``.tolist()`` call — + :meth:`finalize_mlx_result` does that. """ self._ensure_mlx_pool_initialized() @@ -292,15 +423,33 @@ class MlxTpModelWorker(TpModelWorker): reqs = batch.reqs if forward_mode.is_idle(): - return None, [], [], None, "idle" + return MlxLaunch( + lazy_tokens=None, prefills=[], extends=[], decode=None, mode="idle" + ) self._cleanup_stale_rids(forward_mode, {req.rid for req in reqs}) if forward_mode.is_decode(): req_ids = [req.rid for req in reqs] - pending_decode = self._mlx_runner.decode_batch_start(req_ids) - mx.async_eval(pending_decode.lazy_tokens) - return pending_decode.lazy_tokens, [], [], pending_decode, "decode" + pending_decode = self._mlx_runner.decode_batch_start( + req_ids, + edit_rows=self._stacked_edit_rows( + self._build_logit_edit_rows(batch), req_ids + ), + logprob_spec=self._logprob_spec_for(self._logprob_rows(batch), req_ids), + logits_hook=self._custom_logits_hook(batch), + ) + mx.async_eval( + pending_decode.lazy_tokens, + *lazy_logprob_arrays(pending_decode.lazy_logprobs), + ) + return MlxLaunch( + lazy_tokens=pending_decode.lazy_tokens, + prefills=[], + extends=[], + decode=pending_decode, + mode="decode", + ) if forward_mode.is_extend(): # TODO (changminbark): Implement per-batch flushing using prefix_slot_ids @@ -313,18 +462,14 @@ class MlxTpModelWorker(TpModelWorker): f"MLX async runner does not support forward mode: {forward_mode}" ) - def _async_extend_batch(self, batch: ScheduleBatch) -> tuple[ - Union[mx.array, None], - list[MlxPendingPrefill], - list[MlxPendingExtend], - Optional[MlxPendingDecode], - str, - ]: + def _async_extend_batch(self, batch: ScheduleBatch) -> MlxLaunch: """Launch each request in an EXTEND batch lazily and kick GPU work.""" reqs = batch.reqs input_ids_cpu = batch.input_ids.cpu().tolist() out_cache_loc_cpu = batch.out_cache_loc.cpu().tolist() extend_seq_lens = batch.extend_lens + edit_rows = self._build_logit_edit_rows(batch) + logprob_rows = self._logprob_rows(batch) offset = 0 slot_offset = 0 @@ -349,6 +494,9 @@ class MlxTpModelWorker(TpModelWorker): req_id=req.rid, new_token_ids=req_token_ids, new_slot_ids=req_new_slots, + needs_logits=self._chunk_needs_logits(req), + logit_edit_row=edit_rows[req.rid] if edit_rows else None, + logprob_spec=self._logprob_spec_for(logprob_rows, [req.rid]), ) ) elif route == "decode": @@ -365,13 +513,18 @@ class MlxTpModelWorker(TpModelWorker): new_slot_ids=req_new_slots, req_pool_idx=req.req_pool_idx, req=req, + needs_logits=self._chunk_needs_logits(req), + logit_edit_row=edit_rows[req.rid] if edit_rows else None, + logprob_spec=self._logprob_spec_for(logprob_rows, [req.rid]), ) ) pending_mixed_decode: Optional[MlxPendingDecode] = None if mixed_decode_rids: pending_mixed_decode = self._mlx_runner.decode_batch_start( - mixed_decode_rids + mixed_decode_rids, + edit_rows=self._stacked_edit_rows(edit_rows, mixed_decode_rids), + logprob_spec=self._logprob_spec_for(logprob_rows, mixed_decode_rids), ) # Stack lazy tokens so the caller has a single handle to evaluate @@ -389,51 +542,28 @@ class MlxTpModelWorker(TpModelWorker): else: lazy_stacked = None - for p in pending_prefills: - async_args.extend(self._cache_state(p.cache)) - for e in pending_extends: - async_args.extend(self._cache_state(self._mlx_runner._req_caches[e.req_id])) + for pending in (*pending_prefills, *pending_extends): + async_args.extend(self._mlx_runner.cache_state_arrays([pending.cache])) + async_args.extend(lazy_logprob_arrays(pending.lazy_logprobs)) if pending_mixed_decode is not None: async_args.append(pending_mixed_decode.lazy_tokens) - for c_list in pending_mixed_decode.caches: - async_args.extend(self._cache_state(c_list)) + async_args.extend(lazy_logprob_arrays(pending_mixed_decode.lazy_logprobs)) + async_args.extend( + self._mlx_runner.cache_state_arrays(pending_mixed_decode.caches) + ) if async_args: mx.async_eval(*async_args) - return ( - lazy_stacked, - pending_prefills, - pending_extends, - pending_mixed_decode, - "extend", + return MlxLaunch( + lazy_tokens=lazy_stacked, + prefills=pending_prefills, + extends=pending_extends, + decode=pending_mixed_decode, + mode="extend", ) - @staticmethod - def _cache_state(cache_list) -> list[mx.array]: - """Flatten a per-layer cache list to its ``state`` arrays.""" - arrays: list[mx.array] = [] - - def collect(value): - if isinstance(value, mx.array): - arrays.append(value) - elif value is None: - return - elif isinstance(value, (list, tuple)): - for item in value: - collect(item) - elif isinstance(value, dict): - for item in value.values(): - collect(item) - - for cache in cache_list: - collect(getattr(cache, "state", ())) - return arrays - - def async_chained_decode_mlx( - self, - prev_pending: MlxPendingDecode, - ) -> tuple[mx.array, list, list, MlxPendingDecode, str]: + def async_chained_decode_mlx(self, prev_pending: MlxPendingDecode) -> MlxLaunch: """Launch a decode step that chains off a still-lazy previous decode. This is the "no idle gap" pipelining primitive: build the next @@ -455,22 +585,21 @@ class MlxTpModelWorker(TpModelWorker): * ``prev_pending`` should be finalised BEFORE the returned pending, so per-request token lists are appended in order. - Returns a 5-tuple matching - :meth:`async_forward_batch_generation_mlx` for the decode case: - ``(lazy_tokens, [], [], pending_decode, "decode")``. The empty - prefill/extend lists are always absent for chained decodes. + Returns an :class:`MlxLaunch` in ``"decode"`` mode; its prefill + and extend lists are always empty for a chained decode. """ pending = self._mlx_runner.decode_batch_start_chained(prev_pending) mx.async_eval(pending.lazy_tokens) - return pending.lazy_tokens, [], [], pending, "decode" + return MlxLaunch( + lazy_tokens=pending.lazy_tokens, + prefills=[], + extends=[], + decode=pending, + mode="decode", + ) def finalize_mlx_result( - self, - prefills: list[MlxPendingPrefill], - extends: list[MlxPendingExtend], - decode: Optional[MlxPendingDecode], - mode: str, - reqs: list, + self, launch: MlxLaunch, reqs: list ) -> GenerationBatchResult: """Materialise a lazy MLX result into a :class:`GenerationBatchResult`. @@ -480,28 +609,40 @@ class MlxTpModelWorker(TpModelWorker): """ from sglang.srt.layers.logits_processor import LogitsProcessorOutput - if mode == "idle": + decode = launch.decode + if launch.mode == "idle": return GenerationBatchResult( logits_output=LogitsProcessorOutput(next_token_logits=None), can_run_cuda_graph=False, ) - if mode == "decode": + step_logprob_rows: dict[str, tuple] = {} + + if launch.mode == "decode": assert decode is not None next_tokens_list = self._mlx_runner.decode_batch_finalize(decode) + self._collect_step_logprobs( + step_logprob_rows, decode.lazy_logprobs, decode.req_ids + ) - elif mode == "extend": + elif launch.mode == "extend": prefill_map: dict[str, int] = {} - for pending_p in prefills: + for pending_p in launch.prefills: prefill_map[pending_p.req_id] = self._mlx_runner.prefill_finalize( pending_p ) + self._collect_step_logprobs( + step_logprob_rows, pending_p.lazy_logprobs, [pending_p.req_id] + ) extend_map: dict[str, int] = {} - for pending_e in extends: + for pending_e in launch.extends: extend_map[pending_e.req_id] = self._mlx_runner.extend_finalize( pending_e ) + self._collect_step_logprobs( + step_logprob_rows, pending_e.lazy_logprobs, [pending_e.req_id] + ) decode_map: dict[str, int] = {} if decode is not None: @@ -509,6 +650,9 @@ class MlxTpModelWorker(TpModelWorker): decode_map = { rid: tok for rid, tok in zip(decode.req_ids, mixed_tokens) } + self._collect_step_logprobs( + step_logprob_rows, decode.lazy_logprobs, decode.req_ids + ) next_tokens_list = [] for req in reqs: @@ -520,11 +664,16 @@ class MlxTpModelWorker(TpModelWorker): next_tokens_list.append(prefill_map[req.rid]) else: - raise ValueError(f"Unknown MLX async mode: {mode}") + raise ValueError(f"Unknown MLX async mode: {launch.mode}") next_token_ids = torch.tensor(next_tokens_list, dtype=torch.long, device="cpu") + logits_output = ( + self._assemble_logprob_output(step_logprob_rows, reqs) + if step_logprob_rows + else LogitsProcessorOutput(next_token_logits=None) + ) return GenerationBatchResult( - logits_output=LogitsProcessorOutput(next_token_logits=None), + logits_output=logits_output, next_token_ids=next_token_ids, can_run_cuda_graph=False, ) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 3119c7728..3a9de1444 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -276,7 +276,12 @@ from sglang.srt.observability.trace import process_tracing_init, trace_set_threa from sglang.srt.parser.reasoning_parser import ReasoningParser from sglang.srt.platforms import current_platform from sglang.srt.plugins import load_plugins -from sglang.srt.runtime_context import get_context, get_parallel, publish +from sglang.srt.runtime_context import ( + get_context, + get_device, + get_parallel, + publish, +) from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_params import TOP_K_ALL from sglang.srt.server_args import PortArgs, ServerArgs @@ -2605,6 +2610,24 @@ class Scheduler( self._add_request_to_queue(req) return + if ( + get_device().mlx_enable_sampling + and req.return_logprob + and 0 <= req.logprob_start_len < len(req.origin_input_ids) + ): + # The MLX sampling path computes output logprobs only; the + # prefill result carries no input_token_logprobs, so letting + # this through would crash output processing. + error_msg = ( + "Prompt input logprobs (logprob_start_len) are not supported " + "on the MLX sampling path; omit logprob_start_len to get " + "output logprobs." + ) + req.logprob_start_len = -1 + req.set_finish_with_abort(error_msg) + self._add_request_to_queue(req) + return + if recv_req.return_routed_experts: error_msg = None if recv_req.routed_experts_start_len < 0: diff --git a/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py b/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py index 379501f0e..8409c4a21 100644 --- a/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/logprob_result_processor.py @@ -327,7 +327,13 @@ class SchedulerLogprobResultProcessor: else: self._initialize_empty_logprob_containers(req) - if req.logprob.top_logprobs_num > 0: + if ( + req.logprob.top_logprobs_num > 0 + and output.next_token_top_logprobs_val is not None + ): + # Guarded like next_token_logprobs above: a backend may leave the + # top-logprob fields unset even for a request that asked for them + # (indexing None raises TypeError). req.logprob.output_top_logprobs_val.append( output.next_token_top_logprobs_val[i] ) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 0086ef228..c1759dcd9 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1198,6 +1198,25 @@ class ServerArgs: NS("device"), ] = 1 random_seed: A[Optional[int], "The random seed.", NS("device")] = None + mlx_enable_sampling: A[ + bool, + ( + "MLX backend only: sample decode tokens (temperature / top-k / " + "top-p / min-p) instead of greedy argmax. Sampling runs inside " + "the lazy MLX graph, so it works with the overlap scheduler; " + "first tokens from prefill/extend are sampled too. Greedy " + "requests keep exact argmax behavior. Also enables on the MLX " + "path: grammar vocab masks and custom logit processors (these " + "break decode chaining per step; custom processors run on " + "pure-decode steps only), logit_bias, output logprobs (sampled " + "token / top-k / token_ids; prompt input logprobs are not " + "computed), NaN sanitization (SGLANG_SANITIZE_NAN_LOGITS), and " + "per-request sampling_seed under " + "--enable-deterministic-inference (deterministic within MLX " + "only). Penalties are not applied." + ), + NS("device"), + ] = False watchdog_timeout: A[ float, "Set watchdog timeout in seconds. If a forward batch takes longer than this, the server will crash to prevent hanging.", @@ -5359,9 +5378,13 @@ class ServerArgs: elif model_arch in ["GptOssForCausalLM"]: # Attention backend selection + XPU dtype validation moved to the # override registry (arg_groups/overrides.py: _gpt_oss_overrides). - # None of these backends exist on MPS; attention_backend is still - # unset there at this point (the torch_native default fills later). - if not is_mps(): + # None of these backends exist on MPS, and under MLX attention + # runs inside the MLX runner, so attention_backend is still unset + # at this point (the torch_native default fills later). macOS + # *without* MLX is not exempt: it has no runner of its own, so it + # must still be held to the supported-backend list instead of + # silently landing on torch_native (no SWA, no sinks). + if not (is_mps() and use_mlx()): supported_backends = [ "triton", "trtllm_mha", diff --git a/test/registered/mlx/models_e2e/test_gpt_oss_mlx_correctness.py b/test/registered/mlx/models_e2e/test_gpt_oss_mlx_correctness.py index a03dfeea1..e7a3cbe27 100644 --- a/test/registered/mlx/models_e2e/test_gpt_oss_mlx_correctness.py +++ b/test/registered/mlx/models_e2e/test_gpt_oss_mlx_correctness.py @@ -5,8 +5,9 @@ uses per-head attention sinks, so it exercises the MLX backend's sliding-window path end to end. Two guards: 1. ``TestGptOssMlxCorrectness`` — black-box serving smoke against a running - server, including a >128-token prompt so the sliding window actually - engages. + server with the radix cache enabled (the default KV path), including a + >128-token prompt so the sliding window actually engages and a repeated + prompt so a radix prefix hit must reproduce the cold greedy output. 2. ``TestGptOssMlxReferenceCorrectness`` — token-for-token equivalence of ``MlxModelRunner`` greedy decoding against raw, unpatched mlx_lm greedy generation. SGLang keeps full KV and applies banded masks / @@ -117,10 +118,14 @@ class TestGptOssMlxCorrectness(CustomTestCase): cls.base_url, timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, other_args=[ + # Radix cache stays enabled (the default): sliding-window + # layers keep windowed per-request KV, the shared pool holds + # full-attention layers, and prefix hits recompute the + # prefix, so serving must stay correct without + # --disable-radix-cache. "--trust-remote-code", "--tp-size", "1", - "--disable-radix-cache", "--disable-cuda-graph", "--mem-fraction-static", MEM_FRACTION_STATIC, @@ -189,6 +194,24 @@ class TestGptOssMlxCorrectness(CustomTestCase): ) self.assertIn("BLUEBERRY", text.upper()) + def test_radix_prefix_hit_reproduces_greedy_output(self): + # The server runs with the radix cache enabled. Sending the same + # >128-token prompt twice makes the second request hit the cached + # prefix; on sliding-window models the runner recomputes the prefix + # (windowed KV keeps no pool history), and greedy output must be + # identical to the cold request. + messages = [ + {"role": "system", "content": "You are a concise assistant."}, + { + "role": "user", + "content": _NUMBER_LIST + + ". Which number comes right after 41? Answer briefly.", + }, + ] + cold = self._chat(messages, max_tokens=48) + hit = self._chat(messages, max_tokens=48) + self.assertEqual(cold, hit) + @unittest.skipUnless(_HAS_MLX, _SKIP_REASON) class TestGptOssMlxReferenceCorrectness(CustomTestCase): diff --git a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py index 921eeb541..f62f5b414 100644 --- a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py +++ b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py @@ -46,6 +46,7 @@ if _HAS_MLX: MlxPendingJob, SchedulerMlxOverlapMixin, ) + from sglang.srt.hardware_backend.mlx.tp_worker import MlxLaunch from sglang.srt.managers.scheduler_components import ( batch_result_processor as batch_result_processor_module, ) @@ -309,7 +310,7 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase): new_slot_ids=[4], req_pool_idx=0, ) - MlxModelRunner._eval_with_cache(pending.lazy_token, pending.cache) + runner.eval_pending(pending) mx.eval(*runner._attention_kv_pool.all_buffers()) runner.prefill_finalize(pending) @@ -350,7 +351,8 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase): calls.append( (len(caches), batched_input.tolist(), list(helper_req_ids)) ) - return mx.array(list(range(len(caches))), dtype=mx.int32) + # Last-token logits whose argmax is the row index. + return mx.eye(len(caches), 8, dtype=mx.float32) def fail_native(*args, **kwargs): raise AssertionError("dense decode should use batched attention") @@ -386,7 +388,8 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase): def fake_batched(caches, batched_input, helper_req_ids): calls.append((len(caches), batched_input.tolist(), list(helper_req_ids))) - return mx.array([8], dtype=mx.int32) + # Last-token logits whose argmax is token 8. + return mx.arange(9, dtype=mx.float32)[None, :] def fail_native(*args, **kwargs): raise AssertionError("dense chained decode should use batched attention") @@ -502,12 +505,13 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase): ] ] - lazy_tokens = runner._decode_with_batched_attention( + lazy_logits = runner._decode_with_batched_attention( cache, mx.array([[7]], dtype=mx.int32), ["r0"], ) - mx.eval(lazy_tokens, *MlxModelRunner._cache_state_arrays(cache)) + lazy_tokens = mx.argmax(lazy_logits, axis=-1) + mx.eval(lazy_tokens, *MlxModelRunner.cache_state_arrays(cache)) self.assertEqual(lazy_tokens.tolist(), [0]) self.assertEqual(cache[0][0].offset, 1) @@ -558,6 +562,9 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase): req_pool_idx={"r0": 0, "r1": 1}, req_to_token_pool=req_to_token_pool, attention_layer_indices=[0], + # The fused scatter addresses pool buffers by full-attention index, + # so the context requires the map whenever the RoPE kernel is live. + full_kv_pool_index_by_layer={0: 0}, ) self.assertEqual(ctx.seq_lens, [1, 2]) @@ -578,7 +585,8 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase): def fake_hybrid(caches, batched_input, helper_req_ids): calls.append((len(caches), batched_input.tolist(), list(helper_req_ids))) - return mx.array([4, 5], dtype=mx.int32) + # Last-token logits whose argmax is 4 for row 0, 5 for row 1. + return mx.eye(8, dtype=mx.float32)[4:6] def fail_batched(*args, **kwargs): raise AssertionError( @@ -713,7 +721,7 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase): new_slot_ids=[4], req_pool_idx=req.req_pool_idx, ) - MlxModelRunner._eval_with_cache(pending.lazy_token, pending.cache) + runner.eval_pending(pending) runner.prefill_finalize(pending) self.assertEqual(runner.model.seen_inputs, [[[13]]]) @@ -770,7 +778,7 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase): req_pool_idx=req.req_pool_idx, req=req, ) - MlxModelRunner._eval_with_cache(pending.lazy_token, pending.cache) + runner.eval_pending(pending) runner.prefill_finalize(pending) tracked = [FakeNativeCache(), None] runner._req_to_token_pool.auxiliary_state_pool.restore_cache( @@ -832,7 +840,7 @@ class TestMlxAuxiliaryStateRunnerCache(unittest.TestCase): req_pool_idx=req.req_pool_idx, req=req, ) - MlxModelRunner._eval_with_cache(pending.lazy_token, pending.cache) + runner.eval_pending(pending) runner.prefill_finalize(pending) tracked = [FakeNativeCache(), None] runner._req_to_token_pool.auxiliary_state_pool.restore_cache( @@ -1094,11 +1102,13 @@ class TestMlxOverlapScheduler(unittest.TestCase): scheduler.last_batch = stale_batch pending = MlxPendingJob( - lazy_tokens=None, - prefills=["prefill"], - extends=[], - decode=None, - mode="extend", + launch=MlxLaunch( + lazy_tokens=None, + prefills=["prefill"], + extends=[], + decode=None, + mode="extend", + ), batch_copy=batch_copy, schedule_batch=schedule_batch, reqs=[SimpleNamespace(rid="r0")], diff --git a/test/registered/unit/hardware_backend/mlx/test_metal_profiler.py b/test/registered/unit/hardware_backend/mlx/test_metal_profiler.py index bc69ba668..c245a9267 100644 --- a/test/registered/unit/hardware_backend/mlx/test_metal_profiler.py +++ b/test/registered/unit/hardware_backend/mlx/test_metal_profiler.py @@ -203,8 +203,14 @@ class TestSchedulerProfilerManagerMPS(unittest.TestCase): mgr._init_profile(output_dir, None, None, None, None, None, False, "test") return mgr + # MetalCaptureProfiler has two strategies: start_mlx drives + # mx.metal.start_capture, start_mps drives torch.mps.profiler.metal_capture. + # This manager takes the MPS one, so that is the symbol to stand in for -- + # patching mx.metal here leaves the real Metal capture running, which fails + # with "Capture layer is not inserted" unless MTL_CAPTURE_ENABLED=1 is set + # in the environment. def test_start_profile_failure_does_not_crash(self): - import mlx.core as mx + import torch from sglang.srt.hardware_backend.mlx.profiler import ( apply_metal_profiler_patches, @@ -215,20 +221,22 @@ class TestSchedulerProfilerManagerMPS(unittest.TestCase): with tempfile.TemporaryDirectory() as tmp: mgr = self._make_manager(tmp) with patch.object( - mx.metal, - "start_capture", + torch.mps.profiler, + "metal_capture", side_effect=RuntimeError("Capture layer is not inserted"), ): result = mgr._start_profile() self.assertFalse(result.success) + self.assertIn("Capture layer is not inserted", result.message) self.assertFalse(mgr.profile_in_progress) self.assertIsNone(mgr.torch_profiler) def test_start_profile_success_with_mock_capture(self): + from unittest.mock import MagicMock from unittest.mock import patch as mock_patch - import mlx.core as mx + import torch from sglang.srt.hardware_backend.mlx.profiler import ( apply_metal_profiler_patches, @@ -238,14 +246,17 @@ class TestSchedulerProfilerManagerMPS(unittest.TestCase): with tempfile.TemporaryDirectory() as tmp: mgr = self._make_manager(tmp) - with mock_patch.object(mx.metal, "start_capture"), mock_patch.object( - mx.metal, "stop_capture" + capture_ctx = MagicMock() + with mock_patch.object( + torch.mps.profiler, "metal_capture", return_value=capture_ctx ), mock_patch("torch.distributed.barrier"): result = mgr._start_profile() - self.assertTrue(result.success) + self.assertTrue(result.success, result.message) self.assertTrue(mgr.profile_in_progress) + capture_ctx.__enter__.assert_called_once() mgr._stop_profile() self.assertFalse(mgr.profile_in_progress) + capture_ctx.__exit__.assert_called_once() if __name__ == "__main__": diff --git a/test/registered/unit/hardware_backend/mlx/test_mlx_reference_correctness.py b/test/registered/unit/hardware_backend/mlx/test_mlx_reference_correctness.py index 7d2449512..33ba76d89 100644 --- a/test/registered/unit/hardware_backend/mlx/test_mlx_reference_correctness.py +++ b/test/registered/unit/hardware_backend/mlx/test_mlx_reference_correctness.py @@ -202,6 +202,13 @@ class TestMlxReferenceCorrectness(CustomTestCase): self.runner.remove_request(rid) return out + def _truncate_at_eos(self, seq): + """``seq`` up to and including its first EOS (whole seq if none).""" + for i, tok in enumerate(seq): + if tok in self.eos_ids: + return seq[: i + 1] + return list(seq) + def _diff_msg(self, prompt, ref, sgl): horizon = min(len(ref), len(sgl)) first = next((j for j in range(horizon) if ref[j] != sgl[j]), horizon) @@ -246,10 +253,18 @@ class TestMlxReferenceCorrectness(CustomTestCase): for rid in rids: self.runner.remove_request(rid) + # Compare up to and including the first EOS. The horizon is fixed so + # the batch composition never changes mid-run, which walks past EOS on + # short answers -- and there the distribution is near-degenerate, so + # batched and solo argmax can pick different tokens from a numerical + # tie. That is float reduction order (a padded batched SDPA vs an + # unpadded solo one), not state bleed: any cache crossover would show + # up while the model still has an opinion. Measured on this fixture, + # case 1 reaches EOS at index 2 and first differs at index 6. for i, (prompt, _, _) in enumerate(self.cases): - self.assertEqual( - batched[i], solo[i], self._diff_msg(prompt, solo[i], batched[i]) - ) + want = self._truncate_at_eos(solo[i]) + got = self._truncate_at_eos(batched[i]) + self.assertEqual(got, want, self._diff_msg(prompt, want, got)) if __name__ == "__main__": diff --git a/test/registered/unit/hardware_backend/mlx/test_mlx_sampling.py b/test/registered/unit/hardware_backend/mlx/test_mlx_sampling.py new file mode 100644 index 000000000..3cd08b813 --- /dev/null +++ b/test/registered/unit/hardware_backend/mlx/test_mlx_sampling.py @@ -0,0 +1,728 @@ +"""Unit tests for MLX in-graph sampling (hardware_backend/mlx/sampling.py).""" + +from __future__ import annotations + +import importlib.util +import unittest +from collections import Counter + +from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") +register_mlx_ci(est_time=20, suite="stage-a-unit-test-mlx") + +_HAS_MLX = importlib.util.find_spec("mlx") is not None +_SKIP_REASON = "requires mlx" + +if _HAS_MLX: + import mlx.core as mx + + from sglang.srt.hardware_backend.mlx.sampling import ( + DEFAULT_SAMPLING_SEED, + GREEDY_PARAMS, + MAX_BOUNDED_TOP_K, + MlxLogprobSpec, + MlxSamplingParams, + _candidate_width, + _gumbel_noise, + _murmur_hash32, + all_greedy, + compute_logprobs, + sample_tokens, + sanitize_logits, + ) + + +def _reference_murmur3(seed: int, pos: int, col: int) -> int: + """Pure-Python MurmurHash3 mirroring the Triton kernel in + sglang/kernels/ops/sampling/murmur_hash.py: blocks seed_low, + seed_high, position, column; length-16 finalization; fmix32.""" + + def mix(h: int, k: int) -> int: + k = (k * 0xCC9E2D51) & 0xFFFFFFFF + k = ((k << 15) | (k >> 17)) & 0xFFFFFFFF + k = (k * 0x1B873593) & 0xFFFFFFFF + h ^= k + h = ((h << 13) | (h >> 19)) & 0xFFFFFFFF + return (h * 5 + 0xE6546B64) & 0xFFFFFFFF + + seed &= 0xFFFFFFFFFFFFFFFF + h = mix(0, seed & 0xFFFFFFFF) + h = mix(h, (seed >> 32) & 0xFFFFFFFF) + h = mix(h, pos & 0xFFFFFFFF) + h = mix(h, col & 0xFFFFFFFF) + h ^= 16 + h ^= h >> 16 + h = (h * 0x85EBCA6B) & 0xFFFFFFFF + h ^= h >> 13 + h = (h * 0xC2B2AE35) & 0xFFFFFFFF + h ^= h >> 16 + return h + + +def _params(temperature=1.0, top_k=1 << 30, top_p=1.0, min_p=0.0, seed=None): + return MlxSamplingParams( + temperature=temperature, top_k=top_k, top_p=top_p, min_p=min_p, seed=seed + ) + + +@unittest.skipUnless(_HAS_MLX, _SKIP_REASON) +class TestMurmurHashPort(CustomTestCase): + def test_matches_pure_python_reference(self): + """Guards the mx uint32 port of the CUDA murmur kernel: any drift in + wraparound/shift/block-order semantics changes seeded sampling.""" + seeds = [0, 1, 42, 2**31, 2**63 + 12345] + positions = [0, 7, 1023, 2**31 - 1, 5] + vocab = 64 + hashed = _murmur_hash32(seeds=seeds, positions=positions, vocab_size=vocab) + mx.eval(hashed) + for row, (seed, pos) in enumerate(zip(seeds, positions)): + for col in (0, 1, vocab // 2, vocab - 1): + self.assertEqual( + int(hashed[row, col].item()), + _reference_murmur3(seed, pos, col), + msg=f"mismatch at seed={seed} pos={pos} col={col}", + ) + + +@unittest.skipUnless(_HAS_MLX, _SKIP_REASON) +class TestSampleTokens(CustomTestCase): + VOCAB = 32 + + def _logits(self, batch_size: int, key_int: int = 0) -> mx.array: + return ( + mx.random.normal(shape=(batch_size, self.VOCAB), key=mx.random.key(key_int)) + * 3.0 + ) + + def _draw(self, logits, params, positions=None, n=200, key_start=100): + """Sample n times with distinct keys, return per-row token Counters.""" + batch_size = logits.shape[0] + positions = positions if positions is not None else [5] * batch_size + counters = [Counter() for _ in range(batch_size)] + for i in range(n): + toks = sample_tokens( + last_logits=logits, + params=params, + positions=positions, + key=mx.random.key(key_start + i), + ) + mx.eval(toks) + for row, t in enumerate(toks.tolist()): + counters[row][int(t)] += 1 + return counters + + @staticmethod + def _reference_support(probs, top_k, top_p, min_p): + """Independent replica of the mask, in pure Python, on sorted probs.""" + order = sorted(range(len(probs)), key=lambda i: (-probs[i], i)) + keep, cum = [], 0.0 + for rank, idx in enumerate(order): + p = probs[idx] + masked = ( + rank >= min(top_k, len(probs)) + or cum > top_p + or p < probs[order[0]] * min_p + ) + cum += p + if not masked: + keep.append(idx) + return set(keep) + + def _probs(self, logits): + probs = mx.softmax(logits.astype(mx.float32), axis=-1) + mx.eval(probs) + return probs[0].tolist() + + def test_greedy_rows_match_argmax_in_mixed_batch(self): + """A greedy row must return exactly argmax even when other rows in + the batch sample — guards the where() row-select and the sglang + greedy convention (top_k == 1).""" + logits = self._logits(3) + expected = mx.argmax(logits, axis=-1).tolist() + params = [_params(top_k=1), _params(temperature=0.7), _params(top_k=1)] + counters = self._draw(logits, params, n=25) + self.assertEqual(set(counters[0]), {expected[0]}) + self.assertEqual(set(counters[2]), {expected[2]}) + + def test_filter_supports_match_reference(self): + """The sampled support must stay inside the independently computed + mask for each filter and for their combination — guards the rank + mask, the nucleus exclusion, the min_p threshold, and the + sorted->vocab index map.""" + cases = [ + ("top_k=2", dict(top_k=2)), + ("top_p=0.6", dict(top_p=0.6)), + ("min_p=0.3", dict(min_p=0.3)), + ("top_k=8,top_p=0.7,min_p=0.05", dict(top_k=8, top_p=0.7, min_p=0.05)), + ] + for label, kwargs in cases: + with self.subTest(label): + logits = self._logits(1, key_int=3) + support = self._reference_support( + self._probs(logits), + kwargs.get("top_k", 1 << 30), + kwargs.get("top_p", 1.0), + kwargs.get("min_p", 0.0), + ) + drawn = set(self._draw(logits, [_params(**kwargs)], n=400)[0]) + self.assertTrue(drawn <= support, f"{label}: extra {drawn - support}") + self.assertTrue(drawn, f"{label}: nothing sampled") + + def test_seeded_row_is_deterministic_and_key_independent(self): + """A row with sampling_seed must produce the same token regardless + of the RNG key or batch composition — the murmur-gumbel path only + depends on (seed, position, logits).""" + logits = self._logits(2, key_int=5) + seeded = _params(temperature=1.0, seed=1234) + tok_solo = sample_tokens( + last_logits=logits[:1], + params=[seeded], + positions=[9], + key=mx.random.key(0), + ) + tok_other_key = sample_tokens( + last_logits=logits[:1], + params=[seeded], + positions=[9], + key=mx.random.key(999), + ) + tok_in_batch = sample_tokens( + last_logits=logits, + params=[seeded, _params(temperature=0.8)], + positions=[9, 3], + key=mx.random.key(7), + ) + mx.eval(tok_solo, tok_other_key, tok_in_batch) + self.assertEqual(tok_solo.tolist(), tok_other_key.tolist()) + self.assertEqual(int(tok_in_batch[0].item()), int(tok_solo[0].item())) + + def test_seeded_row_unaffected_by_batchmate_filtering(self): + """A seeded row's token must not change when a batchmate triggers + the top-k/top-p sort path — guards the vocab-id-space noise + contract (regression: noise was applied in sorted-rank space when + any row filtered, so batch composition changed seeded tokens).""" + # Near-uniform seeded row: the Gumbel noise decides the token, so + # a change of noise index space is guaranteed to show up. + seeded_logits = ( + mx.random.normal(shape=(1, self.VOCAB), key=mx.random.key(8)) * 0.05 + ) + mate_logits = ( + mx.random.normal(shape=(1, self.VOCAB), key=mx.random.key(9)) * 3.0 + ) + logits = mx.concatenate([seeded_logits, mate_logits], axis=0) + seeded = _params(seed=4321) + solo = sample_tokens( + last_logits=logits[:1], + params=[seeded], + positions=[6], + key=mx.random.key(0), + ) + with_filtering_mate = sample_tokens( + last_logits=logits, + params=[seeded, _params(temperature=1.2, top_k=2)], + positions=[6, 11], + key=mx.random.key(55), + ) + mx.eval(solo, with_filtering_mate) + self.assertEqual(int(with_filtering_mate[0].item()), int(solo[0].item())) + + def test_bounded_top_k_picks_the_same_token_as_the_full_vocab_chain(self): + """The bounded top-K chain is an optimization, not a policy change: + a seeded row must pick the same token whether or not the batch is + eligible for it. A batchmate without a finite top_k pushes the + whole batch back onto the full-vocab chain, so the same seeded row + is sampled both ways here.""" + logits = self._logits(2, key_int=11) + seeded = _params(top_k=4, seed=2024) + bounded = sample_tokens( + last_logits=logits[:1], + params=[seeded], + positions=[6], + key=mx.random.key(0), + ) + # min_p alone leaves top_k at TOP_K_ALL, so this batch falls back. + full_vocab = sample_tokens( + last_logits=logits, + params=[seeded, _params(min_p=0.1)], + positions=[6, 2], + key=mx.random.key(3), + ) + mx.eval(bounded, full_vocab) + self.assertEqual(int(full_vocab[0].item()), int(bounded[0].item())) + + def test_candidate_width_gates_the_bounded_chain(self): + """Only a batch whose widest top_k fits inside both the bound and + the vocabulary may shrink the chain; anything else must return the + full vocab size (which selects the scatter-back path).""" + vocab = 4096 + for label, params, expected in [ + ("under the bound", [_params(top_k=64)], 64), + ("at the bound", [_params(top_k=MAX_BOUNDED_TOP_K)], MAX_BOUNDED_TOP_K), + ("past the bound", [_params(top_k=MAX_BOUNDED_TOP_K + 1)], vocab), + ("no top_k (TOP_K_ALL)", [_params()], vocab), + ("top_k == vocab", [_params(top_k=vocab)], vocab), + ("widest row wins", [_params(top_k=4), _params(top_k=64)], 64), + ("one unbounded row", [_params(top_k=4), _params(top_p=0.9)], vocab), + ]: + with self.subTest(label): + self.assertEqual(_candidate_width(params, vocab), expected) + + def test_seeded_row_varies_with_position(self): + """Positions feed the hash, so a fixed seed must not freeze the + distribution across steps: over many positions the sampled tokens + must not all collapse to one value (vocab of near-uniform probs).""" + logits = mx.zeros((1, self.VOCAB)) # uniform distribution + seeded = [_params(seed=77)] + toks = set() + for pos in range(40): + t = sample_tokens( + last_logits=logits, + params=seeded, + positions=[pos], + key=mx.random.key(0), + ) + mx.eval(t) + toks.add(int(t[0].item())) + self.assertGreater(len(toks), 5, toks) + + def test_seeded_noise_is_finite(self): + """The uniform draw is clamped to [2**-32, 1 - 2**-24] before the + double log: uint32(0xFFFFFFFF) rounds UP to 2**32 in float32, so an + unclamped u can exceed 1 and make log(-log u) NaN — and u == 1 gives + +inf, which would deterministically force that token.""" + # Sanity-check the hazard the clamp exists for. + u_max = mx.array([0xFFFFFFFF], dtype=mx.uint32).astype(mx.float32) / float( + 0xFFFFFFFF + ) + mx.eval(u_max) + self.assertGreaterEqual(float(u_max.item()), 1.0) + self.assertFalse(bool(mx.isfinite(-mx.log(-mx.log(u_max))).item())) + + noise = _gumbel_noise( + params=[_params(seed=1), _params(seed=2**63 - 1)], + positions=[0, 4096], + shape=(2, 1 << 16), + key=mx.random.key(0), + ) + mx.eval(noise) + self.assertTrue(bool(mx.all(mx.isfinite(noise)).item())) + + def test_seed_with_min_p_is_supported(self): + """seed + min_p is well defined here (the pytorch backend asserts on + the combination): Gumbel-max over unnormalized masked weights is + invariant to the missing renormalization, which is exactly the TODO + at layers/sampler.py's multinomial_with_seed path.""" + logits = self._logits(1, key_int=4) + support = self._reference_support(self._probs(logits), 1 << 30, 1.0, 0.3) + tok = sample_tokens( + last_logits=logits, + params=[_params(seed=1234, min_p=0.3)], + positions=[9], + key=mx.random.key(0), + ) + mx.eval(tok) + self.assertIn(int(tok[0].item()), support) + + def test_temperature_sharpens_distribution(self): + """Lower temperature must concentrate mass on the argmax token — + guards the per-row temperature division (e.g. broadcasting bugs + that apply one row's temperature to all rows).""" + logits = self._logits(2, key_int=6) + expected0 = int(mx.argmax(logits[0]).item()) + params = [_params(temperature=0.05), _params(temperature=5.0)] + counters = self._draw(logits, params, n=200) + self.assertGreater(counters[0][expected0] / 200.0, 0.95) + self.assertGreater(len(counters[1]), 5, "high temp should spread mass") + + def test_greedy_helpers(self): + self.assertTrue(all_greedy([GREEDY_PARAMS, _params(top_k=1)])) + self.assertFalse(all_greedy([GREEDY_PARAMS, _params(temperature=0.9)])) + + +@unittest.skipUnless(_HAS_MLX, _SKIP_REASON) +class TestSanitizeAndLogprobs(CustomTestCase): + VOCAB = 16 + + def test_sanitize_matches_nan_to_num_semantics(self): + """Guards the port of sanitize_nan_logits' exact replacement values + (+-1e30, not dtype extremes — temperature division would overflow + dtype extremes back to inf and softmax them to NaN).""" + import struct + + def f32(v): + return struct.unpack("f", struct.pack("f", v))[0] + + x = mx.array([[1.0, float("nan"), float("inf"), -float("inf")]]) + out = sanitize_logits(x) + mx.eval(out) + self.assertEqual(out.tolist(), [[1.0, f32(-1e30), f32(1e30), f32(-1e30)]]) + + def test_logprobs_match_reference_and_row_shapes(self): + """compute_logprobs must equal log_softmax(logits/temp) per row and + cut top-k / token-ids to each row's requested shape — guards the + per-row temperature broadcast and the spec row alignment.""" + import math + + logits = mx.random.normal(shape=(2, self.VOCAB), key=mx.random.key(11)) + params = [_params(temperature=0.5), _params(temperature=2.0)] + tokens = mx.array([3, 7], dtype=mx.uint32) + spec = MlxLogprobSpec(top_ks=(2, 0), token_ids=(None, (1, 4))) + lp = compute_logprobs(logits, params, tokens, spec) + mx.eval(*[a for a in [lp.chosen, lp.top_val, lp.top_idx] if a is not None]) + + raw = logits.tolist() + for row, temp in ((0, 0.5), (1, 2.0)): + scaled = [v / temp for v in raw[row]] + m = max(scaled) + lse = m + math.log(sum(math.exp(v - m) for v in scaled)) + ref = [v - lse for v in scaled] + chosen_token = int(tokens[row].item()) + self.assertAlmostEqual( + float(lp.chosen[row].item()), ref[chosen_token], places=4 + ) + if row == 0: + expect_top = sorted(ref, reverse=True)[:2] + got = lp.top_val[row].tolist()[:2] + for a, b in zip(got, expect_top): + self.assertAlmostEqual(a, b, places=4) + if row == 1: + got = lp.token_ids_val[1].tolist() + self.assertAlmostEqual(got[0], ref[1], places=4) + self.assertAlmostEqual(got[1], ref[4], places=4) + self.assertIsNone(lp.token_ids_val[0]) + + +@unittest.skipUnless(_HAS_MLX, _SKIP_REASON) +class TestRunnerSelectTokens(CustomTestCase): + """_select_tokens_with_logprobs lifecycle on a bare runner (object.__new__).""" + + class _FakeCache: + def __init__(self, offset): + self.offset = offset + + class _FakeLayout: + has_auxiliary_state = False + first_attention_layer_index = 0 + + def _runner(self, enable_sampling): + from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner + + runner = object.__new__(MlxModelRunner) + runner._enable_sampling = enable_sampling + runner._cache_layout = self._FakeLayout() + runner._req_sampling = {} + runner._rng_key = mx.random.key(0) if enable_sampling else None + return runner + + def test_disabled_and_greedy_paths_consume_no_rng(self): + """Flag-off and all-greedy batches must return exact argmax and + leave the RNG key untouched — guards the byte-exact greedy + contract that the e2e temp=0 test relies on.""" + logits = mx.random.normal(shape=(2, 16), key=mx.random.key(1)) + expected = mx.argmax(logits, axis=-1).tolist() + caches = [[self._FakeCache(4)], [self._FakeCache(9)]] + + disabled = self._runner(enable_sampling=False) + toks = disabled._select_tokens_with_logprobs(logits, ["a", "b"], caches)[0] + self.assertEqual(toks.tolist(), expected) + + enabled = self._runner(enable_sampling=True) + enabled._req_sampling = {"a": GREEDY_PARAMS, "b": _params(top_k=1)} + key_before = enabled._rng_key + toks = enabled._select_tokens_with_logprobs(logits, ["a", "b"], caches)[0] + self.assertEqual(toks.tolist(), expected) + self.assertIs(enabled._rng_key, key_before) + + def test_discarded_chunk_without_trunk_stays_greedy_and_consumes_no_rng(self): + """A needs_logits=False chunk on a model without a headless trunk + must not sample: consuming RNG for a discarded token would make + the final output depend on prefill chunking.""" + + def full_model_only(input_ids, cache=None): + return mx.zeros((1, input_ids.shape[1], 16)) + + runner = self._runner(enable_sampling=True) + runner.model = full_model_only # no .model attr -> no trunk + runner._req_sampling = {"a": _params(temperature=1.0)} + key_before = runner._rng_key + tok, lazy_logprobs = runner._forward_lazy_token( + mx.array([[3, 4]], dtype=mx.int32), + [self._FakeCache(2)], + needs_logits=False, + req_id="a", + ) + self.assertIsNone(lazy_logprobs) + mx.eval(tok) + self.assertEqual(tok.tolist(), [0]) # argmax of zeros + self.assertIs(runner._rng_key, key_before) + + def test_logit_edits_gate_greedy_and_sampled_and_logprobs(self): + """An additive -inf edit row must exclude a token from greedy argmax, + from sampling, AND from the reported logprob distribution — guards + the edits-before-selection ordering (a regression that samples raw + logits would pass every other test on near-uniform inputs).""" + runner = self._runner(enable_sampling=True) + runner._req_sampling = {"g": GREEDY_PARAMS, "s": _params(temperature=1.0)} + caches = [[self._FakeCache(4)], [self._FakeCache(4)]] + logits = mx.zeros((2, 8)) + logits = mx.put_along_axis( + logits, + mx.array([[7], [7]], dtype=mx.uint32), + mx.array([[5.0], [5.0]]), + axis=-1, + ) # token 7 dominates both rows + edits = mx.zeros((2, 8)) + edits = mx.put_along_axis( + edits, + mx.array([[7], [7]], dtype=mx.uint32), + mx.array([[-float("inf")], [-float("inf")]]), + axis=-1, + ) # ...but is masked out for both + spec = MlxLogprobSpec(top_ks=(1, 1), token_ids=(None, None)) + for _ in range(10): + tokens, lp = runner._select_tokens_with_logprobs( + logits, ["g", "s"], caches, edits, spec + ) + mx.eval(tokens, lp.chosen, lp.top_val) + self.assertNotIn(7, tokens.tolist()) + self.assertNotIn(7, [row[0] for row in lp.top_idx.tolist()]) + + def test_seed_is_gated_on_deterministic_inference(self): + """Upstream seed contract: SamplingBatchInfo only populates + sampling_seed under --enable-deterministic-inference, and then seeds + every row (default 42). A per-request seed outside that flag is + ignored by every other backend, so it is ignored here too.""" + from types import SimpleNamespace + + def make_req(sampling_seed): + return SimpleNamespace( + sampling_params=SimpleNamespace( + temperature=0.8, + top_k=1 << 30, + top_p=1.0, + min_p=0.0, + sampling_seed=sampling_seed, + frequency_penalty=0.0, + presence_penalty=0.0, + repetition_penalty=1.0, + ) + ) + + self.assertIsNone(MlxSamplingParams.from_req(make_req(7)).seed) + self.assertIsNone( + MlxSamplingParams.from_req(make_req(None), deterministic_seeding=False).seed + ) + self.assertEqual( + MlxSamplingParams.from_req(make_req(7), deterministic_seeding=True).seed, 7 + ) + self.assertEqual( + MlxSamplingParams.from_req(make_req(None), deterministic_seeding=True).seed, + DEFAULT_SAMPLING_SEED, + ) + + def test_chained_decode_keeps_logit_bias(self): + """A chained decode step must keep applying the batch's static + logit_bias rows — regression: the chained path passed edits=None, + silently dropping the bias after the first (fresh) step.""" + runner = self._runner(enable_sampling=True) + runner._req_sampling = {"a": GREEDY_PARAMS} + runner._req_caches = {"a": [self._FakeCache(3)]} + runner._req_token_ids = {"a": [1]} + # token 2 dominates; the edit row bans it -> argmax must fall to 1 + logits = mx.array([[0.0, 3.0, 5.0, 0.0]]) + runner._decode_with_batched_attention = lambda caches, x, rids: logits + edits = mx.array([[0.0, 0.0, -float("inf"), 0.0]]) + + fresh = runner.decode_batch_start(["a"], edit_rows=edits) + chained = runner.decode_batch_start_chained(fresh) + mx.eval(fresh.lazy_tokens, chained.lazy_tokens) + self.assertEqual(fresh.lazy_tokens.tolist(), [1]) + self.assertEqual(chained.lazy_tokens.tolist(), [1]) + + def test_logits_hook_bridge_roundtrip(self): + """The custom-logit-processor hook must see materialized float32 + logits and its in-place edits must re-enter the graph — guards the + mx->numpy->mx bridge (a copy-semantics change would drop edits).""" + runner = self._runner(enable_sampling=True) + logits = mx.zeros((1, 8), dtype=mx.bfloat16) + + def hook(arr): + assert arr.dtype.name == "float32" + arr[0, 5] = 99.0 + return arr + + edited = runner._run_logits_hook(logits, hook) + mx.eval(edited) + self.assertEqual(int(mx.argmax(edited, axis=-1)[0].item()), 5) + + def test_sampling_path_advances_rng_key(self): + """Consecutive sampling builds must consume distinct keys, or every + chained decode step would draw identical noise.""" + logits = mx.zeros((1, 16)) + runner = self._runner(enable_sampling=True) + runner._req_sampling = {"a": _params(temperature=1.0)} + caches = [[self._FakeCache(3)]] + toks = set() + for _ in range(20): + t = runner._select_tokens_with_logprobs(logits, ["a"], caches)[0] + mx.eval(t) + toks.add(int(t[0].item())) + self.assertGreater(len(toks), 3, toks) + + +@unittest.skipUnless(_HAS_MLX, _SKIP_REASON) +class TestWorkerSamplingExtras(CustomTestCase): + """Worker-side builders: logit-edit rows, logprob specs, output assembly.""" + + VOCAB = 8 + + @classmethod + def setUpClass(cls): + from sglang.srt.runtime_context import get_context + + # The worker reads --mlx-enable-sampling off the device config bag, + # which fails closed before a publish. + cls._config = get_context().override_server_args(mlx_enable_sampling=True) + cls._config.install() + cls.addClassCleanup(cls._config.restore) + + @staticmethod + def _worker(): + from sglang.srt.hardware_backend.mlx.tp_worker import MlxTpModelWorker + + return MlxTpModelWorker.__new__(MlxTpModelWorker) + + def _batch(self, sinfo, n=2, return_logprob=False, has_grammar=False): + from types import SimpleNamespace + + return SimpleNamespace( + reqs=[ + SimpleNamespace( + rid=f"r{i}", return_logprob=return_logprob, grammar=None + ) + for i in range(n) + ], + sampling_info=sinfo, + return_logprob=return_logprob, + top_logprobs_nums=None, + token_ids_logprobs=None, + has_grammar=has_grammar, + ) + + def test_edit_rows_combine_grammar_mask_and_bias(self): + """The grammar mask must be applied through the backend's own + apply_vocab_mask on a zeros base and summed with logit_bias — + guards the backend-agnostic zeros-trick, the combine order, and the + ForwardBatch.init_new grammars-population mirror (regression: the + MLX paths never build a ForwardBatch, so sinfo.grammars stayed None + and live grammar objects produced no mask at all).""" + from types import SimpleNamespace + + import torch + + class FakeGrammar: + def apply_vocab_mask(self, logits, vocab_mask): + logits[0, 3] = -float("inf") # row 0 forbids token 3 + + sinfo = SimpleNamespace( + grammars=None, # not yet populated, as on the real MLX path + logit_bias=torch.zeros(2, self.VOCAB).index_put_( + (torch.tensor([1]), torch.tensor([5])), torch.tensor([2.5]) + ), + vocab_size=self.VOCAB, + grammar_mask=None, + ) + + def update_mask(): + sinfo.grammar_mask = SimpleNamespace( + grammar=FakeGrammar(), vocab_mask=torch.zeros(2, 1) + ) + + sinfo.update_regex_vocab_mask = update_mask + + batch = self._batch(sinfo, has_grammar=True) + batch.reqs[0].grammar = object() + rows = self._worker()._build_logit_edit_rows(batch) + self.assertEqual( + [g is not None for g in sinfo.grammars], + [True, False], + "worker must mirror ForwardBatch.init_new's grammars population", + ) + mx.eval(rows["r0"], rows["r1"]) + self.assertEqual(rows["r0"].tolist()[3], -float("inf")) + self.assertEqual(rows["r1"].tolist()[5], 2.5) + self.assertIsNone(sinfo.grammar_mask, "mask must be released after use") + + def test_edit_rows_none_when_nothing_to_edit(self): + from types import SimpleNamespace + + sinfo = SimpleNamespace(grammars=None, logit_bias=None, vocab_size=self.VOCAB) + self.assertIsNone(self._worker()._build_logit_edit_rows(self._batch(sinfo))) + + def test_logprob_spec_subset_alignment(self): + """Spec rows must align to the rid subset order, not batch order — + guards mixed-batch decode sub-batches.""" + from sglang.srt.hardware_backend.mlx.tp_worker import MlxTpModelWorker + + rows = {"a": (3, None), "c": (0, (7, 9))} + spec = MlxTpModelWorker._logprob_spec_for(rows, ["c", "b", "a"]) + self.assertEqual(spec.top_ks, (0, 0, 3)) + self.assertEqual(spec.token_ids, ((7, 9), None, None)) + self.assertIsNone(MlxTpModelWorker._logprob_spec_for(rows, ["x"])) + + @unittest.skipUnless( + importlib.util.find_spec("xgrammar") is not None, "requires xgrammar" + ) + def test_xgrammar_wrapper_supports_cpu_logits(self): + """The MLX edit-row builder feeds CPU logits to the grammar + backend's apply_vocab_mask — regression: the xgrammar wrapper + raised 'Unsupported device: cpu' (its dispatch stopped at + cuda/xpu/musa/npu), so every grammar request crashed the worker.""" + import math + + import numpy as np + import torch + + from sglang.srt.constrained.xgrammar_backend import XGrammarGrammar + + logits = torch.zeros(1, 40) + blocks = math.ceil(40 / 32) + bitmask = torch.full((1, blocks), -1, dtype=torch.int32) + bitmask[0, 0] = int(np.int32(np.uint32(0xFFFFFFFF & ~(1 << 7)))) + XGrammarGrammar.apply_vocab_mask(None, logits, bitmask) + self.assertEqual(logits[0, 7].item(), -float("inf")) + self.assertEqual(logits[0, 6].item(), 0.0) + + def test_assemble_logprob_output_matches_scheduler_contract(self): + """Field shapes must survive the scheduler's move_logprobs_to_cpu + (`.tolist()` on the batch tensor and on every per-row val/idx entry) + and add_logprob_return_values indexing — guards the external + LogitsProcessorOutput consumption contract, including rows without + logprob requests getting empty-but-tolistable fills.""" + from types import SimpleNamespace + + from sglang.srt.hardware_backend.mlx.tp_worker import MlxTpModelWorker + + step_rows = {"a": (-1.5, [-0.1, -0.2], [4, 2], [-3.0], [9])} + reqs = [SimpleNamespace(rid="a"), SimpleNamespace(rid="b")] + out = MlxTpModelWorker._assemble_logprob_output(step_rows, reqs) + + self.assertEqual(out.next_token_logprobs.tolist(), [-1.5, 0.0]) + self.assertEqual( + [v.tolist() for v in out.next_token_top_logprobs_val], + [[-0.10000000149011612, -0.20000000298023224], []], + ) + self.assertEqual( + [v.tolist() for v in out.next_token_top_logprobs_idx], [[4, 2], []] + ) + self.assertEqual( + [v.tolist() for v in out.next_token_token_ids_logprobs_val], + [[-3.0], []], + ) + self.assertEqual(out.next_token_token_ids_logprobs_idx, [[9], []]) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/hardware_backend/mlx/test_scheduler_mixin.py b/test/registered/unit/hardware_backend/mlx/test_scheduler_mixin.py index 4fbbd5912..fafcbb751 100644 --- a/test/registered/unit/hardware_backend/mlx/test_scheduler_mixin.py +++ b/test/registered/unit/hardware_backend/mlx/test_scheduler_mixin.py @@ -129,6 +129,7 @@ class TestOverlapLoopStampsLaunchTs(unittest.TestCase): from sglang.srt.hardware_backend.mlx.scheduler_mixin import ( SchedulerMlxOverlapMixin, ) + from sglang.srt.hardware_backend.mlx.tp_worker import MlxLaunch scheduler = self._make_scheduler(recv_side_effect=[[], _StopLoop()]) @@ -148,7 +149,13 @@ class TestOverlapLoopStampsLaunchTs(unittest.TestCase): scheduler.tp_worker.async_forward_batch_generation_mlx.side_effect = ( lambda _batch: ( events.append("forward"), - (None, [], [], None, "extend"), + MlxLaunch( + lazy_tokens=None, + prefills=[], + extends=[], + decode=None, + mode="extend", + ), )[1] ) @@ -173,6 +180,7 @@ class TestOverlapLoopStampsLaunchTs(unittest.TestCase): from sglang.srt.hardware_backend.mlx.scheduler_mixin import ( SchedulerMlxOverlapMixin, ) + from sglang.srt.hardware_backend.mlx.tp_worker import MlxLaunch # Iteration 1: fresh decode launch. Iteration 2: chain a second # decode on top of it. Iteration 3: stop. @@ -193,16 +201,22 @@ class TestOverlapLoopStampsLaunchTs(unittest.TestCase): scheduler.get_next_batch_to_run.return_value = plan pending_decode = MagicMock() - scheduler.tp_worker.async_forward_batch_generation_mlx.return_value = ( - MagicMock(), - [], - [], - pending_decode, - "decode", + scheduler.tp_worker.async_forward_batch_generation_mlx.return_value = MlxLaunch( + lazy_tokens=MagicMock(), + prefills=[], + extends=[], + decode=pending_decode, + mode="decode", ) scheduler.tp_worker.async_chained_decode_mlx.side_effect = lambda _decode: ( events.append("chained_forward"), - (MagicMock(), [], [], MagicMock(), "decode"), + MlxLaunch( + lazy_tokens=MagicMock(), + prefills=[], + extends=[], + decode=MagicMock(), + mode="decode", + ), )[1] launch_times = iter((1.0, 2.0)) diff --git a/test/registered/unit/hardware_backend/mlx/test_sliding_window_attention.py b/test/registered/unit/hardware_backend/mlx/test_sliding_window_attention.py index 36e05bac3..7a18427b8 100644 --- a/test/registered/unit/hardware_backend/mlx/test_sliding_window_attention.py +++ b/test/registered/unit/hardware_backend/mlx/test_sliding_window_attention.py @@ -7,10 +7,13 @@ pin the three seams that make such models work on the MLX backend: 1. The attention contract accepts ``sm_scale`` and exposes per-layer window sizes read from the mlx-lm container convention (``layer_types`` + ``window_size``). -2. The cache shims' ``make_mask`` mirrors mlx_lm's - ``cache.create_attention_mask`` exactly — in particular ``window_size`` +2. The cache shims' ``make_mask`` matches mlx_lm's + ``cache.create_attention_mask`` semantically — in particular ``window_size`` must produce a banded mask (including for N == 1) instead of being silently dropped, or sliding-window layers degrade to full attention. + Where the window provably cannot bind (``offset + N <= window_size``) the + band equals plain causal, and the shims return the cheap form instead, as + mlx_lm's own ``RotatingKVCache.make_mask`` does. 3. ``MLXAttentionWrapper._batched_decode`` applies the window by truncating each request's KV to the trailing window, passes ``sinks`` through, and uses the contract scale helper. @@ -173,7 +176,22 @@ class TestGptOssAttentionContract(CustomTestCase): @unittest.skipUnless(_HAS_MLX, _SKIP_REASON) class TestShimMakeMask(CustomTestCase): - """The shims must return exactly what mlx_lm's own KVCache.make_mask returns.""" + """The shims must be semantically equal to mlx_lm's own KVCache.make_mask. + + Equal *content*, not equal representation: where the window provably + cannot bind the shims return the cheap ``"causal"`` / ``None`` form that + mlx_lm's RotatingKVCache.make_mask also returns, so the comparison + densifies both sides first. + """ + + def _dense(self, mask, N, offset): + """Dense boolean form of any of the three mask representations.""" + if mask is None: + return mx.ones((N, offset + N), dtype=mx.bool_) + if isinstance(mask, str): + self.assertEqual(mask, "causal") + return create_causal_mask(N, offset) + return mask def _shims(self, offset): contig = ContiguousAttentionKVCache( @@ -185,14 +203,13 @@ class TestShimMakeMask(CustomTestCase): ) return (AttentionOffsetCache(offset=offset), contig, pool_backed) - def _assert_same_mask(self, got, ref, msg): - if ref is None or isinstance(ref, str): - self.assertEqual(got, ref, msg) - else: - self.assertTrue( - isinstance(got, mx.array) and mx.array_equal(got, ref).item(), - msg, - ) + def _assert_same_mask(self, got, ref, msg, N, offset): + self.assertTrue( + mx.array_equal( + self._dense(got, N, offset), self._dense(ref, N, offset) + ).item(), + msg, + ) def test_shims_match_mlx_lm_reference(self): cases = [ @@ -213,8 +230,34 @@ class TestShimMakeMask(CustomTestCase): ref, f"{type(shim).__name__} mismatch for N={N} offset={offset} " f"window={window} return_array={return_array}", + N, + offset, ) + def test_non_binding_window_returns_the_cheap_mask(self): + # offset + N <= window: no query can reach past the window, so the + # band equals plain causal and materialising it only costs time (a + # mask array forces sdpa off its fused causal path, ~2x per layer). + self.assertIsNone(make_attention_mask(1, 0, window_size=4)) + self.assertIsNone(make_attention_mask(1, 3, window_size=4)) + self.assertEqual(make_attention_mask(4, 0, window_size=4), "causal") + # ...and one position past the boundary the band is required again. + self.assertIsInstance(make_attention_mask(4, 1, window_size=4), mx.array) + self.assertIsInstance(make_attention_mask(1, 4, window_size=4), mx.array) + + def test_non_binding_window_matches_the_band_it_replaces(self): + # The shortcut is only legal because the two forms are elementwise + # identical; pin that against mlx_lm's own band builder. + for N, offset, window in ((1, 0, 4), (1, 3, 4), (4, 0, 4), (8, 0, 16)): + band = create_causal_mask(N, offset, window_size=window) + cheap = self._dense( + make_attention_mask(N, offset, window_size=window), N, offset + ) + self.assertTrue( + mx.array_equal(band, cheap).item(), + f"N={N} offset={offset} window={window}", + ) + def test_windowed_mask_is_banded_including_self(self): # Query at absolute position 6 with W=4 may attend to keys 3..6 # (j in [i - W + 1, i], the window includes the query itself). diff --git a/test/registered/unit/hardware_backend/mlx/test_swa_radix_pool.py b/test/registered/unit/hardware_backend/mlx/test_swa_radix_pool.py new file mode 100644 index 000000000..2acd4a2a7 --- /dev/null +++ b/test/registered/unit/hardware_backend/mlx/test_swa_radix_pool.py @@ -0,0 +1,380 @@ +"""Unit tests for sliding-window layers on the MLX radix/pool KV path. + +The shared ``MlxAttentionKVPool`` stores full-attention layers only, +sliding-window layers keep window-bounded per-request storage, and a +radix prefix hit on an SWA model recomputes the whole prefix. Scheduler +bookkeeping stays in the unclamped coordinates, so these tests drive +``MlxModelRunner`` directly with hand-built slot ids, mirroring how the +tp_worker calls it. +""" + +from __future__ import annotations + +import importlib.util +import unittest + +from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=10, suite="base-a-test-cpu") +register_mlx_ci(est_time=10, suite="stage-a-unit-test-mlx") + +_HAS_MLX = ( + importlib.util.find_spec("mlx") is not None + and importlib.util.find_spec("mlx_lm") is not None +) +_SKIP_REASON = "requires mlx + mlx_lm" + +if _HAS_MLX: + import mlx.core as mx + from mlx_lm.models import gpt_oss + + from sglang.srt.hardware_backend.mlx.aot import ( + MlxAOTKernelContext, + MlxAOTKernelSet, + MlxAOTRoPEContext, + MlxAOTRoPEKernel, + ) + from sglang.srt.hardware_backend.mlx.kv_cache import ( + BatchedDecodeContext, + ContiguousAttentionKVCache, + MlxAttentionKVPool, + MLXAttentionWrapper, + WindowedAttentionKVCache, + find_attention_layers, + get_layer_window_sizes, + patch_model_attention, + ) + from sglang.srt.hardware_backend.mlx.kv_cache.layout import MlxModelCacheLayout + from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner + +TINY_WINDOW = 8 + + +def _tiny_gpt_oss_model(): + """Randomly initialized 4-layer gpt_oss with alternating sliding/full layers. + + Mirrors test_windowed_kv_cache.py's builder (kept local: the registered + unit-test directory is not an importable package). + """ + args = gpt_oss.ModelArgs( + num_hidden_layers=4, + num_local_experts=8, + num_experts_per_tok=2, + vocab_size=128, + hidden_size=64, + intermediate_size=64, + head_dim=16, + num_attention_heads=4, + num_key_value_heads=2, + sliding_window=TINY_WINDOW, + rope_theta=150000, + rope_scaling={ + "rope_type": "yarn", + "factor": 32.0, + "beta_fast": 32.0, + "beta_slow": 1.0, + "original_max_position_embeddings": 4096, + "truncate": False, + }, + ) + return gpt_oss.Model(args) + + +def _stub_runner(model, disable_radix_cache, pool_size=64): + """Surgically build a runner around an already-loaded tiny model.""" + layers, attrs = find_attention_layers(model) + runner = MlxModelRunner.__new__(MlxModelRunner) + runner.model = model + runner.disable_radix_cache = disable_radix_cache + runner._cache_layout = MlxModelCacheLayout.from_attention_discovery( + layers, attrs, layer_window_sizes=get_layer_window_sizes(model) + ) + runner._max_seq_len = 64 + runner._cache_pool = [] + runner._req_caches = {} + runner._req_token_ids = {} + runner._req_sampling = {} + runner._req_pool_idx = {} + runner._req_synced_offset = {} + runner._req_to_token_pool = None + runner._attention_kv_pool = None + runner._decode_step_ct = 0 + runner._clear_steps = 0 + runner._aot_kernels = MlxAOTKernelSet() + runner._pool_size = pool_size + if not disable_radix_cache: + runner.init_cache_pools(None) + return runner + + +@unittest.skipUnless(_HAS_MLX, _SKIP_REASON) +class TestSwaLayoutAndPoolContract(CustomTestCase): + def _layout(self, with_windows=True): + model = _tiny_gpt_oss_model() + layers, attrs = find_attention_layers(model) + return MlxModelCacheLayout.from_attention_discovery( + layers, attrs, get_layer_window_sizes(model) if with_windows else None + ) + + def test_partition_and_dense_full_pool_index(self): + layout = self._layout() + self.assertEqual(layout.attention_layer_indices, (0, 1, 2, 3)) + self.assertEqual(layout.swa_attention_layer_indices, (0, 2)) + self.assertEqual(layout.full_attention_layer_indices, (1, 3)) + # Dense over full layers only, so it differs from the cache index. + self.assertEqual(layout.full_kv_pool_index_by_layer, {1: 0, 3: 1}) + self.assertEqual(layout.attention_pool_index_by_layer, {0: 0, 1: 1, 2: 2, 3: 3}) + with self.assertRaises(KeyError): + layout.full_kv_pool_index(0) + + # Without a window map the two indices coincide (pre-SWA behavior). + plain = self._layout(with_windows=False) + self.assertFalse(plain.has_sliding_window_layers) + self.assertEqual( + plain.full_kv_pool_index_by_layer, plain.attention_pool_index_by_layer + ) + + def test_sliding_window_model_gets_no_pool(self): + # An SWA prefix hit recomputes the prefix instead of gathering it, so + # the shared pool would have no reader. Allocating it would burn the + # whole auto-sized KV budget on a write-only buffer. + runner = _stub_runner(_tiny_gpt_oss_model(), disable_radix_cache=False) + self.assertTrue(runner._cache_layout.has_sliding_window_layers) + self.assertIsNone(runner._attention_kv_pool) + # The layer-type split still resolves -- it is the seam the shared + # window-aware SWA pool will build on. + self.assertEqual(runner._cache_layout.full_kv_pool_index_by_layer, {1: 0, 3: 1}) + + def test_pool_covers_every_layer_without_sliding_windows(self): + model = _tiny_gpt_oss_model() + layers, attrs = find_attention_layers(model) + runner = _stub_runner(model, disable_radix_cache=True) + runner._cache_layout = MlxModelCacheLayout.from_attention_discovery( + layers, attrs + ) + runner.disable_radix_cache = False + runner.init_cache_pools(None) + self.assertEqual(runner._attention_kv_pool.num_layers, 4) + self.assertEqual(runner._attention_kv_pool.pool_size, 65) + + def test_all_sliding_window_model_gets_no_pool(self): + # An all-SWA model has nothing to pool. Pool construction must skip + # out, and pool sizing must still land on a finite slot count rather + # than dividing by zero bytes per slot. + runner = _stub_runner(_tiny_gpt_oss_model(), disable_radix_cache=False) + layers, attrs = find_attention_layers(runner.model) + runner._cache_layout = MlxModelCacheLayout.from_attention_discovery( + layers, attrs, {idx: TINY_WINDOW for idx in range(4)} + ) + self.assertEqual(runner._cache_layout.full_attention_layer_indices, ()) + self.assertEqual(runner._cache_layout.full_kv_pool_index_by_layer, {}) + + runner._attention_kv_pool = None + runner.init_cache_pools(None) + self.assertIsNone(runner._attention_kv_pool) + + runner._mem_fraction_static = 0.5 + self.assertGreater(runner._compute_pool_size(None), 0) + + def test_sliding_flag_without_window_map_still_rejected(self): + model = _tiny_gpt_oss_model() + patch_model_attention(model) + model.model.layers[0].self_attn._inner.is_sliding = True + runner = _stub_runner(model, disable_radix_cache=True) + # With the container window map the flagged layer is bounded: fine. + runner._get_attn_config() + # Without a resolvable window the layer cannot be bounded: reject. + layers, attrs = find_attention_layers(model) + runner._cache_layout = MlxModelCacheLayout.from_attention_discovery( + layers, attrs + ) + with self.assertRaises(NotImplementedError): + runner._get_attn_config() + + def test_sync_writes_full_layers_only(self): + runner = _stub_runner(_tiny_gpt_oss_model(), disable_radix_cache=False) + # init_cache_pools skips the pool on an SWA model (see + # test_sliding_window_model_gets_no_pool), so attach one by hand: the + # layer-type filtering in _sync_new_kv_to_pool is what the shared + # window-aware SWA pool will rely on, and it must stay correct. + self.assertIsNone(runner._attention_kv_pool) + runner._attention_kv_pool = MlxAttentionKVPool( + pool_size=runner._pool_size + 1, + num_layers=runner._cache_layout.num_full_attention_layers, + n_kv_heads=2, + head_dim=16, + dtype=mx.float32, + ) + cache = runner._new_native_cache() + per_layer_k = {} + for layer_idx in range(4): + k = mx.full((1, 2, 5, 16), float(layer_idx + 1)) + cache[layer_idx].update_and_fetch(k, -k) + per_layer_k[layer_idx] = k + slot_ids = [7, 9, 11] + runner._sync_new_kv_to_pool(cache, cache_start=2, slot_ids=slot_ids) + for layer_idx, pool_idx in ((1, 0), (3, 1)): + got_k, got_v = runner._attention_kv_pool.get_kv( + pool_idx, mx.array(slot_ids, dtype=mx.int32) + ) + want = per_layer_k[layer_idx][0, :, 2:5, :].transpose(1, 0, 2) + self.assertTrue(mx.array_equal(got_k, want).item()) + self.assertTrue(mx.array_equal(got_v, -want).item()) + # Untouched pool slots stay zero (nothing wrote outside the slots). + rest_k, _ = runner._attention_kv_pool.get_kv( + 0, mx.array([1, 2, 3], dtype=mx.int32) + ) + self.assertEqual(mx.abs(rest_k).max().item(), 0.0) + + +@unittest.skipUnless(_HAS_MLX, _SKIP_REASON) +class TestSwaRadixPath(CustomTestCase): + """Radix-path prefill/decode is token-identical to the per-request path. + + Both runners share one tiny gpt-oss (same random weights). The + reference runner uses the ``disable_radix_cache`` per-request path + pinned by test_windowed_kv_cache.py; the radix runner replays the + same requests through pool sync, prefix hits, and prefix recomputes. + """ + + DECODE_STEPS = 6 + + @classmethod + def setUpClass(cls): + mx.random.seed(7) + cls.model = _tiny_gpt_oss_model() + patch_model_attention(cls.model) + + def setUp(self): + self.reference = _stub_runner(self.model, disable_radix_cache=True) + self.radix = _stub_runner(self.model, disable_radix_cache=False) + + def _greedy(self, runner, rid, full_ids, new_ids, prefix_slots, new_slots): + tokens = [ + runner.prefill( + req_id=rid, + new_token_ids=list(new_ids), + full_token_ids=list(full_ids), + prefix_slot_ids=list(prefix_slots), + new_slot_ids=list(new_slots), + req_pool_idx=0, + ) + ] + for _ in range(self.DECODE_STEPS): + tokens.extend(runner.decode_batch([rid])) + return tokens + + def _reference_stream(self, prompt): + tokens = self._greedy( + self.reference, "ref", prompt, prompt, prefix_slots=(), new_slots=() + ) + self.reference.remove_request("ref") + return tokens + + def _seed(self, prompt, slots): + self._greedy(self.radix, "seed", prompt, prompt, (), slots) + self.radix.remove_request("seed") + + def _assert_windowed_bounded(self, rid): + cache = self.radix._req_caches[rid] + for layer_idx in self.radix._cache_layout.swa_attention_layer_indices: + entry = cache[layer_idx] + self.assertIsInstance(entry, WindowedAttentionKVCache) + self.assertLessEqual(entry.get_kv()[0].shape[2], 2 * TINY_WINDOW) + + def test_cold_prefill_matches_reference(self): + prompt = [(i * 7 + 3) % 128 for i in range(20)] + want = self._reference_stream(prompt) + got = self._greedy( + self.radix, "cold", prompt, prompt, (), range(1, len(prompt) + 1) + ) + self.assertEqual(got, want) + self._assert_windowed_bounded("cold") + + def test_partial_prefix_hit_recomputes_exactly(self): + # Prefix (20) is well past the window (8): the hit must recompute the + # whole prefix rather than gather it, a chunked extend continues on + # top, and the stream must match one cold reference over the same + # tokens with every cache left at the unclamped absolute position. + prefix = [(i * 7 + 3) % 128 for i in range(20)] + chunk_a, chunk_b = [9, 42, 77, 5], [11, 13, 17] + prefix_slots = list(range(1, len(prefix) + 1)) + self._seed(prefix, prefix_slots) + want = self._reference_stream(prefix + chunk_a + chunk_b) + + gathers = [] + self.radix._cache_with_pool_backed_attention = lambda slots, n: gathers.append( + n + ) + self.radix.prefill( + req_id="hit", + new_token_ids=chunk_a, + full_token_ids=prefix + chunk_a, + prefix_slot_ids=prefix_slots, + new_slot_ids=list(range(30, 34)), + req_pool_idx=0, + ) + self.assertEqual(gathers, [], "SWA prefix hits must recompute, not gather") + + got = [self.radix.extend("hit", chunk_b, list(range(34, 37)))] + for _ in range(self.DECODE_STEPS): + got.extend(self.radix.decode_batch(["hit"])) + self.assertEqual(got, want) + self._assert_windowed_bounded("hit") + expected = len(prefix + chunk_a + chunk_b) + self.DECODE_STEPS + for layer_idx in range(4): + self.assertEqual(self.radix._req_caches["hit"][layer_idx].offset, expected) + + def test_full_prefix_hit_without_new_tokens(self): + # An exact hit leaves no extend tokens: the prefix rebuild supplies + # run tokens ending on the last prefix token, whose logits predict + # the next token. + prompt = [(i * 5 + 11) % 128 for i in range(20)] + prefix_slots = list(range(1, len(prompt) + 1)) + self._seed(prompt, prefix_slots) + want = self._reference_stream(prompt) + got = self._greedy(self.radix, "exact", prompt, [], prefix_slots, ()) + self.assertEqual(got, want) + + def test_fused_aot_kernel_serves_full_layers_by_full_pool_index(self): + # The fused RoPE+pool-scatter kernel must skip sliding-window layers + # and address the pool by the full-attention index, not the cache one. + sliding_wrapper = self.model.model.layers[0].self_attn + full_wrapper = self.model.model.layers[1].self_attn + + recorded = [] + original = MLXAttentionWrapper._rope_custom_aot + + def _recording_rope(queries, keys, values, positions, pool_idx, rope_ctx): + recorded.append(pool_idx) + return queries, keys + + MLXAttentionWrapper._rope_custom_aot = staticmethod(_recording_rope) + try: + win = WindowedAttentionKVCache(TINY_WINDOW) + contig = ContiguousAttentionKVCache( + n_kv_heads=2, head_dim=16, max_seq_len=32, dtype=mx.float32 + ) + ctx = BatchedDecodeContext( + batch_size=1, + seq_lens=[0], + attention_layer_caches=[[win], [contig]], + attention_pool_index_by_layer={0: 0, 1: 1}, + full_kv_pool_index_by_layer={1: 0}, + aot=MlxAOTKernelContext( + rope=MlxAOTRoPEContext(kernel=MlxAOTRoPEKernel(), kv_pool=None) + ), + ) + x = mx.random.normal((1, 1, 64)) + mx.eval(sliding_wrapper._batched_decode(x, ctx)) + self.assertEqual(recorded, [], "SWA layer must not hit the fused kernel") + mx.eval(full_wrapper._batched_decode(x, ctx)) + # Cache index for layer 1 is 1; its full-pool index is 0. + self.assertEqual(recorded, [0], "full layer needs the full-pool index") + finally: + MLXAttentionWrapper._rope_custom_aot = original + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/hardware_backend/mlx/test_tp_worker_routing.py b/test/registered/unit/hardware_backend/mlx/test_tp_worker_routing.py index 379dd8b5d..cd29832d8 100644 --- a/test/registered/unit/hardware_backend/mlx/test_tp_worker_routing.py +++ b/test/registered/unit/hardware_backend/mlx/test_tp_worker_routing.py @@ -15,11 +15,15 @@ correct discriminator is ``batch.decoding_reqs``, not the chunk length. The routing decision was duplicated across the sync and async paths (the bug therefore existed in both). It now lives in the shared -``MlxTpModelWorker._route_extend_request`` helper. These tests cover: +``MlxTpModelWorker._route_extend_request`` helper, and the sync entry point +launches through the async one rather than re-implementing it. These tests +cover: - * the helper decision directly (both paths delegate to it); - * the sync wiring, by driving ``_forward_batch_generation_mlx``; - * the async wiring, by driving ``_async_extend_batch``. + * the helper decision directly; + * the async wiring, by driving ``_async_extend_batch``; + * the sync entry point, by driving ``_forward_batch_generation_mlx`` -- + which also guards the delegation, since a divergence there would show up + as a routing or token-ordering difference between the two. They mock the MLX runner and load no model. Apple-Silicon-only because ``tp_worker`` imports ``mlx.core`` at module load. @@ -35,7 +39,9 @@ from types import SimpleNamespace import torch from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.srt.runtime_context import get_context from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.test_utils import CustomTestCase # CPU marker is AST-parsed "this test exists"; actual CPU-side execution is # gated by the @skipUnless guard below. MLX marker runs for real on the MLX @@ -49,11 +55,15 @@ _SKIP_REASON = "Apple-Silicon-only (tp_worker imports mlx.core at module load)" class _FakeRunner: - """Records which routing path each request took (sync + async surfaces).""" + """Records which routing path each request took (both worker paths + drive the runner through the same start/finalize surface).""" def __init__(self, known_rids): self._known = set(known_rids) self.calls: list[tuple[str, str]] = [] # (op, rid) + # (op, rid) -> needs_logits as received; guards the worker's + # chunk-finality derivation reaching the runner intact. + self.logits_flags: dict[tuple[str, str], bool] = {} self._req_caches: dict[str, list] = {} self._counter = 0 @@ -73,37 +83,27 @@ class _FakeRunner: return SimpleNamespace(state=[mx.array([0.0], dtype=mx.float32)]) - # --- sync surface --- - def extend(self, rid, new_token_ids, new_slot_ids): - self.calls.append(("extend", rid)) - self._counter += 1 - return 1000 + self._counter - - def decode_batch(self, rids): - for rid in rids: - self.calls.append(("decode", rid)) - return [2000 + i for i in range(len(rids))] - - def prefill( + # --- start/finalize surface (shared by the sync and async worker paths) --- + def extend_start( self, req_id, new_token_ids, - full_token_ids, - prefix_slot_ids, new_slot_ids, - req_pool_idx, - req=None, + needs_logits=True, + logit_edit_row=None, + logprob_spec=None, ): - self.calls.append(("prefill", req_id)) - return 3000 - - # --- async surface --- - def extend_start(self, req_id, new_token_ids, new_slot_ids): import mlx.core as mx self.calls.append(("extend_start", req_id)) + self.logits_flags[("extend_start", req_id)] = needs_logits self._req_caches[req_id] = [self._fake_cache_layer()] - return SimpleNamespace(lazy_token=mx.array([0], dtype=mx.int32), req_id=req_id) + return SimpleNamespace( + lazy_token=mx.array([0], dtype=mx.int32), + cache=self._req_caches[req_id], + req_id=req_id, + lazy_logprobs=None, + ) def prefill_start( self, @@ -114,17 +114,24 @@ class _FakeRunner: new_slot_ids, req_pool_idx, req=None, + needs_logits=True, + logit_edit_row=None, + logprob_spec=None, ): import mlx.core as mx self.calls.append(("prefill_start", req_id)) + self.logits_flags[("prefill_start", req_id)] = needs_logits return SimpleNamespace( lazy_token=mx.array([0], dtype=mx.int32), cache=[self._fake_cache_layer()], req_id=req_id, + lazy_logprobs=None, ) - def decode_batch_start(self, rids): + def decode_batch_start( + self, rids, edit_rows=None, logprob_spec=None, logits_hook=None + ): import mlx.core as mx for rid in rids: @@ -133,8 +140,29 @@ class _FakeRunner: lazy_tokens=mx.array([0] * len(rids), dtype=mx.int32), caches=[[self._fake_cache_layer()] for _ in rids], req_ids=list(rids), + lazy_logprobs=None, ) + def prefill_finalize(self, pending): + return 3000 + + def extend_finalize(self, pending): + self._counter += 1 + return 1000 + self._counter + + def decode_batch_finalize(self, pending): + return [2000 + i for i in range(len(pending.req_ids))] + + def collect_logprobs(self, lazy_logprobs): + return None + + def eval_pending(self, pending): + pass + + @staticmethod + def cache_state_arrays(caches): + return [s for cache_list in caches for c in cache_list for s in c.state] + class _FakeReq: def __init__(self, rid, req_pool_idx=0): @@ -142,6 +170,11 @@ class _FakeReq: self.prefix_indices = torch.empty(0, dtype=torch.long) self.fill_ids = [0] self.req_pool_idx = req_pool_idx + # Mirrors Req's chunk-finality contract read by + # MlxTpModelWorker._chunk_needs_logits: extend_range=None means + # "not truncated" (final chunk / plain prefill). + self.extend_range = None + self.full_untruncated_fill_ids = self.fill_ids def get_fill_ids(self): return self.fill_ids @@ -154,15 +187,26 @@ class _FakeBatch: self.reqs = reqs self.extend_lens = list(extend_lens) self.decoding_reqs = decoding_reqs + self.sampling_info = None + self.return_logprob = False # Arbitrary but correctly-sized token / slot arrays. self.input_ids = torch.arange(total, dtype=torch.long) self.out_cache_loc = torch.arange(total, dtype=torch.long) @unittest.skipUnless(_IS_APPLE_SILICON and _HAS_MLX, _SKIP_REASON) -class TestMlxExtendRouting(unittest.TestCase): +class TestMlxExtendRouting(CustomTestCase): """Routing contract for MlxTpModelWorker: shared helper + sync + async.""" + @classmethod + def setUpClass(cls): + # The worker reads --mlx-enable-sampling off the device config bag, + # which fails closed before a publish. Routing itself is orthogonal + # to sampling, so pin it off for the whole case. + cls._config = get_context().override_server_args(mlx_enable_sampling=False) + cls._config.install() + cls.addClassCleanup(cls._config.restore) + @staticmethod def _worker(known_rids): from sglang.srt.hardware_backend.mlx.tp_worker import MlxTpModelWorker @@ -170,6 +214,10 @@ class TestMlxExtendRouting(unittest.TestCase): worker = MlxTpModelWorker.__new__(MlxTpModelWorker) worker._mlx_runner = _FakeRunner(known_rids) worker._mlx_active_rids = set() + # The sync entry point delegates to the async launch, which guards + # pool creation behind this flag; forward_batch_generation has + # already run it for real by the time either path is reached. + worker._mlx_pool_initialized = True return worker # ---------- the shared decision helper ---------- @@ -200,13 +248,28 @@ class TestMlxExtendRouting(unittest.TestCase): def test_sync_one_token_continuation_routes_to_extend(self): """THE REGRESSION (sync): a 1-token continuation must extend, not decode.""" runner = self._run_sync([_FakeReq("r1")], [1], {"r1"}, None, ForwardMode.EXTEND) - self.assertEqual(runner.ops_for("r1"), ["extend"]) + self.assertEqual(runner.ops_for("r1"), ["extend_start"]) + # Untruncated (extend_range None) => final chunk => logits required. + self.assertIs(runner.logits_flags[("extend_start", "r1")], True) + + def test_sync_non_final_chunk_skips_logits(self): + """Head-skip derivation: a scheduler-truncated chunk (extend_range.end + below the request's full untruncated length) reaches the runner with + needs_logits=False; its next-token output is popped as the stale + intermediate token, so computing the vocab head for it is pure waste. + Everything else about routing is unchanged.""" + req = _FakeReq("r1") + req.full_untruncated_fill_ids = list(range(8)) + req.extend_range = SimpleNamespace(start=0, end=4) # 4 < 8: non-final + runner = self._run_sync([req], [4], {"r1"}, None, ForwardMode.EXTEND) + self.assertEqual(runner.ops_for("r1"), ["extend_start"]) + self.assertIs(runner.logits_flags[("extend_start", "r1")], False) def test_sync_genuine_mixed_decode_routes_to_decode(self): p, d = _FakeReq("p1"), _FakeReq("d1") runner = self._run_sync([p, d], [4, 1], {"d1"}, [d], ForwardMode.MIXED) - self.assertEqual(runner.ops_for("p1"), ["prefill"]) - self.assertEqual(runner.ops_for("d1"), ["decode"]) + self.assertEqual(runner.ops_for("p1"), ["prefill_start"]) + self.assertEqual(runner.ops_for("d1"), ["decode_start"]) # ---------- async path: _async_extend_batch ---------- @@ -216,26 +279,34 @@ class TestMlxExtendRouting(unittest.TestCase): worker = MlxTpModelWorker.__new__(MlxTpModelWorker) worker._mlx_runner = _FakeRunner(known_rids) batch = _FakeBatch(forward_mode, reqs, extend_lens, decoding_reqs) - # returns (lazy_stacked, pending_prefills, pending_extends, - # pending_mixed_decode, mode) - result = worker._async_extend_batch(batch) - return worker._mlx_runner, result + launch = worker._async_extend_batch(batch) + return worker._mlx_runner, launch def test_async_one_token_continuation_routes_to_extend(self): """THE REGRESSION (async): a 1-token continuation must extend, not decode.""" - runner, result = self._run_async( + runner, launch = self._run_async( [_FakeReq("r1")], [1], {"r1"}, None, ForwardMode.EXTEND ) self.assertEqual(runner.ops_for("r1"), ["extend_start"]) - self.assertEqual(len(result[2]), 1) # one pending extend - self.assertIsNone(result[3]) # no mixed decode + self.assertIs(runner.logits_flags[("extend_start", "r1")], True) + self.assertEqual(len(launch.extends), 1) # one pending extend + self.assertIsNone(launch.decode) # no mixed decode + + def test_async_non_final_chunk_skips_logits(self): + """Async twin of the head-skip derivation guard.""" + req = _FakeReq("r1") + req.full_untruncated_fill_ids = list(range(8)) + req.extend_range = SimpleNamespace(start=0, end=4) + runner, _ = self._run_async([req], [4], {"r1"}, None, ForwardMode.EXTEND) + self.assertEqual(runner.ops_for("r1"), ["extend_start"]) + self.assertIs(runner.logits_flags[("extend_start", "r1")], False) def test_async_genuine_mixed_decode_routes_to_decode(self): p, d = _FakeReq("p1"), _FakeReq("d1") - runner, result = self._run_async([p, d], [4, 1], {"d1"}, [d], ForwardMode.MIXED) + runner, launch = self._run_async([p, d], [4, 1], {"d1"}, [d], ForwardMode.MIXED) self.assertEqual(runner.ops_for("p1"), ["prefill_start"]) self.assertEqual(runner.ops_for("d1"), ["decode_start"]) - self.assertIsNotNone(result[3]) # pending mixed decode present + self.assertIsNotNone(launch.decode) # pending mixed decode present if __name__ == "__main__": diff --git a/test/registered/unit/hardware_backend/mlx/test_windowed_kv_cache.py b/test/registered/unit/hardware_backend/mlx/test_windowed_kv_cache.py new file mode 100644 index 000000000..60a1f0d87 --- /dev/null +++ b/test/registered/unit/hardware_backend/mlx/test_windowed_kv_cache.py @@ -0,0 +1,321 @@ +"""Unit tests for the MLX windowed per-request attention KV cache. + +``WindowedAttentionKVCache`` keeps only the trailing ``window`` tokens of a +sliding-window layer. Every level is pinned against the full-history path it +replaces: the cache arrays against a ``ContiguousAttentionKVCache`` trailing +slice, the container forward against full-history caches, and +``MLXAttentionWrapper`` batched decode against the same wrapper driven by +full-history caches. + +Sliding-window layers use this storage on both KV paths; how it composes +with the shared pool and radix prefix hits is pinned in +test_swa_radix_pool.py. +""" + +from __future__ import annotations + +import importlib.util +import unittest + +from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=10, suite="base-a-test-cpu") +register_mlx_ci(est_time=10, suite="stage-a-unit-test-mlx") + +_HAS_MLX = ( + importlib.util.find_spec("mlx") is not None + and importlib.util.find_spec("mlx_lm") is not None +) +_SKIP_REASON = "requires mlx + mlx_lm" + +if _HAS_MLX: + import mlx.core as mx + from mlx_lm.models import gpt_oss + from mlx_lm.models.base import create_causal_mask + + from sglang.srt.hardware_backend.mlx.kv_cache import ( + BatchedDecodeContext, + ContiguousAttentionKVCache, + MLXAttentionWrapper, + WindowedAttentionKVCache, + find_attention_layers, + get_layer_window_sizes, + make_attention_mask, + ) + from sglang.srt.hardware_backend.mlx.kv_cache.layout import MlxModelCacheLayout + +WINDOW = 8 +HIDDEN, N_KV_HEADS, HEAD_DIM = 64, 2, 16 + + +def _dense_mask(mask, n_queries: int, offset: int): + """Densify the cheap ``"causal"`` / ``None`` mask forms. + + ``make_attention_mask`` returns those instead of a materialised band + whenever the window cannot bind, so their width lives in the key tensor + rather than in the mask. Densifying keeps width and content checkable + for both forms. + """ + if mask is None or isinstance(mask, str): + return create_causal_mask(n_queries, offset) + return mask + + +def _tiny_gpt_oss_model(): + """Random-weight 4-layer gpt_oss, alternating sliding/full layers.""" + return gpt_oss.Model( + gpt_oss.ModelArgs( + num_hidden_layers=4, + num_local_experts=8, + num_experts_per_tok=2, + vocab_size=128, + hidden_size=HIDDEN, + intermediate_size=64, + head_dim=HEAD_DIM, + num_attention_heads=4, + num_key_value_heads=N_KV_HEADS, + sliding_window=WINDOW, + ) + ) + + +@unittest.skipUnless(_HAS_MLX, _SKIP_REASON) +class TestWindowedCacheEquivalence(CustomTestCase): + """Storage equivalence with the full-history trailing slice. + + Both caches receive identical K/V and only copy it, so the comparisons + are exact (``mx.array_equal``), not tolerance-based. + """ + + W, H, D = 4, 2, 8 + + def _kv(self, S): + return mx.random.normal((1, self.H, S, self.D)) + + def test_chunk_patterns_match_full_trailing_slice(self): + """Windowed storage == trailing slice of full history, everywhere. + + Also pins what the forward pass depends on: the mask built before + ``update_and_fetch`` is exactly as wide as the keys it then returns, + the kept prefix still covers a full window, ``offset`` stays + absolute, and the decode buffer stays bounded by ``2 * window``. + """ + mx.random.seed(0) + for chunks in [ + (3,), # stays inside the window + (4,), # lands exactly on the window + (5,), # first chunk already crosses the window + (6, 3), # second chunk forces prefix normalisation + (2, 2, 2, 2), # repeated small chunks + (1, 1, 1), # degenerate single-token chunks + (10, 1, 7), # chunk larger than 2*window, then mixed + ]: + full = ContiguousAttentionKVCache(max_seq_len=128) + win = WindowedAttentionKVCache(self.W) + for S in chunks: + k, v = self._kv(S), self._kv(S) + mask = win.make_mask(S, window_size=self.W) # runs before update + fk, fv = full.update_and_fetch(k, v) + wk, wv = win.update_and_fetch(k, v) + at = f"chunk S={S} of {chunks}" + self.assertTrue(mx.array_equal(wk, fk[:, :, -wk.shape[2] :, :]), at) + self.assertTrue(mx.array_equal(wv, fv[:, :, -wv.shape[2] :, :]), at) + # When the window cannot bind, make_mask returns the cheap + # "causal"/None form whose width is implicit in the key + # tensor; densify so the invariant stays checkable either way. + dense = _dense_mask(mask, S, wk.shape[2] - S) + self.assertEqual(dense.shape[-1], wk.shape[2], f"mask width, {at}") + self.assertGreaterEqual( + wk.shape[2] - S, min(win.offset - S, self.W), f"prefix, {at}" + ) + self.assertEqual(win.offset, full.offset, at) + + for step in range(4 * self.W): + k, v = self._kv(1), self._kv(1) + full.write_token(k, v) + win.write_token(k, v) + fk, _ = full.get_kv() + wk, _ = win.get_kv() + t = min(win.offset, self.W) + at = f"decode step {step} after {chunks}" + self.assertTrue(mx.array_equal(wk[:, :, -t:, :], fk[:, :, -t:, :]), at) + self.assertEqual(win.offset, full.offset, at) + self.assertLessEqual(win.keys.shape[2], 2 * self.W, at) + + def test_decode_reallocates_amortised_not_per_token(self): + """Compaction must stay amortised O(1) on both write paths.""" + for write in ("write_token", "update_and_fetch"): + win = WindowedAttentionKVCache(self.W) + big = self._kv(5 * self.W) + win.update_and_fetch(big, big) + buf, reallocs = win.keys, 0 + for _ in range(10 * self.W): + getattr(win, write)(self._kv(1), self._kv(1)) + if win.keys is not buf: + buf, reallocs = win.keys, reallocs + 1 + self.assertLessEqual(reallocs, 12, f"{write} reallocated {reallocs}x") + + def test_full_context_mask_raises_once_history_is_unservable(self): + win = WindowedAttentionKVCache(self.W) + win.update_and_fetch(self._kv(3), self._kv(3)) + self.assertEqual(win.make_mask(2), "causal") # nothing dropped yet + # One oversized chunk is enough: the next update normalises the + # prefix to the window, so full context is already unservable. + win.update_and_fetch(self._kv(10), self._kv(10)) + with self.assertRaises(RuntimeError): + win.make_mask(2) + + def test_reset_keeps_buffers_and_replays(self): + win = WindowedAttentionKVCache(self.W) + win.update_and_fetch(self._kv(6), self._kv(6)) + win.reset() + self.assertEqual(win.offset, 0) + self.assertIsNotNone(win.keys) # buffer kept for reuse + k = self._kv(2) + out, _ = win.update_and_fetch(k, k) + self.assertEqual(out.shape[2], 2) # no stale prefix survived + self.assertTrue(mx.array_equal(out, k)) + + +@unittest.skipUnless(_HAS_MLX, _SKIP_REASON) +class TestWindowedModelForward(CustomTestCase): + """Container path: chunked prefill + greedy decode on a tiny gpt-oss.""" + + def test_chunked_prefill_and_greedy_decode_match_full_history(self): + mx.random.seed(0) + model = _tiny_gpt_oss_model() + windows = get_layer_window_sizes(model) + self.assertEqual([windows[i] for i in range(4)], [WINDOW, None, WINDOW, None]) + ids = (mx.arange(20) * 7 + 3) % 128 # 2.5x window + split = 12 # second chunk starts beyond the window + + tokens = {} + for name in ("windowed", "full"): + cache = [ + ( + WindowedAttentionKVCache(windows[i]) + if name == "windowed" and windows[i] is not None + else ContiguousAttentionKVCache(max_seq_len=64) + ) + for i in range(4) + ] + model(ids[None, :split], cache=cache) + out = model(ids[None, split:], cache=cache) + seq = [] + for _ in range(2 * WINDOW): # crosses the compaction boundary + token = mx.argmax(out[:, -1, :], axis=-1) + seq.append(token.item()) + out = model(token[None], cache=cache) + tokens[name] = seq + + self.assertEqual( + tokens["windowed"], + tokens["full"], + "windowed caches diverge from full history", + ) + + +@unittest.skipUnless(_HAS_MLX, _SKIP_REASON) +class TestWindowedBatchedDecode(CustomTestCase): + """The production decode path (``MLXAttentionWrapper._batched_decode``). + + Windowed and full-history caches are driven through the *same* wrapper + with the same inputs. The wrapper slices the trailing window off + whatever ``get_kv`` returns, so both runs must build byte-identical SDPA + inputs and the outputs must be bit-equal, not merely close. + """ + + def test_chained_decode_across_compaction_boundary(self): + """Decode steps built in still-lazy pairs, riding a compaction. + + Pairs mirror ``decode_batch_start_chained``: step N+1's graph is + built before step N materialises. Compaction allocates a fresh + buffer instead of mutating in place, so step N's returned views must + stay valid. Prefill lengths 15 and 6 put the oversized-chunk shrink + in the first pair and a steady-state rebuild in a later one. + """ + mx.random.seed(1) + attn = _tiny_gpt_oss_model().model.layers[0].self_attn + wrapper = MLXAttentionWrapper(attn, layer_idx=0, window_size=WINDOW) + wins, fulls = [], [] + for length in (15, 6): + x = mx.random.normal((1, length, HIDDEN)) + win = WindowedAttentionKVCache(WINDOW) + full = ContiguousAttentionKVCache(max_seq_len=64) + for cache in (win, full): + attn(x, make_attention_mask(length, 0, window_size=WINDOW), cache=cache) + wins.append(win) + fulls.append(full) + + def build_step(x_step, caches): + ctx = BatchedDecodeContext( + batch_size=len(caches), + seq_lens=[c.offset for c in caches], + attention_layer_caches=[caches], + ) + return wrapper._batched_decode(x_step, ctx) + + for pair in range(6): + steps = [mx.random.normal((len(wins), 1, HIDDEN)) for _ in range(2)] + # Both graphs are built before either materialises. + got = [(build_step(x, wins), build_step(x, fulls)) for x in steps] + mx.eval(got) + for tag, (windowed, full) in zip("ab", got): + self.assertTrue( + mx.array_equal(windowed, full), f"pair {pair}{tag} diverges" + ) + for win, full in zip(wins, fulls): + self.assertEqual(win.offset, full.offset) + + +@unittest.skipUnless(_HAS_MLX, _SKIP_REASON) +class TestModelRunnerCacheWiring(CustomTestCase): + """``_new_native_cache``/``_acquire_cache`` wiring without loading weights.""" + + def _stub_runner(self, window_map): + from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner + + layers, attrs = find_attention_layers(_tiny_gpt_oss_model()) + runner = MlxModelRunner.__new__(MlxModelRunner) + runner._cache_layout = MlxModelCacheLayout.from_attention_discovery( + layers, attrs, layer_window_sizes=window_map + ) + runner._max_seq_len = 4096 + runner._cache_pool = [] + return runner + + def test_windowed_only_for_sliding_layers_and_reset_on_reuse(self): + runner = self._stub_runner(get_layer_window_sizes(_tiny_gpt_oss_model())) + cache = runner._new_native_cache() + self.assertEqual( + [type(c) for c in cache], + [ + WindowedAttentionKVCache, + ContiguousAttentionKVCache, + WindowedAttentionKVCache, + ContiguousAttentionKVCache, + ], + ) + self.assertEqual(cache[0].window, WINDOW) + + # Models without container windows have an empty map, so every + # attention layer keeps a contiguous full-history cache. + for c in self._stub_runner({})._new_native_cache(): + self.assertIsInstance(c, ContiguousAttentionKVCache) + + k = mx.random.normal((1, N_KV_HEADS, 10, HEAD_DIM)) + cache[0].update_and_fetch(k, k) + cache[1].update_and_fetch(k, k) + runner._release_cache(cache) + reused = runner._acquire_cache() + self.assertIs(reused, cache) + for c in reused: + self.assertEqual(c.offset, 0) + # A stale local buffer would prepend the previous request's KV. + out, _ = reused[0].update_and_fetch(k, k) + self.assertEqual(out.shape[2], 10) + + +if __name__ == "__main__": + unittest.main()