[MLX] Support gpt-oss: sliding-window attention, attention sinks, sm_scale (#30050)
Co-authored-by: R0CKSTAR <yeahdongcn@gmail.com> Co-authored-by: Alex Nails <alex.nails@radixark.ai>
This commit is contained in:
co-authored by
R0CKSTAR
Alex Nails
parent
68b961e9fb
commit
553dc0f936
@@ -51,6 +51,7 @@ from sglang.srt.utils.common import (
|
|||||||
is_gfx95_supported,
|
is_gfx95_supported,
|
||||||
is_hip,
|
is_hip,
|
||||||
is_mnnvl_fabric_device,
|
is_mnnvl_fabric_device,
|
||||||
|
is_mps,
|
||||||
is_musa,
|
is_musa,
|
||||||
is_npu,
|
is_npu,
|
||||||
is_sm90_supported,
|
is_sm90_supported,
|
||||||
@@ -921,7 +922,8 @@ def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict:
|
|||||||
overrides["attention_backend"] = "intel_xpu"
|
overrides["attention_backend"] = "intel_xpu"
|
||||||
elif is_hip():
|
elif is_hip():
|
||||||
overrides["attention_backend"] = "aiter"
|
overrides["attention_backend"] = "aiter"
|
||||||
else:
|
elif not is_mps():
|
||||||
|
# No triton on macOS; MPS keeps the platform default.
|
||||||
overrides["attention_backend"] = "triton"
|
overrides["attention_backend"] = "triton"
|
||||||
if is_xpu():
|
if is_xpu():
|
||||||
# Check for bf16 dtype on Intel XPU. Reads the pristine dtype request,
|
# Check for bf16 dtype on Intel XPU. Reads the pristine dtype request,
|
||||||
|
|||||||
@@ -139,7 +139,23 @@ def _build_rope_kernel(inputs: MlxAOTKernelBuildInputs) -> MlxAOTRoPEKernel:
|
|||||||
# AOT kernel currently requires rope_dim == head_dim.
|
# AOT kernel currently requires rope_dim == head_dim.
|
||||||
return MlxAOTRoPEKernel()
|
return MlxAOTRoPEKernel()
|
||||||
|
|
||||||
base = float(getattr(rope, "base", 10000.0))
|
# The kernel computes vanilla RoPE from a scalar base. Scaled variants
|
||||||
|
# such as YarnRoPE/Llama3RoPE/SuScaledRoPE expose no ``base`` and bake
|
||||||
|
# their scaling into precomputed ``_freqs`` (plus an ``mscale`` factor
|
||||||
|
# applied outside mx.fast.rope), while linear scaling keeps ``base`` but
|
||||||
|
# sets ``scale != 1`` on nn.RoPE. The kernel has inputs for none of
|
||||||
|
# these, so they must fall back to mx.fast.rope.
|
||||||
|
base = getattr(rope, "base", None)
|
||||||
|
if base is None:
|
||||||
|
return MlxAOTRoPEKernel()
|
||||||
|
if getattr(rope, "_freqs", None) is not None:
|
||||||
|
return MlxAOTRoPEKernel()
|
||||||
|
if float(getattr(rope, "mscale", 1.0)) != 1.0:
|
||||||
|
return MlxAOTRoPEKernel()
|
||||||
|
if float(getattr(rope, "scale", 1.0)) != 1.0:
|
||||||
|
return MlxAOTRoPEKernel()
|
||||||
|
base = float(base)
|
||||||
|
|
||||||
num_qo_heads = get_num_heads(sample_attn)
|
num_qo_heads = get_num_heads(sample_attn)
|
||||||
if num_qo_heads is None:
|
if num_qo_heads is None:
|
||||||
return MlxAOTRoPEKernel()
|
return MlxAOTRoPEKernel()
|
||||||
|
|||||||
@@ -1,7 +1,10 @@
|
|||||||
"""Cache components for the MLX backend."""
|
"""Cache components for the MLX backend."""
|
||||||
|
|
||||||
from sglang.srt.hardware_backend.mlx.kv_cache.attention_contract import (
|
from sglang.srt.hardware_backend.mlx.kv_cache.attention_contract import (
|
||||||
|
get_attention_scale,
|
||||||
|
get_container_window_size,
|
||||||
get_head_dim,
|
get_head_dim,
|
||||||
|
get_layer_window_sizes,
|
||||||
get_num_heads,
|
get_num_heads,
|
||||||
get_num_kv_heads,
|
get_num_kv_heads,
|
||||||
is_attention_module,
|
is_attention_module,
|
||||||
@@ -11,6 +14,7 @@ from sglang.srt.hardware_backend.mlx.kv_cache.attention_kv_cache import (
|
|||||||
AttentionOffsetCache,
|
AttentionOffsetCache,
|
||||||
ContiguousAttentionKVCache,
|
ContiguousAttentionKVCache,
|
||||||
PoolBackedAttentionKVCache,
|
PoolBackedAttentionKVCache,
|
||||||
|
make_attention_mask,
|
||||||
)
|
)
|
||||||
from sglang.srt.hardware_backend.mlx.kv_cache.attention_kv_pool import (
|
from sglang.srt.hardware_backend.mlx.kv_cache.attention_kv_pool import (
|
||||||
MlxAttentionKVPool,
|
MlxAttentionKVPool,
|
||||||
@@ -40,12 +44,16 @@ __all__ = [
|
|||||||
"AttentionOffsetCache",
|
"AttentionOffsetCache",
|
||||||
"ContiguousAttentionKVCache",
|
"ContiguousAttentionKVCache",
|
||||||
"find_attention_layers",
|
"find_attention_layers",
|
||||||
|
"get_attention_scale",
|
||||||
|
"get_container_window_size",
|
||||||
"get_head_dim",
|
"get_head_dim",
|
||||||
"get_context",
|
"get_context",
|
||||||
|
"get_layer_window_sizes",
|
||||||
"get_num_layers",
|
"get_num_layers",
|
||||||
"get_num_heads",
|
"get_num_heads",
|
||||||
"get_num_kv_heads",
|
"get_num_kv_heads",
|
||||||
"is_attention_module",
|
"is_attention_module",
|
||||||
|
"make_attention_mask",
|
||||||
"MLXAttentionWrapper",
|
"MLXAttentionWrapper",
|
||||||
"MlxAttentionKVPool",
|
"MlxAttentionKVPool",
|
||||||
"MlxAuxiliaryStateComponent",
|
"MlxAuxiliaryStateComponent",
|
||||||
|
|||||||
@@ -4,10 +4,12 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from typing import Any, Iterable
|
from typing import Any, Iterable
|
||||||
|
|
||||||
# ``rope`` and ``scale`` are required by MLXAttentionWrapper. Keeping them in
|
# ``rope`` and a softmax scale are required by MLXAttentionWrapper. Keeping
|
||||||
# the contract also prevents recurrent mixers such as DeltaNet from being
|
# them in the contract also prevents recurrent mixers such as DeltaNet from
|
||||||
# mistaken for softmax attention just because they expose projection layers.
|
# being mistaken for softmax attention just because they expose projections.
|
||||||
ATTENTION_API_ATTRS = ("q_proj", "k_proj", "v_proj", "o_proj", "rope", "scale")
|
ATTENTION_API_ATTRS = ("q_proj", "k_proj", "v_proj", "o_proj", "rope")
|
||||||
|
# Any one of these satisfies the scale requirement (gpt_oss uses ``sm_scale``).
|
||||||
|
SCALE_ATTRS = ("scale", "sm_scale")
|
||||||
NUM_HEAD_ATTRS = ("n_heads", "num_heads", "num_attention_heads")
|
NUM_HEAD_ATTRS = ("n_heads", "num_heads", "num_attention_heads")
|
||||||
NUM_KV_HEAD_ATTRS = ("n_kv_heads", "num_k_heads", "num_kv_heads", "num_key_value_heads")
|
NUM_KV_HEAD_ATTRS = ("n_kv_heads", "num_k_heads", "num_kv_heads", "num_key_value_heads")
|
||||||
SLIDING_ATTENTION_ATTRS = (
|
SLIDING_ATTENTION_ATTRS = (
|
||||||
@@ -17,6 +19,9 @@ SLIDING_ATTENTION_ATTRS = (
|
|||||||
"use_sliding_window",
|
"use_sliding_window",
|
||||||
"is_swa",
|
"is_swa",
|
||||||
)
|
)
|
||||||
|
# mlx-lm containers name their scalar sliding window either ``window_size``
|
||||||
|
# (gpt_oss, gemma4) or ``sliding_window`` (olmo3, llama SWA variants, ...).
|
||||||
|
WINDOW_SIZE_ATTRS = ("window_size", "sliding_window")
|
||||||
|
|
||||||
|
|
||||||
def first_present_attr(module: Any, names: Iterable[str]) -> Any | None:
|
def first_present_attr(module: Any, names: Iterable[str]) -> Any | None:
|
||||||
@@ -50,14 +55,47 @@ def get_head_dim(module: Any) -> int | None:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def get_attention_scale(module: Any) -> float | None:
|
||||||
|
return first_present_attr(module, SCALE_ATTRS)
|
||||||
|
|
||||||
|
|
||||||
def is_attention_module(module: Any) -> bool:
|
def is_attention_module(module: Any) -> bool:
|
||||||
return (
|
return (
|
||||||
all(hasattr(module, attr) for attr in ATTENTION_API_ATTRS)
|
all(hasattr(module, attr) for attr in ATTENTION_API_ATTRS)
|
||||||
|
and any(hasattr(module, attr) for attr in SCALE_ATTRS)
|
||||||
and get_num_heads(module) is not None
|
and get_num_heads(module) is not None
|
||||||
and get_num_kv_heads(module) is not None
|
and get_num_kv_heads(module) is not None
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
|
|
||||||
|
def get_layer_window_sizes(model: Any) -> dict[int, int | None]:
|
||||||
|
"""Per-layer sliding-window sizes from the mlx-lm container convention.
|
||||||
|
|
||||||
|
Containers such as gpt_oss or olmo3 expose ``layer_types`` (one entry
|
||||||
|
per layer, ``"sliding_attention"`` marking windowed layers) plus a
|
||||||
|
scalar window (see ``WINDOW_SIZE_ATTRS``). Returns
|
||||||
|
``{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)
|
||||||
|
layer_types = getattr(container, "layer_types", None)
|
||||||
|
window_size = get_container_window_size(model)
|
||||||
|
if not layer_types or window_size is None:
|
||||||
|
return {}
|
||||||
|
return {
|
||||||
|
idx: window_size if layer_type == "sliding_attention" else None
|
||||||
|
for idx, layer_type in enumerate(layer_types)
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def uses_sliding_window_attention(*modules: Any) -> bool:
|
def uses_sliding_window_attention(*modules: Any) -> bool:
|
||||||
return any(
|
return any(
|
||||||
bool(getattr(module, attr, False))
|
bool(getattr(module, attr, False))
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ from __future__ import annotations
|
|||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
|
from mlx_lm.models.base import create_causal_mask
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.hardware_backend.mlx.kv_cache.attention_kv_pool import (
|
from sglang.srt.hardware_backend.mlx.kv_cache.attention_kv_pool import (
|
||||||
@@ -12,6 +13,23 @@ if TYPE_CHECKING:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def make_attention_mask(N, offset, return_array=False, window_size=None):
|
||||||
|
"""Mirror mlx_lm ``cache.create_attention_mask`` for cache shims.
|
||||||
|
|
||||||
|
Containers delegate mask creation to ``cache.make_mask`` whenever the
|
||||||
|
cache exposes it, so the shims must honor ``window_size`` (sliding-window
|
||||||
|
layers pass it, including for N == 1) or windowed models silently fall
|
||||||
|
back to full attention.
|
||||||
|
"""
|
||||||
|
if window_size is not None:
|
||||||
|
return create_causal_mask(N, offset, window_size=window_size)
|
||||||
|
if N == 1:
|
||||||
|
return None
|
||||||
|
if return_array:
|
||||||
|
return create_causal_mask(N, offset)
|
||||||
|
return "causal"
|
||||||
|
|
||||||
|
|
||||||
class AttentionOffsetCache:
|
class AttentionOffsetCache:
|
||||||
"""Data-free shim satisfying mlx-lm's cache protocol.
|
"""Data-free shim satisfying mlx-lm's cache protocol.
|
||||||
|
|
||||||
@@ -25,8 +43,10 @@ class AttentionOffsetCache:
|
|||||||
def state(self):
|
def state(self):
|
||||||
return () # Empty — safe for mx.eval unpacking
|
return () # Empty — safe for mx.eval unpacking
|
||||||
|
|
||||||
def make_mask(self, N, **kwargs):
|
def make_mask(self, N, return_array=False, window_size=None, **kwargs):
|
||||||
return None if N == 1 else "causal"
|
return make_attention_mask(
|
||||||
|
N, self.offset, return_array=return_array, window_size=window_size
|
||||||
|
)
|
||||||
|
|
||||||
def update_and_fetch(self, keys, values):
|
def update_and_fetch(self, keys, values):
|
||||||
raise RuntimeError("AttentionOffsetCache should not store data")
|
raise RuntimeError("AttentionOffsetCache should not store data")
|
||||||
@@ -60,6 +80,11 @@ class ContiguousAttentionKVCache:
|
|||||||
self.offset = 0
|
self.offset = 0
|
||||||
self.max_seq_len = max_seq_len
|
self.max_seq_len = max_seq_len
|
||||||
|
|
||||||
|
def make_mask(self, N, return_array=False, window_size=None, **kwargs):
|
||||||
|
return make_attention_mask(
|
||||||
|
N, self.offset, return_array=return_array, window_size=window_size
|
||||||
|
)
|
||||||
|
|
||||||
def _allocate(self, keys: mx.array) -> None:
|
def _allocate(self, keys: mx.array) -> None:
|
||||||
"""Allocate buffers matching the first key tensor's shape."""
|
"""Allocate buffers matching the first key tensor's shape."""
|
||||||
B, n_kv_heads, _, head_dim = keys.shape
|
B, n_kv_heads, _, head_dim = keys.shape
|
||||||
@@ -77,9 +102,6 @@ class ContiguousAttentionKVCache:
|
|||||||
return ()
|
return ()
|
||||||
return (self.keys, self.values)
|
return (self.keys, self.values)
|
||||||
|
|
||||||
def make_mask(self, N, **kwargs):
|
|
||||||
return None if N == 1 else "causal"
|
|
||||||
|
|
||||||
def _grow(self, required: int) -> None:
|
def _grow(self, required: int) -> None:
|
||||||
"""Double the buffer until it can hold *required* tokens."""
|
"""Double the buffer until it can hold *required* tokens."""
|
||||||
new_max = self.max_seq_len
|
new_max = self.max_seq_len
|
||||||
@@ -173,8 +195,10 @@ class PoolBackedAttentionKVCache:
|
|||||||
return (self._full_keys, self._full_values)
|
return (self._full_keys, self._full_values)
|
||||||
return ()
|
return ()
|
||||||
|
|
||||||
def make_mask(self, N, **kwargs):
|
def make_mask(self, N, return_array=False, window_size=None, **kwargs):
|
||||||
return None if N == 1 else "causal"
|
return make_attention_mask(
|
||||||
|
N, self.offset, return_array=return_array, window_size=window_size
|
||||||
|
)
|
||||||
|
|
||||||
def update_and_fetch(
|
def update_and_fetch(
|
||||||
self, keys: mx.array, values: mx.array
|
self, keys: mx.array, values: mx.array
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ from sglang.srt.hardware_backend.mlx.aot import (
|
|||||||
MlxAOTRoPEContext,
|
MlxAOTRoPEContext,
|
||||||
)
|
)
|
||||||
from sglang.srt.hardware_backend.mlx.kv_cache.attention_contract import (
|
from sglang.srt.hardware_backend.mlx.kv_cache.attention_contract import (
|
||||||
|
get_attention_scale,
|
||||||
get_head_dim,
|
get_head_dim,
|
||||||
get_num_heads,
|
get_num_heads,
|
||||||
get_num_kv_heads,
|
get_num_kv_heads,
|
||||||
@@ -120,12 +121,29 @@ class MLXAttentionWrapper(nn.Module):
|
|||||||
|
|
||||||
When ``BatchedDecodeContext`` is set, performs per-request RoPE,
|
When ``BatchedDecodeContext`` is set, performs per-request RoPE,
|
||||||
cache writes, and batched SDPA. Otherwise delegates to inner module.
|
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.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, inner: nn.Module, layer_idx: int):
|
def __init__(
|
||||||
|
self, inner: nn.Module, layer_idx: int, window_size: int | None = None
|
||||||
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
object.__setattr__(self, "_inner", inner)
|
object.__setattr__(self, "_inner", inner)
|
||||||
object.__setattr__(self, "_layer_idx", layer_idx)
|
object.__setattr__(self, "_layer_idx", layer_idx)
|
||||||
|
object.__setattr__(self, "_window_size", window_size)
|
||||||
|
# Resolved once at patch time (weights are loaded before patching and
|
||||||
|
# the inner module is never swapped afterwards), keeping the decode
|
||||||
|
# hot path free of attribute scans and failing fast on a bad module.
|
||||||
|
scale = get_attention_scale(inner)
|
||||||
|
if scale is None:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Cannot determine attention scale for {type(inner).__name__}"
|
||||||
|
)
|
||||||
|
object.__setattr__(self, "_scale", scale)
|
||||||
|
object.__setattr__(self, "_sinks", getattr(inner, "sinks", None))
|
||||||
|
|
||||||
def __call__(self, x: mx.array, mask: Any = None, cache: Any = None) -> mx.array:
|
def __call__(self, x: mx.array, mask: Any = None, cache: Any = None) -> mx.array:
|
||||||
ctx = get_context()
|
ctx = get_context()
|
||||||
@@ -200,7 +218,17 @@ class MLXAttentionWrapper(nn.Module):
|
|||||||
keys = inner.rope(keys, offset=offsets)
|
keys = inner.rope(keys, offset=offsets)
|
||||||
|
|
||||||
layer_caches = ctx.attention_layer_caches[attention_pool_idx]
|
layer_caches = ctx.attention_layer_caches[attention_pool_idx]
|
||||||
pad_sizes = ctx.pad_sizes
|
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]
|
||||||
|
|
||||||
# TODO: replace per-request loop with native batched/ragged
|
# TODO: replace per-request loop with native batched/ragged
|
||||||
# attention once mx.fast.scaled_dot_product_attention supports
|
# attention once mx.fast.scaled_dot_product_attention supports
|
||||||
@@ -212,6 +240,9 @@ class MLXAttentionWrapper(nn.Module):
|
|||||||
layer_caches[i].write_token(keys[i : i + 1], values[i : i + 1])
|
layer_caches[i].write_token(keys[i : i + 1], values[i : i + 1])
|
||||||
|
|
||||||
k_all, v_all = layer_caches[i].get_kv()
|
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:, :]
|
||||||
|
|
||||||
pad = pad_sizes[i]
|
pad = pad_sizes[i]
|
||||||
if pad > 0:
|
if pad > 0:
|
||||||
@@ -226,17 +257,34 @@ class MLXAttentionWrapper(nn.Module):
|
|||||||
keys_b = mx.concatenate(all_k, axis=0)
|
keys_b = mx.concatenate(all_k, axis=0)
|
||||||
values_b = mx.concatenate(all_v, 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
|
attn_mask = None
|
||||||
if ctx.needs_padding:
|
if pad_mask is not None:
|
||||||
mask_bool = ctx.positions[None, :] >= ctx.valid_lens[:, None]
|
|
||||||
attn_mask = mx.where(
|
attn_mask = mx.where(
|
||||||
mask_bool[:, None, None, :],
|
pad_mask[:, None, None, :],
|
||||||
mx.array(mx.finfo(queries.dtype).min, dtype=queries.dtype),
|
mx.array(mx.finfo(queries.dtype).min, dtype=queries.dtype),
|
||||||
mx.array(0.0, 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(
|
output = mx.fast.scaled_dot_product_attention(
|
||||||
queries, keys_b, values_b, scale=inner.scale, mask=attn_mask
|
queries,
|
||||||
|
keys_b,
|
||||||
|
values_b,
|
||||||
|
scale=self._scale,
|
||||||
|
mask=attn_mask,
|
||||||
|
**sink_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
output = output.transpose(0, 2, 1, 3).reshape(B, 1, -1)
|
output = output.transpose(0, 2, 1, 3).reshape(B, 1, -1)
|
||||||
|
|||||||
@@ -1,16 +1,21 @@
|
|||||||
"""Model introspection and attention patching."""
|
"""Model introspection and attention patching."""
|
||||||
|
|
||||||
|
import logging
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
import mlx.nn as nn
|
import mlx.nn as nn
|
||||||
|
|
||||||
from sglang.srt.hardware_backend.mlx.kv_cache.attention_contract import (
|
from sglang.srt.hardware_backend.mlx.kv_cache.attention_contract import (
|
||||||
|
get_container_window_size,
|
||||||
|
get_layer_window_sizes,
|
||||||
is_attention_module,
|
is_attention_module,
|
||||||
)
|
)
|
||||||
from sglang.srt.hardware_backend.mlx.kv_cache.attention_wrapper import (
|
from sglang.srt.hardware_backend.mlx.kv_cache.attention_wrapper import (
|
||||||
MLXAttentionWrapper,
|
MLXAttentionWrapper,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
def _find_attention_attr(layer: Any) -> str | None:
|
def _find_attention_attr(layer: Any) -> str | None:
|
||||||
"""Return the direct child name that satisfies the attention contract."""
|
"""Return the direct child name that satisfies the attention contract."""
|
||||||
@@ -43,6 +48,19 @@ def patch_model_attention(model: Any) -> int:
|
|||||||
is set, so it is always installed and never removed.
|
is set, so it is always installed and never removed.
|
||||||
"""
|
"""
|
||||||
layer_list, attn_attrs = find_attention_layers(model)
|
layer_list, attn_attrs = find_attention_layers(model)
|
||||||
|
window_sizes = get_layer_window_sizes(model)
|
||||||
|
if not window_sizes and get_container_window_size(model) is not None:
|
||||||
|
# e.g. gemma3-style containers derive per-layer windows from a
|
||||||
|
# pattern instead of ``layer_types``. Prefill masks (delegated to
|
||||||
|
# the container) honor the window, but batched decode cannot
|
||||||
|
# without a per-layer map, so outputs would diverge past the
|
||||||
|
# window. Surface it instead of silently splitting semantics.
|
||||||
|
logger.warning(
|
||||||
|
"Model %s declares a sliding window but no per-layer "
|
||||||
|
"layer_types map; MLX batched decode will not apply the "
|
||||||
|
"window and long-context output may be incorrect.",
|
||||||
|
type(model).__name__,
|
||||||
|
)
|
||||||
patched = 0
|
patched = 0
|
||||||
for idx, (layer, attn_attr) in enumerate(zip(layer_list, attn_attrs)):
|
for idx, (layer, attn_attr) in enumerate(zip(layer_list, attn_attrs)):
|
||||||
if attn_attr is None:
|
if attn_attr is None:
|
||||||
@@ -50,7 +68,11 @@ def patch_model_attention(model: Any) -> int:
|
|||||||
attn = getattr(layer, attn_attr)
|
attn = getattr(layer, attn_attr)
|
||||||
if isinstance(attn, MLXAttentionWrapper):
|
if isinstance(attn, MLXAttentionWrapper):
|
||||||
continue
|
continue
|
||||||
setattr(layer, attn_attr, MLXAttentionWrapper(attn, idx))
|
setattr(
|
||||||
|
layer,
|
||||||
|
attn_attr,
|
||||||
|
MLXAttentionWrapper(attn, idx, window_size=window_sizes.get(idx)),
|
||||||
|
)
|
||||||
patched += 1
|
patched += 1
|
||||||
return patched
|
return patched
|
||||||
|
|
||||||
|
|||||||
@@ -308,3 +308,12 @@ class MlxModelRunnerStub(ModelRunner):
|
|||||||
def alloc_memory_pool(self, memory_pool_config=None):
|
def alloc_memory_pool(self, memory_pool_config=None):
|
||||||
"""No-op: MLX manages its own KV cache."""
|
"""No-op: MLX manages its own KV cache."""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
def init_attention_backends(self):
|
||||||
|
"""No-op: attention runs inside the MLX runner.
|
||||||
|
|
||||||
|
The backend named by ``server_args.attention_backend`` would never
|
||||||
|
be used, and building one can crash: some backends read real KV
|
||||||
|
buffers in ``__init__``, which this stub never allocates.
|
||||||
|
"""
|
||||||
|
self.attn_backend = None
|
||||||
|
|||||||
@@ -5359,28 +5359,30 @@ class ServerArgs:
|
|||||||
elif model_arch in ["GptOssForCausalLM"]:
|
elif model_arch in ["GptOssForCausalLM"]:
|
||||||
# Attention backend selection + XPU dtype validation moved to the
|
# Attention backend selection + XPU dtype validation moved to the
|
||||||
# override registry (arg_groups/overrides.py: _gpt_oss_overrides).
|
# override registry (arg_groups/overrides.py: _gpt_oss_overrides).
|
||||||
|
# None of these backends exist on MPS; attention_backend is still
|
||||||
supported_backends = [
|
# unset there at this point (the torch_native default fills later).
|
||||||
"triton",
|
if not is_mps():
|
||||||
"trtllm_mha",
|
supported_backends = [
|
||||||
"fa3",
|
"triton",
|
||||||
"fa4",
|
"trtllm_mha",
|
||||||
"ascend",
|
"fa3",
|
||||||
"intel_amx",
|
"fa4",
|
||||||
"intel_xpu",
|
"ascend",
|
||||||
"aiter",
|
"intel_amx",
|
||||||
]
|
"intel_xpu",
|
||||||
prefill_attn_backend, decode_attn_backend = (
|
"aiter",
|
||||||
self._resolved_attention_backends()
|
]
|
||||||
)
|
prefill_attn_backend, decode_attn_backend = (
|
||||||
assert (
|
self._resolved_attention_backends()
|
||||||
prefill_attn_backend in supported_backends
|
)
|
||||||
and decode_attn_backend in supported_backends
|
assert (
|
||||||
), (
|
prefill_attn_backend in supported_backends
|
||||||
f"GptOssForCausalLM requires one of {supported_backends} attention backend, but got the following backends\n"
|
and decode_attn_backend in supported_backends
|
||||||
f"- Prefill: {prefill_attn_backend}\n"
|
), (
|
||||||
f"- Decode: {decode_attn_backend}\n"
|
f"GptOssForCausalLM requires one of {supported_backends} attention backend, but got the following backends\n"
|
||||||
)
|
f"- Prefill: {prefill_attn_backend}\n"
|
||||||
|
f"- Decode: {decode_attn_backend}\n"
|
||||||
|
)
|
||||||
|
|
||||||
quant_method = get_quantization_config(hf_config)
|
quant_method = get_quantization_config(hf_config)
|
||||||
is_mxfp4_quant_format = quant_method == "mxfp4"
|
is_mxfp4_quant_format = quant_method == "mxfp4"
|
||||||
|
|||||||
@@ -0,0 +1,371 @@
|
|||||||
|
"""Correctness tests for gpt-oss served on the SGLang MLX backend.
|
||||||
|
|
||||||
|
gpt-oss interleaves sliding-window (window=128) and full-attention layers and
|
||||||
|
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.
|
||||||
|
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 /
|
||||||
|
trailing-window truncation, while vanilla mlx_lm uses RotatingKVCache
|
||||||
|
for sliding layers — mathematically equivalent, so tokens must match
|
||||||
|
exactly.
|
||||||
|
|
||||||
|
Both follow the structure of the qwen MoE MLX correctness tests
|
||||||
|
(PR #29440).
|
||||||
|
|
||||||
|
Prompt length matters for both: sequences up to 128 tokens never engage the
|
||||||
|
window (banded and causal masks coincide), so a short-prompt test passes even
|
||||||
|
if window handling is completely broken. Prompts here are >128 tokens. They
|
||||||
|
also stay well under 2048 tokens: past mlx_lm's prefill chunk size the
|
||||||
|
RotatingKVCache reference trims differently and exact token equality no
|
||||||
|
longer holds by construction.
|
||||||
|
|
||||||
|
MLX-gated like its siblings: registered on the CPU suite but skipped wherever
|
||||||
|
``mlx`` is absent (all current CI runners); runs for real only on Apple
|
||||||
|
Silicon. The default 20B model needs ~11 GB of weights — override with
|
||||||
|
``SGLANG_MLX_TEST_MODEL`` (e.g. a local download of
|
||||||
|
``mlx-community/gpt-oss-20b-MXFP4-Q8``).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import gc
|
||||||
|
import importlib.util
|
||||||
|
import os
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import requests
|
||||||
|
|
||||||
|
from sglang.srt.utils import kill_process_tree
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci
|
||||||
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
DEFAULT_URL_FOR_TEST,
|
||||||
|
CustomTestCase,
|
||||||
|
popen_launch_server,
|
||||||
|
try_cached_model,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Registered on the CPU suite but skipped wherever mlx is absent; runs for real
|
||||||
|
# only on Apple Silicon. Also registered under stage-b-e2e-mlx, which the
|
||||||
|
# macOS CI lane (pr-test-mlx.yml) only dispatches via a gated workflow_dispatch.
|
||||||
|
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||||
|
register_mlx_ci(est_time=1, suite="stage-b-e2e-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 (Apple Silicon only)"
|
||||||
|
|
||||||
|
MODEL_PATH = os.environ.get(
|
||||||
|
"SGLANG_MLX_TEST_MODEL", "mlx-community/gpt-oss-20b-MXFP4-Q8"
|
||||||
|
)
|
||||||
|
MEM_FRACTION_STATIC = os.environ.get("SGLANG_MLX_TEST_MEM_FRACTION", "0.9")
|
||||||
|
# Skip (do NOT crash) unless this much system memory is free; an MLX Metal
|
||||||
|
# OOM is uncatchable and can reboot the machine. ~12 GB suits the default
|
||||||
|
# 20B MXFP4-Q8 repo (11 GB of weights).
|
||||||
|
MIN_FREE_GB = float(os.environ.get("SGLANG_MLX_TEST_MIN_FREE_GB", "12"))
|
||||||
|
|
||||||
|
# Filler that pushes every prompt past 128 tokens (the gpt-oss sliding
|
||||||
|
# window) while staying far below 2048. The question at the end keeps greedy
|
||||||
|
# answers short and deterministic.
|
||||||
|
_NUMBER_LIST = "The following is a list of numbers: " + ", ".join(
|
||||||
|
str(i) for i in range(1, 121)
|
||||||
|
)
|
||||||
|
LONG_PROMPTS = [
|
||||||
|
_NUMBER_LIST + ". Which number comes right after 57? Answer briefly.",
|
||||||
|
_NUMBER_LIST + ". What is the sum of the first three numbers? Answer briefly.",
|
||||||
|
]
|
||||||
|
MAX_NEW_TOKENS = 64 # equivalence horizon; analysis-channel tokens count too
|
||||||
|
BATCH_HORIZON = 24 # fixed step count for the batching-isolation test
|
||||||
|
|
||||||
|
|
||||||
|
def _available_gb():
|
||||||
|
try:
|
||||||
|
import psutil
|
||||||
|
|
||||||
|
return psutil.virtual_memory().available / 1024**3
|
||||||
|
except Exception:
|
||||||
|
return None # psutil absent -> skip the pre-flight check
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(_HAS_MLX, _SKIP_REASON)
|
||||||
|
class TestGptOssMlxCorrectness(CustomTestCase):
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
avail = _available_gb()
|
||||||
|
if avail is not None and avail < MIN_FREE_GB:
|
||||||
|
raise unittest.SkipTest(
|
||||||
|
f"insufficient free memory: {avail:.1f} GB < {MIN_FREE_GB} GB "
|
||||||
|
f"needed to safely serve {MODEL_PATH} "
|
||||||
|
f"(override SGLANG_MLX_TEST_MIN_FREE_GB)"
|
||||||
|
)
|
||||||
|
|
||||||
|
cls.model = try_cached_model(MODEL_PATH)
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
|
||||||
|
env = os.environ.copy()
|
||||||
|
env["SGLANG_USE_MLX"] = "1"
|
||||||
|
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=[
|
||||||
|
"--trust-remote-code",
|
||||||
|
"--tp-size",
|
||||||
|
"1",
|
||||||
|
"--disable-radix-cache",
|
||||||
|
"--disable-cuda-graph",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
MEM_FRACTION_STATIC,
|
||||||
|
"--max-running-requests",
|
||||||
|
"1",
|
||||||
|
"--context-length",
|
||||||
|
"2048",
|
||||||
|
],
|
||||||
|
env=env,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
if hasattr(cls, "process") and cls.process is not None:
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
def _chat(self, messages, max_tokens=64, temperature=0):
|
||||||
|
resp = requests.post(
|
||||||
|
f"{self.base_url}/v1/chat/completions",
|
||||||
|
json={
|
||||||
|
"model": MODEL_PATH,
|
||||||
|
"messages": messages,
|
||||||
|
"temperature": temperature,
|
||||||
|
"max_tokens": max_tokens,
|
||||||
|
},
|
||||||
|
timeout=300,
|
||||||
|
)
|
||||||
|
resp.raise_for_status()
|
||||||
|
return resp.json()["choices"][0]["message"]["content"].strip()
|
||||||
|
|
||||||
|
def test_basic_generation_nonempty(self):
|
||||||
|
text = self._chat(
|
||||||
|
[
|
||||||
|
{"role": "system", "content": "You are a concise assistant."},
|
||||||
|
{"role": "user", "content": "Say hello briefly."},
|
||||||
|
],
|
||||||
|
max_tokens=32,
|
||||||
|
)
|
||||||
|
self.assertIsInstance(text, str)
|
||||||
|
self.assertGreater(len(text), 0)
|
||||||
|
|
||||||
|
def test_simple_arithmetic(self):
|
||||||
|
text = self._chat(
|
||||||
|
[
|
||||||
|
{"role": "system", "content": "You are a concise assistant."},
|
||||||
|
{"role": "user", "content": "What is 2+2? Reply with just the number."},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
self.assertIn("4", text)
|
||||||
|
|
||||||
|
def test_long_prompt_engages_sliding_window(self):
|
||||||
|
# >128 prompt tokens: prefill and decode both run with the sliding
|
||||||
|
# window engaged on half the layers. The needle sits near the end of
|
||||||
|
# the prompt, inside the window of the final positions.
|
||||||
|
text = self._chat(
|
||||||
|
[
|
||||||
|
{"role": "system", "content": "You are a concise assistant."},
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": (
|
||||||
|
_NUMBER_LIST + ". The secret word is BLUEBERRY. "
|
||||||
|
"What is the secret word? Answer briefly."
|
||||||
|
),
|
||||||
|
},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
self.assertIn("BLUEBERRY", text.upper())
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(_HAS_MLX, _SKIP_REASON)
|
||||||
|
class TestGptOssMlxReferenceCorrectness(CustomTestCase):
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
import mlx.core as mx
|
||||||
|
from mlx_lm import load
|
||||||
|
|
||||||
|
avail = _available_gb()
|
||||||
|
if avail is not None and avail < MIN_FREE_GB:
|
||||||
|
raise unittest.SkipTest(
|
||||||
|
f"insufficient free memory: {avail:.1f} GB < {MIN_FREE_GB} GB needed "
|
||||||
|
f"to safely load {MODEL_PATH} (override SGLANG_MLX_TEST_MIN_FREE_GB)"
|
||||||
|
)
|
||||||
|
|
||||||
|
model_path = try_cached_model(MODEL_PATH)
|
||||||
|
|
||||||
|
# --- Phase 1: reference tokens from UNPATCHED mlx_lm (one copy resident) ---
|
||||||
|
try:
|
||||||
|
ref_model, cls.tokenizer = load(
|
||||||
|
model_path, tokenizer_config={"trust_remote_code": True}
|
||||||
|
)
|
||||||
|
except Exception as exc: # not cached / offline / bad path
|
||||||
|
raise unittest.SkipTest(f"could not load {MODEL_PATH}: {exc}")
|
||||||
|
|
||||||
|
eos = getattr(cls.tokenizer, "eos_token_ids", None) or {
|
||||||
|
cls.tokenizer.eos_token_id
|
||||||
|
}
|
||||||
|
cls.eos_ids = set(eos)
|
||||||
|
|
||||||
|
cls.cases = [] # (prompt, prompt_ids, reference_token_ids)
|
||||||
|
for prompt in LONG_PROMPTS:
|
||||||
|
prompt_ids = list(
|
||||||
|
cls.tokenizer.apply_chat_template(
|
||||||
|
[{"role": "user", "content": prompt}], add_generation_prompt=True
|
||||||
|
)
|
||||||
|
)
|
||||||
|
# The whole point of this test: the sliding window only engages
|
||||||
|
# past 128 tokens, and the RotatingKVCache reference only stays
|
||||||
|
# trim-free below mlx_lm's prefill chunking threshold.
|
||||||
|
assert 128 < len(prompt_ids) <= 2048, (
|
||||||
|
f"prompt must be >128 and <=2048 tokens to exercise the "
|
||||||
|
f"sliding window, got {len(prompt_ids)}"
|
||||||
|
)
|
||||||
|
ref_ids = cls._reference_greedy(
|
||||||
|
ref_model, cls.tokenizer, prompt_ids, MAX_NEW_TOKENS
|
||||||
|
)
|
||||||
|
cls.cases.append((prompt, prompt_ids, ref_ids))
|
||||||
|
|
||||||
|
# --- Release the reference BEFORE building the runner (cap peak at 1x) ---
|
||||||
|
del ref_model
|
||||||
|
gc.collect()
|
||||||
|
mx.clear_cache()
|
||||||
|
active_gb = mx.get_active_memory() / 1024**3
|
||||||
|
if active_gb > 2.0:
|
||||||
|
raise unittest.SkipTest(
|
||||||
|
f"reference model not released (active={active_gb:.1f} GB); "
|
||||||
|
"skipping to avoid a double-resident OOM"
|
||||||
|
)
|
||||||
|
|
||||||
|
# --- Phase 2: SGLang runner (one copy resident) ---
|
||||||
|
from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner
|
||||||
|
|
||||||
|
cls.runner = MlxModelRunner(
|
||||||
|
model_path=model_path,
|
||||||
|
trust_remote_code=True,
|
||||||
|
disable_radix_cache=True, # per-request contiguous caches; no big pool
|
||||||
|
mem_fraction_static=float(MEM_FRACTION_STATIC),
|
||||||
|
)
|
||||||
|
cls.runner.init_cache_pools(req_to_token_pool=None)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
runner = getattr(cls, "runner", None)
|
||||||
|
if runner is not None:
|
||||||
|
runner.clear()
|
||||||
|
cls.runner = None
|
||||||
|
gc.collect()
|
||||||
|
try:
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
mx.clear_cache()
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# --- helpers ----------------------------------------------------------
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _reference_greedy(model, tokenizer, prompt_ids, max_new):
|
||||||
|
"""Ground-truth token ids from raw, unpatched mlx_lm greedy generation."""
|
||||||
|
import mlx.core as mx
|
||||||
|
from mlx_lm import stream_generate
|
||||||
|
from mlx_lm.sample_utils import make_sampler
|
||||||
|
|
||||||
|
sampler = make_sampler(temp=0.0) # greedy / argmax
|
||||||
|
out = []
|
||||||
|
for resp in stream_generate(
|
||||||
|
model, tokenizer, mx.array(prompt_ids), max_tokens=max_new, sampler=sampler
|
||||||
|
):
|
||||||
|
out.append(int(resp.token))
|
||||||
|
return out
|
||||||
|
|
||||||
|
def _prefill(self, rid, prompt_ids):
|
||||||
|
return int(
|
||||||
|
self.runner.prefill(
|
||||||
|
req_id=rid,
|
||||||
|
new_token_ids=list(prompt_ids),
|
||||||
|
full_token_ids=list(prompt_ids),
|
||||||
|
prefix_slot_ids=[],
|
||||||
|
new_slot_ids=[],
|
||||||
|
req_pool_idx=0,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
def _decode(self, rids):
|
||||||
|
return [int(t) for t in self.runner.decode_batch(rids)]
|
||||||
|
|
||||||
|
def _sglang_greedy(self, rid, prompt_ids, max_new):
|
||||||
|
"""SGLang MLX greedy generation, stopping at EOS like the reference."""
|
||||||
|
tok = self._prefill(rid, prompt_ids)
|
||||||
|
out = [tok]
|
||||||
|
while len(out) < max_new and tok not in self.eos_ids:
|
||||||
|
tok = self._decode([rid])[0]
|
||||||
|
out.append(tok)
|
||||||
|
self.runner.remove_request(rid)
|
||||||
|
return out
|
||||||
|
|
||||||
|
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)
|
||||||
|
return (
|
||||||
|
f"\nprompt: {prompt[:80]!r}..."
|
||||||
|
f"\n first divergence @ index {first} (len ref={len(ref)} sgl={len(sgl)})"
|
||||||
|
f"\n ref text: {self.tokenizer.decode(ref)!r}"
|
||||||
|
f"\n sgl text: {self.tokenizer.decode(sgl)!r}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# --- tests ------------------------------------------------------------
|
||||||
|
|
||||||
|
def test_greedy_matches_reference_exact(self):
|
||||||
|
"""SGLang MLX greedy output == unpatched mlx_lm greedy output, token-for-token."""
|
||||||
|
for i, (prompt, prompt_ids, ref) in enumerate(self.cases):
|
||||||
|
sgl = self._sglang_greedy(f"ref-{i}", prompt_ids, MAX_NEW_TOKENS)
|
||||||
|
self.assertEqual(sgl, ref, self._diff_msg(prompt, ref, sgl))
|
||||||
|
|
||||||
|
def test_batched_decode_matches_solo(self):
|
||||||
|
"""A request's tokens are identical whether decoded alone or in a batch.
|
||||||
|
|
||||||
|
Pins slot/cache isolation for the sliding-window decode path: the
|
||||||
|
per-request trailing-window truncation and locally rebuilt padding
|
||||||
|
mask must not let one request's state bleed into another's.
|
||||||
|
"""
|
||||||
|
ids_list = [prompt_ids for (_, prompt_ids, _) in self.cases]
|
||||||
|
|
||||||
|
# Solo: prefill, decode a fixed horizon, remove -- one request at a time.
|
||||||
|
solo = []
|
||||||
|
for i, ids in enumerate(ids_list):
|
||||||
|
seq = [self._prefill(f"solo-{i}", ids)]
|
||||||
|
for _ in range(BATCH_HORIZON - 1):
|
||||||
|
seq.append(self._decode([f"solo-{i}"])[0])
|
||||||
|
self.runner.remove_request(f"solo-{i}")
|
||||||
|
solo.append(seq)
|
||||||
|
|
||||||
|
# Batched: prefill all, then advance them together in one decode_batch.
|
||||||
|
rids = [f"batch-{i}" for i in range(len(ids_list))]
|
||||||
|
batched = [[self._prefill(rid, ids)] for rid, ids in zip(rids, ids_list)]
|
||||||
|
for _ in range(BATCH_HORIZON - 1):
|
||||||
|
for j, t in enumerate(self._decode(rids)):
|
||||||
|
batched[j].append(t)
|
||||||
|
for rid in rids:
|
||||||
|
self.runner.remove_request(rid)
|
||||||
|
|
||||||
|
for i, (prompt, _, _) in enumerate(self.cases):
|
||||||
|
self.assertEqual(
|
||||||
|
batched[i], solo[i], self._diff_msg(prompt, solo[i], batched[i])
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main(verbosity=3)
|
||||||
@@ -1,4 +1,4 @@
|
|||||||
"""Guard the MLX stub's ``alloc_memory_pool`` override against drift.
|
"""Guard the MLX stub's ``ModelRunner`` overrides against drift.
|
||||||
|
|
||||||
The base ``ModelRunner.alloc_memory_pool`` runs ``_init_pools`` which
|
The base ``ModelRunner.alloc_memory_pool`` runs ``_init_pools`` which
|
||||||
asserts ``is_draft_worker`` (model_runner_kv_cache_mixin.py:409); the
|
asserts ``is_draft_worker`` (model_runner_kv_cache_mixin.py:409); the
|
||||||
@@ -6,6 +6,11 @@ MLX stub manages its own KV cache via ``MlxAttentionKVPool`` and must
|
|||||||
short-circuit that GPU-allocation path. If the override is lost, every
|
short-circuit that GPU-allocation path. If the override is lost, every
|
||||||
MLX startup crashes inside ``Scheduler.init_target_memory_pool``.
|
MLX startup crashes inside ``Scheduler.init_target_memory_pool``.
|
||||||
|
|
||||||
|
Similarly, the base ``init_attention_backends`` constructs the torch
|
||||||
|
attention backend named by ``server_args.attention_backend``; MLX never
|
||||||
|
uses one, and some backends read real KV buffers in ``__init__``, which
|
||||||
|
crashes on ``_DummyKVCache``.
|
||||||
|
|
||||||
The checks are signature/identity-only and MLX-gated because importing
|
The checks are signature/identity-only and MLX-gated because importing
|
||||||
the stub pulls in ``mlx.core``.
|
the stub pulls in ``mlx.core``.
|
||||||
"""
|
"""
|
||||||
@@ -75,6 +80,31 @@ class TestMlxRunnerPoolContract(unittest.TestCase):
|
|||||||
f"optional MemoryPoolConfig argument: {exc}"
|
f"optional MemoryPoolConfig argument: {exc}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_stub_overrides_base_init_attention_backends(self):
|
||||||
|
self.assertIn(
|
||||||
|
"init_attention_backends",
|
||||||
|
vars(MlxModelRunnerStub),
|
||||||
|
msg=(
|
||||||
|
"MlxModelRunnerStub lost its init_attention_backends "
|
||||||
|
"override. The base implementation constructs the backend "
|
||||||
|
"named by server_args.attention_backend; some backends "
|
||||||
|
"read real KV buffers in __init__, which crashes on "
|
||||||
|
"_DummyKVCache. MLX never uses a torch attention backend "
|
||||||
|
"— re-add the override that keeps attn_backend = None."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
self.assertIsNot(
|
||||||
|
MlxModelRunnerStub.init_attention_backends,
|
||||||
|
ModelRunner.init_attention_backends,
|
||||||
|
msg="init_attention_backends must be overridden on the MLX "
|
||||||
|
"stub, not inherited from ModelRunner.",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_stub_init_attention_backends_keeps_attn_backend_none(self):
|
||||||
|
runner = object.__new__(MlxModelRunnerStub)
|
||||||
|
MlxModelRunnerStub.init_attention_backends(runner)
|
||||||
|
self.assertIsNone(runner.attn_backend)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -0,0 +1,487 @@
|
|||||||
|
"""Unit tests for MLX sliding-window attention support (gpt-oss style models).
|
||||||
|
|
||||||
|
gpt-oss interleaves sliding-window and full-attention layers, names its
|
||||||
|
softmax scale ``sm_scale``, and adds per-head attention sinks. These tests
|
||||||
|
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``
|
||||||
|
must produce a banded mask (including for N == 1) instead of being
|
||||||
|
silently dropped, or sliding-window layers degrade to full attention.
|
||||||
|
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.
|
||||||
|
|
||||||
|
The AOT RoPE kernel gating is also pinned: YarnRoPE (used by gpt-oss) bakes
|
||||||
|
its base and scaling into precomputed ``_freqs`` plus an ``mscale`` factor,
|
||||||
|
so the vanilla-RoPE Metal kernel must reject it rather than silently compute
|
||||||
|
with base=10000 and no yarn scaling.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import importlib.util
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
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=6, suite="base-a-test-cpu")
|
||||||
|
register_mlx_ci(est_time=6, 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 mlx_lm.models.cache import KVCache
|
||||||
|
|
||||||
|
import sglang.srt.hardware_backend.mlx.aot as mlx_aot
|
||||||
|
from sglang.srt.hardware_backend.mlx.kv_cache import (
|
||||||
|
AttentionOffsetCache,
|
||||||
|
BatchedDecodeContext,
|
||||||
|
ContiguousAttentionKVCache,
|
||||||
|
MLXAttentionWrapper,
|
||||||
|
PoolBackedAttentionKVCache,
|
||||||
|
find_attention_layers,
|
||||||
|
get_attention_scale,
|
||||||
|
get_layer_window_sizes,
|
||||||
|
is_attention_module,
|
||||||
|
make_attention_mask,
|
||||||
|
patch_model_attention,
|
||||||
|
)
|
||||||
|
|
||||||
|
TINY_WINDOW = 8
|
||||||
|
|
||||||
|
|
||||||
|
def _tiny_gpt_oss_model():
|
||||||
|
"""Randomly initialized 4-layer gpt_oss with alternating sliding/full layers."""
|
||||||
|
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)
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(_HAS_MLX, _SKIP_REASON)
|
||||||
|
class TestGptOssAttentionContract(CustomTestCase):
|
||||||
|
def test_gpt_oss_attention_passes_contract(self):
|
||||||
|
model = _tiny_gpt_oss_model()
|
||||||
|
attn = model.model.layers[0].self_attn
|
||||||
|
|
||||||
|
self.assertFalse(hasattr(attn, "scale"))
|
||||||
|
self.assertTrue(hasattr(attn, "sm_scale"))
|
||||||
|
self.assertTrue(is_attention_module(attn))
|
||||||
|
|
||||||
|
layers, attrs = find_attention_layers(model)
|
||||||
|
self.assertEqual(len(layers), 4)
|
||||||
|
self.assertEqual(attrs, ["self_attn"] * 4)
|
||||||
|
|
||||||
|
def test_module_without_any_scale_attr_fails_contract(self):
|
||||||
|
attn = _tiny_gpt_oss_model().model.layers[0].self_attn
|
||||||
|
scaleless = SimpleNamespace(
|
||||||
|
q_proj=attn.q_proj,
|
||||||
|
k_proj=attn.k_proj,
|
||||||
|
v_proj=attn.v_proj,
|
||||||
|
o_proj=attn.o_proj,
|
||||||
|
rope=attn.rope,
|
||||||
|
num_attention_heads=4,
|
||||||
|
num_key_value_heads=2,
|
||||||
|
)
|
||||||
|
self.assertFalse(is_attention_module(scaleless))
|
||||||
|
|
||||||
|
def test_get_attention_scale_prefers_scale_over_sm_scale(self):
|
||||||
|
self.assertEqual(get_attention_scale(SimpleNamespace(scale=0.5)), 0.5)
|
||||||
|
self.assertEqual(get_attention_scale(SimpleNamespace(sm_scale=0.25)), 0.25)
|
||||||
|
self.assertEqual(
|
||||||
|
get_attention_scale(SimpleNamespace(scale=0.5, sm_scale=0.25)), 0.5
|
||||||
|
)
|
||||||
|
self.assertIsNone(get_attention_scale(SimpleNamespace()))
|
||||||
|
|
||||||
|
def test_get_layer_window_sizes_reads_gpt_oss_container(self):
|
||||||
|
windows = get_layer_window_sizes(_tiny_gpt_oss_model())
|
||||||
|
self.assertEqual(windows, {0: TINY_WINDOW, 1: None, 2: TINY_WINDOW, 3: None})
|
||||||
|
|
||||||
|
def test_get_layer_window_sizes_reads_sliding_window_alias(self):
|
||||||
|
# olmo3/llama-style containers name the scalar ``sliding_window``.
|
||||||
|
model = SimpleNamespace(
|
||||||
|
model=SimpleNamespace(
|
||||||
|
layer_types=["sliding_attention", "full_attention"],
|
||||||
|
sliding_window=16,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.assertEqual(get_layer_window_sizes(model), {0: 16, 1: None})
|
||||||
|
|
||||||
|
def test_get_layer_window_sizes_defaults_to_empty(self):
|
||||||
|
self.assertEqual(get_layer_window_sizes(SimpleNamespace()), {})
|
||||||
|
no_window = SimpleNamespace(
|
||||||
|
model=SimpleNamespace(layer_types=["sliding_attention"])
|
||||||
|
)
|
||||||
|
self.assertEqual(get_layer_window_sizes(no_window), {})
|
||||||
|
|
||||||
|
def test_patch_warns_when_window_declared_but_unmapped(self):
|
||||||
|
# A container that declares a scalar window without layer_types
|
||||||
|
# (gemma3-style pattern models): prefill masks honor the window but
|
||||||
|
# batched decode cannot; the mismatch must be surfaced.
|
||||||
|
model = _tiny_gpt_oss_model()
|
||||||
|
model.model.layer_types = []
|
||||||
|
with self.assertLogs(
|
||||||
|
"sglang.srt.hardware_backend.mlx.kv_cache.model_patching",
|
||||||
|
level="WARNING",
|
||||||
|
) as logs:
|
||||||
|
patch_model_attention(model)
|
||||||
|
self.assertTrue(any("sliding window" in msg for msg in logs.output))
|
||||||
|
wrappers = [layer.self_attn for layer in model.model.layers]
|
||||||
|
self.assertEqual([w._window_size for w in wrappers], [None] * 4)
|
||||||
|
|
||||||
|
def test_patch_model_attention_assigns_window_sizes(self):
|
||||||
|
model = _tiny_gpt_oss_model()
|
||||||
|
self.assertEqual(patch_model_attention(model), 4)
|
||||||
|
wrappers = [layer.self_attn for layer in model.model.layers]
|
||||||
|
self.assertTrue(all(isinstance(w, MLXAttentionWrapper) for w in wrappers))
|
||||||
|
self.assertEqual(
|
||||||
|
[w._window_size for w in wrappers],
|
||||||
|
[TINY_WINDOW, None, TINY_WINDOW, None],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(_HAS_MLX, _SKIP_REASON)
|
||||||
|
class TestShimMakeMask(CustomTestCase):
|
||||||
|
"""The shims must return exactly what mlx_lm's own KVCache.make_mask returns."""
|
||||||
|
|
||||||
|
def _shims(self, offset):
|
||||||
|
contig = ContiguousAttentionKVCache(
|
||||||
|
n_kv_heads=1, head_dim=2, max_seq_len=64, dtype=mx.float32
|
||||||
|
)
|
||||||
|
contig.offset = offset
|
||||||
|
pool_backed = PoolBackedAttentionKVCache(
|
||||||
|
pool=None, layer_idx=0, slots=None, prefix_len=offset
|
||||||
|
)
|
||||||
|
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 test_shims_match_mlx_lm_reference(self):
|
||||||
|
cases = [
|
||||||
|
(N, offset, window, return_array)
|
||||||
|
for N in (1, 4)
|
||||||
|
for offset in (0, 3, 9)
|
||||||
|
for window in (None, 4)
|
||||||
|
for return_array in (False, True)
|
||||||
|
]
|
||||||
|
for N, offset, window, return_array in cases:
|
||||||
|
reference = KVCache()
|
||||||
|
reference.offset = offset
|
||||||
|
ref = reference.make_mask(N, return_array=return_array, window_size=window)
|
||||||
|
for shim in self._shims(offset):
|
||||||
|
got = shim.make_mask(N, return_array=return_array, window_size=window)
|
||||||
|
self._assert_same_mask(
|
||||||
|
got,
|
||||||
|
ref,
|
||||||
|
f"{type(shim).__name__} mismatch for N={N} offset={offset} "
|
||||||
|
f"window={window} return_array={return_array}",
|
||||||
|
)
|
||||||
|
|
||||||
|
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).
|
||||||
|
mask = make_attention_mask(1, 6, window_size=4)
|
||||||
|
self.assertEqual(mask.shape, (1, 7))
|
||||||
|
self.assertEqual(
|
||||||
|
[bool(v) for v in mask[0]],
|
||||||
|
[False, False, False, True, True, True, True],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_single_token_windowed_mask_is_not_none(self):
|
||||||
|
# N == 1 must still produce a banded mask when a window is set —
|
||||||
|
# returning None here silently disables the window during decode.
|
||||||
|
mask = make_attention_mask(1, 200, window_size=128)
|
||||||
|
self.assertIsNotNone(mask)
|
||||||
|
self.assertEqual(mask.shape, (1, 201))
|
||||||
|
self.assertEqual(mx.sum(mask).item(), 128)
|
||||||
|
|
||||||
|
def test_prefill_mask_bands_each_query_row(self):
|
||||||
|
# N=5 rows starting at offset 7, W=4: row i allows [i+4, i+7].
|
||||||
|
offset, N, window = 7, 5, 4
|
||||||
|
mask = make_attention_mask(N, offset, window_size=window)
|
||||||
|
self.assertEqual(mask.shape, (N, offset + N))
|
||||||
|
for i in range(N):
|
||||||
|
allowed = {j for j in range(offset + N) if bool(mask[i, j])}
|
||||||
|
expected = set(range(offset + i - window + 1, offset + i + 1))
|
||||||
|
self.assertEqual(allowed, expected, f"row {i}")
|
||||||
|
|
||||||
|
def test_defaults_without_window_are_unchanged(self):
|
||||||
|
self.assertIsNone(make_attention_mask(1, 5))
|
||||||
|
self.assertEqual(make_attention_mask(4, 5), "causal")
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(_HAS_MLX, _SKIP_REASON)
|
||||||
|
class TestBatchedDecodeSlidingWindow(CustomTestCase):
|
||||||
|
"""Batched decode must match a hand-built decode-step reference.
|
||||||
|
|
||||||
|
The reference recomputes the decode step from the inner module's own
|
||||||
|
projections and RoPE, attending over the full (untruncated) KV with
|
||||||
|
mlx_lm's own ``create_causal_mask`` band. The wrapper instead truncates
|
||||||
|
KV to the trailing window and pads ragged requests — mathematically
|
||||||
|
identical, so the comparison is float-tight.
|
||||||
|
|
||||||
|
The reference deliberately mirrors the wrapper's tensor shapes
|
||||||
|
(projections and RoPE batched over B, not per request, and no
|
||||||
|
full-sequence forward): MLX matmul/SDPA kernels pick different code
|
||||||
|
paths per input shape, and e.g. a (2, 1, H) vs (1, 1, H) linear alone
|
||||||
|
differs by ~1e-3 on this tiny model — far above the tolerance that
|
||||||
|
makes this test able to catch real bugs.
|
||||||
|
"""
|
||||||
|
|
||||||
|
HIDDEN = 64
|
||||||
|
N_KV_HEADS = 2
|
||||||
|
HEAD_DIM = 16
|
||||||
|
|
||||||
|
def _prefill_cache(self, attn, x, window):
|
||||||
|
cache = ContiguousAttentionKVCache(
|
||||||
|
n_kv_heads=self.N_KV_HEADS,
|
||||||
|
head_dim=self.HEAD_DIM,
|
||||||
|
max_seq_len=32,
|
||||||
|
dtype=mx.float32,
|
||||||
|
)
|
||||||
|
prefix = x[:, :-1, :]
|
||||||
|
attn(
|
||||||
|
prefix,
|
||||||
|
make_attention_mask(prefix.shape[1], 0, window_size=window),
|
||||||
|
cache=cache,
|
||||||
|
)
|
||||||
|
return cache
|
||||||
|
|
||||||
|
def _project_last_tokens(self, attn, xs, offsets):
|
||||||
|
B, D = len(xs), self.HEAD_DIM
|
||||||
|
x_last = mx.concatenate([x[:, -1:, :] for x in xs], axis=0)
|
||||||
|
q = attn.q_proj(x_last).reshape(B, 1, -1, D).transpose(0, 2, 1, 3)
|
||||||
|
k = attn.k_proj(x_last).reshape(B, 1, -1, D).transpose(0, 2, 1, 3)
|
||||||
|
v = attn.v_proj(x_last).reshape(B, 1, -1, D).transpose(0, 2, 1, 3)
|
||||||
|
off = mx.array(offsets, dtype=mx.int32)
|
||||||
|
return attn.rope(q, offset=off), attn.rope(k, offset=off), v
|
||||||
|
|
||||||
|
def _reference_decode(self, attn, xs, caches, window):
|
||||||
|
"""Full-KV banded-mask decode; must be called before the wrapper
|
||||||
|
writes the decode token into the shared caches."""
|
||||||
|
offsets = [x.shape[1] - 1 for x in xs]
|
||||||
|
q, k_new, v_new = self._project_last_tokens(attn, xs, offsets)
|
||||||
|
outs = []
|
||||||
|
for i, cache in enumerate(caches):
|
||||||
|
k_prefix, v_prefix = cache.get_kv()
|
||||||
|
k = mx.concatenate([k_prefix, k_new[i : i + 1]], axis=2)
|
||||||
|
v = mx.concatenate([v_prefix, v_new[i : i + 1]], axis=2)
|
||||||
|
mask = (
|
||||||
|
create_causal_mask(1, offsets[i], window_size=window)
|
||||||
|
if window is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
outs.append(
|
||||||
|
mx.fast.scaled_dot_product_attention(
|
||||||
|
q[i : i + 1],
|
||||||
|
k,
|
||||||
|
v,
|
||||||
|
scale=attn.sm_scale,
|
||||||
|
mask=mask,
|
||||||
|
sinks=attn.sinks,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
out = mx.concatenate(outs, axis=0)
|
||||||
|
out = out.transpose(0, 2, 1, 3).reshape(len(xs), 1, -1)
|
||||||
|
return attn.o_proj(out)
|
||||||
|
|
||||||
|
def _wrapper_decode(self, attn, window, xs, caches):
|
||||||
|
wrapper = MLXAttentionWrapper(attn, layer_idx=0, window_size=window)
|
||||||
|
ctx = BatchedDecodeContext(
|
||||||
|
batch_size=len(xs),
|
||||||
|
seq_lens=[x.shape[1] - 1 for x in xs],
|
||||||
|
attention_layer_caches=[caches],
|
||||||
|
)
|
||||||
|
x_last = mx.concatenate([x[:, -1:, :] for x in xs], axis=0)
|
||||||
|
out = wrapper._batched_decode(x_last, ctx)
|
||||||
|
mx.eval(out)
|
||||||
|
return out
|
||||||
|
|
||||||
|
def _assert_matches_reference(self, attn, window, lens):
|
||||||
|
mx.random.seed(0)
|
||||||
|
xs = [mx.random.normal((1, L, self.HIDDEN)) for L in lens]
|
||||||
|
caches = [self._prefill_cache(attn, x, window) for x in xs]
|
||||||
|
ref = self._reference_decode(attn, xs, caches, window)
|
||||||
|
got = self._wrapper_decode(attn, window, xs, caches)
|
||||||
|
for i in range(len(xs)):
|
||||||
|
diff = mx.abs(got[i : i + 1] - ref[i : i + 1]).max().item()
|
||||||
|
self.assertLess(
|
||||||
|
diff,
|
||||||
|
1e-5,
|
||||||
|
f"request {i} (len={lens[i]}, window={window}) diverges "
|
||||||
|
f"from the manual decode reference by {diff}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_sliding_layer_ragged_batch_matches_reference(self):
|
||||||
|
# Request 0 crosses the window (12 > 8), request 1 stays inside (5 < 8):
|
||||||
|
# covers trailing-window truncation and the local padding mask at once.
|
||||||
|
model = _tiny_gpt_oss_model()
|
||||||
|
attn = model.model.layers[0].self_attn
|
||||||
|
self._assert_matches_reference(attn, TINY_WINDOW, lens=[12, 5])
|
||||||
|
|
||||||
|
def test_sliding_layer_all_past_window_matches_reference(self):
|
||||||
|
# Both requests exceed the window with unequal true lengths: the
|
||||||
|
# shared context reports padding but the windowed lengths are all
|
||||||
|
# equal, so the correct local pad is zero. Reusing the full-length
|
||||||
|
# ctx metadata here would inject spurious padding.
|
||||||
|
model = _tiny_gpt_oss_model()
|
||||||
|
attn = model.model.layers[0].self_attn
|
||||||
|
self._assert_matches_reference(attn, TINY_WINDOW, lens=[12, 10])
|
||||||
|
|
||||||
|
def test_sliding_layer_single_request_matches_reference(self):
|
||||||
|
# B=1 with truncation: the windowed no-padding branch (mask stays None).
|
||||||
|
model = _tiny_gpt_oss_model()
|
||||||
|
attn = model.model.layers[0].self_attn
|
||||||
|
self._assert_matches_reference(attn, TINY_WINDOW, lens=[12])
|
||||||
|
|
||||||
|
def test_full_attention_layer_matches_reference(self):
|
||||||
|
# Full-attention gpt_oss layer: sinks + sm_scale on the unwindowed path.
|
||||||
|
model = _tiny_gpt_oss_model()
|
||||||
|
attn = model.model.layers[1].self_attn
|
||||||
|
self._assert_matches_reference(attn, None, lens=[12, 5])
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(_HAS_MLX, _SKIP_REASON)
|
||||||
|
class TestSdpaSinksSemantics(CustomTestCase):
|
||||||
|
def test_sdpa_sinks_match_manual_softmax_with_sink_column(self):
|
||||||
|
# Pin mx.fast.scaled_dot_product_attention(sinks=...) to the reference
|
||||||
|
# semantics gpt-oss relies on: append one per-head sink logit to the
|
||||||
|
# softmax and drop its probability column afterwards.
|
||||||
|
mx.random.seed(0)
|
||||||
|
B, H, Lq, Lk, D = 1, 4, 5, 9, 16
|
||||||
|
scale = D**-0.5
|
||||||
|
q = mx.random.normal((B, H, Lq, D))
|
||||||
|
k = mx.random.normal((B, H, Lk, D))
|
||||||
|
v = mx.random.normal((B, H, Lk, D))
|
||||||
|
sinks = mx.random.normal((H,))
|
||||||
|
offset, window = Lk - Lq, 4
|
||||||
|
|
||||||
|
rinds = mx.arange(Lk)
|
||||||
|
linds = mx.arange(offset, offset + Lq)
|
||||||
|
mask = (linds[:, None] >= rinds[None]) & (linds[:, None] < rinds[None] + window)
|
||||||
|
|
||||||
|
out_fast = mx.fast.scaled_dot_product_attention(
|
||||||
|
q, k, v, scale=scale, mask=mask, sinks=sinks
|
||||||
|
)
|
||||||
|
|
||||||
|
scores = (q * scale) @ k.transpose(0, 1, 3, 2)
|
||||||
|
scores = mx.where(mask, scores, mx.finfo(mx.float32).min)
|
||||||
|
sink_col = mx.broadcast_to(sinks[None, :, None, None], (B, H, Lq, 1))
|
||||||
|
probs = mx.softmax(mx.concatenate([scores, sink_col], axis=-1), axis=-1)
|
||||||
|
out_manual = probs[..., :-1] @ v
|
||||||
|
|
||||||
|
diff = mx.abs(out_fast - out_manual).max().item()
|
||||||
|
self.assertLess(diff, 1e-6)
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(_HAS_MLX, _SKIP_REASON)
|
||||||
|
class TestAotRopeKernelGating(CustomTestCase):
|
||||||
|
"""The vanilla-RoPE Metal kernel must reject scaled RoPE variants."""
|
||||||
|
|
||||||
|
def _build_kernel(self, attn, head_dim=2, n_kv_heads=1):
|
||||||
|
original_loader = mlx_aot._load_metal_rope_pool_fused
|
||||||
|
mlx_aot._load_metal_rope_pool_fused = lambda: object()
|
||||||
|
try:
|
||||||
|
return mlx_aot._build_rope_kernel(
|
||||||
|
mlx_aot.MlxAOTKernelBuildInputs(
|
||||||
|
sample_attn=attn,
|
||||||
|
n_kv_heads=n_kv_heads,
|
||||||
|
head_dim=head_dim,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
mlx_aot._load_metal_rope_pool_fused = original_loader
|
||||||
|
|
||||||
|
def test_vanilla_rope_is_accepted(self):
|
||||||
|
attn = SimpleNamespace(
|
||||||
|
n_heads=2,
|
||||||
|
rope=SimpleNamespace(dims=2, traditional=False, base=10000.0),
|
||||||
|
)
|
||||||
|
self.assertTrue(self._build_kernel(attn).enabled)
|
||||||
|
|
||||||
|
def test_gpt_oss_yarn_rope_is_rejected(self):
|
||||||
|
# YarnRoPE has no ``base`` (it is baked into ``_freqs``) and applies
|
||||||
|
# mscale outside mx.fast.rope; the kernel would silently compute
|
||||||
|
# vanilla RoPE with base=10000.
|
||||||
|
attn = _tiny_gpt_oss_model().model.layers[0].self_attn
|
||||||
|
kernel = self._build_kernel(attn, head_dim=attn.head_dim, n_kv_heads=2)
|
||||||
|
self.assertFalse(kernel.enabled)
|
||||||
|
|
||||||
|
def test_missing_base_is_rejected(self):
|
||||||
|
attn = SimpleNamespace(
|
||||||
|
n_heads=2, rope=SimpleNamespace(dims=2, traditional=False)
|
||||||
|
)
|
||||||
|
self.assertFalse(self._build_kernel(attn).enabled)
|
||||||
|
|
||||||
|
def test_precomputed_freqs_are_rejected(self):
|
||||||
|
attn = SimpleNamespace(
|
||||||
|
n_heads=2,
|
||||||
|
rope=SimpleNamespace(
|
||||||
|
dims=2, traditional=False, base=10000.0, _freqs=mx.ones(1)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
self.assertFalse(self._build_kernel(attn).enabled)
|
||||||
|
|
||||||
|
def test_nontrivial_mscale_is_rejected(self):
|
||||||
|
attn = SimpleNamespace(
|
||||||
|
n_heads=2,
|
||||||
|
rope=SimpleNamespace(dims=2, traditional=False, base=10000.0, mscale=1.5),
|
||||||
|
)
|
||||||
|
self.assertFalse(self._build_kernel(attn).enabled)
|
||||||
|
|
||||||
|
def test_linear_scale_is_rejected(self):
|
||||||
|
# rope_scaling type "linear" yields nn.RoPE(..., scale=1/factor); the
|
||||||
|
# kernel computes unscaled positions and must fall back, while the
|
||||||
|
# nn.RoPE default scale of exactly 1.0 stays accepted.
|
||||||
|
def attn(scale):
|
||||||
|
return SimpleNamespace(
|
||||||
|
n_heads=2,
|
||||||
|
rope=SimpleNamespace(
|
||||||
|
dims=2, traditional=False, base=10000.0, scale=scale
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertFalse(self._build_kernel(attn(0.25)).enabled)
|
||||||
|
self.assertTrue(self._build_kernel(attn(1.0)).enabled)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user