From aea78d1e73dea4512406c5d16ea6a05907b9652b Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Sun, 9 Aug 2026 23:45:11 -0700 Subject: [PATCH] [misc] Pass FP8 scales in FlashInfer SWA prefill, autotune fp8 on SM120, and tighten `is_image_understandable_model` (#34217) Co-authored-by: Brayden Zhong --- python/sglang/srt/arg_groups/overrides.py | 8 +- python/sglang/srt/configs/model_config.py | 4 +- .../srt/constrained/xgrammar_backend.py | 8 +- .../mlx/kv_cache/attention_kv_cache.py | 33 ++-- .../mlx/kv_cache/attention_wrapper.py | 21 +-- .../srt/hardware_backend/mlx/model_runner.py | 70 +++----- .../srt/hardware_backend/mlx/sampling.py | 161 +++++------------- .../hardware_backend/mlx/scheduler_mixin.py | 5 +- .../srt/hardware_backend/mlx/tp_worker.py | 27 +-- .../layers/attention/flashinfer_backend.py | 3 + .../runner/flashinfer_autotune.py | 7 +- python/sglang/srt/server_args.py | 10 +- 12 files changed, 115 insertions(+), 242 deletions(-) diff --git a/python/sglang/srt/arg_groups/overrides.py b/python/sglang/srt/arg_groups/overrides.py index ca931e80c..5e6f5a579 100644 --- a/python/sglang/srt/arg_groups/overrides.py +++ b/python/sglang/srt/arg_groups/overrides.py @@ -935,11 +935,9 @@ def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict: elif is_hip(): overrides["attention_backend"] = "aiter" 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. + # Exempt MLX only -- it owns attention in its own runner. macOS + # without MLX still falls through to triton and fails fast below, + # rather than landing on torch_native (no sliding window, no sinks). 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/configs/model_config.py b/python/sglang/srt/configs/model_config.py index 8784771ce..3045fad88 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -450,10 +450,12 @@ class ModelConfig: self.hf_config.architectures ) # TODO: requires further polishing + # Key on the tower, not the attribute: several config classes default + # vision_config to None, which presence alone would read as image-capable. self.is_image_understandable_model = ( enable_multimodal and not self.is_lm_only - and hasattr(self.hf_config, "vision_config") + and getattr(self.hf_config, "vision_config", None) is not None ) # Models expose audio_config at different nesting levels: diff --git a/python/sglang/srt/constrained/xgrammar_backend.py b/python/sglang/srt/constrained/xgrammar_backend.py index e4ddcb55c..58b3f1540 100644 --- a/python/sglang/srt/constrained/xgrammar_backend.py +++ b/python/sglang/srt/constrained/xgrammar_backend.py @@ -59,11 +59,9 @@ MAX_ROLLBACK_TOKENS = 200 def _allocate_token_bitmask(vocab_size: int, batch_size: int) -> torch.Tensor: - # 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. + # Pin where pinning exists, so the later H2D can be a genuine non_blocking + # copy (a pageable source silently downgrades it). MPS torch has no + # pin-memory kernel and asserts on pin_memory=True. return torch.full( get_bitmask_shape(batch_size, vocab_size), -1, 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 6f00e46a1..78f371bc1 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 @@ -23,17 +23,13 @@ def make_attention_mask(N, offset, return_array=False, window_size=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. + # Either no window, or a window that cannot bind: the lowest query position + # is ``offset``, so ``offset + N <= window_size`` puts every causally + # visible key inside the band and the banded mask is elementwise identical + # to a plain causal one (the same shortcut mlx_lm's RotatingKVCache takes). + # 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. if N == 1: return None if return_array: @@ -156,8 +152,7 @@ class ContiguousAttentionKVCache: """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. + attend to; slicing it here costs one op instead of two. """ start = 0 if window is None else max(0, self.offset - window) return ( @@ -209,11 +204,10 @@ class WindowedAttentionKVCache: "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 + # No N == 1 shortcut here: 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. + # exactly W keys, while this buffer returns W + 1 once ``kept == + # window`` -- the trailing window plus the token just written. return make_attention_mask( N, kept, return_array=return_array, window_size=window_size ) @@ -221,9 +215,8 @@ class WindowedAttentionKVCache: 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``. + Split out of ``update_and_fetch`` so the decode path can skip building + the two return slices, which it discards in favour of ``get_kv``. """ S = keys.shape[2] kept = min(self._local, self.window) 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 89b134651..689b8fa1c 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 @@ -57,10 +57,9 @@ 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 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. _padding_by_window: dict = field(init=False, default_factory=dict) def __post_init__(self) -> None: @@ -77,10 +76,9 @@ class BatchedDecodeContext: 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. + # The fused scatter addresses pool buffers by full-attention index; + # defaulting to the cache index would write the wrong buffer + # whenever sliding-window layers are interleaved. raise ValueError( "BatchedDecodeContext requires full_kv_pool_index_by_layer " "when the fused AOT RoPE + pool-scatter kernel is active" @@ -99,7 +97,6 @@ class BatchedDecodeContext: 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: @@ -144,10 +141,8 @@ class 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). + # Any attention layer will do: ``offset`` is the ABSOLUTE sequence + # position on every cache, even a windowed one storing far fewer tokens. seq_lens = [ caches[i][attention_layer_indices[0]].offset for i in range(batch_size) ] diff --git a/python/sglang/srt/hardware_backend/mlx/model_runner.py b/python/sglang/srt/hardware_backend/mlx/model_runner.py index d2bb58382..3c91ab47a 100644 --- a/python/sglang/srt/hardware_backend/mlx/model_runner.py +++ b/python/sglang/srt/hardware_backend/mlx/model_runner.py @@ -699,23 +699,10 @@ class MlxModelRunner: 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. + # Allocating it anyway would burn the whole auto-sized KV budget + # (_compute_pool_size fills mem_fraction_static) on a write-only + # buffer. This also disables the fused AOT RoPE + pool-scatter + # kernel, which is opt-in (SGLANG_MLX_USE_CUSTOM_ROPE, default off). logger.info( "Model has %d sliding-window attention layers; skipping the " "shared attention KV pool (an SWA prefix hit recomputes the " @@ -1115,11 +1102,9 @@ class MlxModelRunner: 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). + vocab-sized float32 logits for every chunk position — the largest + transient allocation in the process. Returns None when the model + exposes no headless trunk, and the caller runs the full forward. """ if self._trunk is None: return None @@ -1165,25 +1150,20 @@ class MlxModelRunner: ) -> 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. + Greedy behavior (sampling disabled, or every row greedy with no logit + edits) is exactly the pre-sampling ``mx.argmax``. ``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; 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). + # Shared by sampling and logprobs; None when neither needs it. scaled = ( scale_by_temperature(edited, params) if not greedy or logprob_spec is not None @@ -1238,10 +1218,9 @@ class MlxModelRunner: 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. + Used for custom logit processors (arbitrary torch callables) — the one + edit that cannot be expressed lazily. Synchronizes the graph, so + callers gate it to fresh, pure-decode launches. """ logits32 = last_logits.astype(mx.float32) mx.eval(logits32) @@ -1279,9 +1258,7 @@ class MlxModelRunner: """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. + write-then-read ordering lands in a single kernel submission. """ if isinstance(pending, MlxPendingDecode): tokens, caches = pending.lazy_tokens, pending.caches @@ -1297,11 +1274,10 @@ class MlxModelRunner: 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). + 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`` is what keeps the trunk in the lazy graph handed to + ``mx.eval`` / ``mx.async_eval``. """ return (hidden[:, -1, 0] * 0).astype(mx.int32) diff --git a/python/sglang/srt/hardware_backend/mlx/sampling.py b/python/sglang/srt/hardware_backend/mlx/sampling.py index fd40e2234..92ada6422 100644 --- a/python/sglang/srt/hardware_backend/mlx/sampling.py +++ b/python/sglang/srt/hardware_backend/mlx/sampling.py @@ -2,82 +2,26 @@ 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 +the forward pass. That 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. +sampled tokens as step N+1's input ids, 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 ``top_k_top_p_min_p_sampling_from_probs_torch`` / +``multinomial_with_seed`` in ``sglang/srt/layers/sampler.py``. -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``): +Gaps against that backend: -* ``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. +* Penalties (frequency/presence/repetition) are not applied (warned once + per process). * 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. + match the full ``sampling_info``). +* Logprobs cover the sampled token, top-k, and requested token ids; + prompt/input logprobs (``logprob_start_len``) are not computed. +* 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. """ from __future__ import annotations @@ -220,11 +164,9 @@ def scale_by_temperature( ) -> 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. + Both :func:`sample_tokens` and :func:`compute_logprobs` start here and take + the result as their ``scaled`` argument: MLX does not eliminate common + subexpressions, so a step doing both would otherwise divide twice. """ temps = mx.array([p.temperature for p in params], dtype=mx.float32)[:, None] return last_logits.astype(mx.float32) / temps @@ -243,9 +185,6 @@ def compute_logprobs( ``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) @@ -284,12 +223,9 @@ def sample_tokens( """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`. + are the absolute sequence positions being sampled (seeded rows only). + Assumes at least one row samples; callers shortcut to ``mx.argmax`` when + ``all_greedy(params)``. """ batch_size, vocab_size = last_logits.shape logits32 = last_logits.astype(mx.float32) @@ -339,11 +275,9 @@ def sample_tokens( ) 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. + # Nothing is masked, so log(softmax(scaled)) is just scaled minus a + # per-row constant: feeding the scaled logits straight to the argmax + # below drops two full-vocab passes (softmax, log). log_weights = scaled noise = _gumbel_noise( @@ -353,12 +287,9 @@ def sample_tokens( 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. + # 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 what keeps a seeded row well-defined under min_p. 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) @@ -366,9 +297,9 @@ def sample_tokens( 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. + # Greedy rows still ran the sampled path above (the row exists either way); + # overwriting them with the unnoised argmax is what keeps their token + # independent of the noise. return mx.where(mx.array(greedy), mx.argmax(logits32, axis=-1), sampled) @@ -376,14 +307,12 @@ 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``). + ``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 can run on + ``[B, K]`` instead of ``[B, vocab]`` and still pick the same token. + Falls back to the full vocabulary as soon as one row wants more 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: @@ -416,11 +345,9 @@ def _gumbel_noise( 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. + # uint32(0xFFFFFFFF) rounds UP to 2**32 in float32, so the quotient can land + # just above 1.0 (log(-log(u)) -> NaN); u == 1.0 gives +inf, which would + # deterministically force that token. u = mx.clip(u, 2.0**-32, 1.0 - 2.0**-24) hash_noise = -mx.log(-mx.log(u)) @@ -449,15 +376,11 @@ def _murmur_hash32( ) -> 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. + Blocks mixed in kernel order: seed_low, seed_high, position, column. The + first three are per-row scalars, folded exactly on the CPU with Python + ints; only the column block and finalization run as vectorized uint32 ops. + ``columns`` hashes an explicit [B, K] set of token ids instead of every id + in ``[0, vocab_size)``. """ row_states = [] for seed, pos in zip(seeds, positions): diff --git a/python/sglang/srt/hardware_backend/mlx/scheduler_mixin.py b/python/sglang/srt/hardware_backend/mlx/scheduler_mixin.py index aecf0ba8b..ee2877ee2 100644 --- a/python/sglang/srt/hardware_backend/mlx/scheduler_mixin.py +++ b/python/sglang/srt/hardware_backend/mlx/scheduler_mixin.py @@ -61,10 +61,7 @@ class MlxPendingJob: 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. + # See SchedulerMlxOverlapMixin._mlx_batch_chain_safe. 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 diff --git a/python/sglang/srt/hardware_backend/mlx/tp_worker.py b/python/sglang/srt/hardware_backend/mlx/tp_worker.py index 438936aa7..2766e03f5 100644 --- a/python/sglang/srt/hardware_backend/mlx/tp_worker.py +++ b/python/sglang/srt/hardware_backend/mlx/tp_worker.py @@ -54,19 +54,10 @@ class MlxLaunch: 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"``. + :meth:`MlxTpModelWorker.finalize_mlx_result`. Evaluating ``lazy_tokens`` + materialises the whole batch. ``decode`` covers both full decode mode and + single-token decodes mixed into an extend batch; ``mode`` is one of + ``"idle"``, ``"decode"``, ``"extend"``. """ lazy_tokens: Optional[mx.array] @@ -389,13 +380,9 @@ class MlxTpModelWorker(TpModelWorker): 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. + Launching and finalising back-to-back builds the same lazy graph, so + routing, logit edits, logprob collection and chunk-head skipping keep + one implementation instead of two. """ launch = self.async_forward_batch_generation_mlx(batch) return self.finalize_mlx_result(launch, batch.reqs) diff --git a/python/sglang/srt/layers/attention/flashinfer_backend.py b/python/sglang/srt/layers/attention/flashinfer_backend.py index 324ca3f33..34c14f21f 100644 --- a/python/sglang/srt/layers/attention/flashinfer_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_backend.py @@ -1382,6 +1382,9 @@ class FlashInferAttnBackend(AttentionBackend): sm_scale=layer.scaling, window_left=swa_window_left, logits_soft_cap=logits_soft_cap, + # Must use _float to avoid device-to-host copy that breaks cuda graph capture. + k_scale=layer.k_scale_float, + v_scale=layer.v_scale_float, ) o, _ = _safe_merge_state(o1, s1, o2, s2) diff --git a/python/sglang/srt/model_executor/runner/flashinfer_autotune.py b/python/sglang/srt/model_executor/runner/flashinfer_autotune.py index e656e6c3f..f01256fd4 100644 --- a/python/sglang/srt/model_executor/runner/flashinfer_autotune.py +++ b/python/sglang/srt/model_executor/runner/flashinfer_autotune.py @@ -101,15 +101,18 @@ def should_run_flashinfer_autotune( from sglang.srt.layers.quantization.fp8_utils import ( get_fp8_gemm_runner_backend, ) - from sglang.srt.utils import is_sm100_supported + from sglang.srt.utils import is_sm100_supported, is_sm120_supported model_uses_modelopt_fp8 = model_quantization in ( "modelopt", "modelopt_fp8", "modelopt_mixed", ) + # SM120 satisfies is_blackwell_supported(), so resolve_mxfp8_dense_gemm_backend + # sends it to the same tunable FlashInfer CUTLASS MXFP8 dense GEMM as SM100; + # without this the kernel always runs at tactic=-1. fp8_gemm_needs_autotune = get_fp8_gemm_runner_backend().is_flashinfer_cutlass() or ( - model_uses_modelopt_fp8 and is_sm100_supported() + model_uses_modelopt_fp8 and (is_sm100_supported() or is_sm120_supported()) ) if not (moe_needs_autotune or fp4_gemm_needs_autotune or fp8_gemm_needs_autotune): diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index c1759dcd9..cbef48884 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -5378,12 +5378,10 @@ 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, 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). + # Exempt MLX only: none of these backends exist on MPS, and MLX runs + # attention inside its own runner, so attention_backend is still + # unset here. Plain macOS stays on the list -- torch_native has + # neither sliding window nor attention sinks. if not (is_mps() and use_mlx()): supported_backends = [ "triton",