[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:
Siming Deng
2026-08-09 18:32:36 -07:00
committed by GitHub
co-authored by R0CKSTAR Alex Nails
parent 68b961e9fb
commit 553dc0f936
12 changed files with 1100 additions and 43 deletions
+3 -1
View File
@@ -51,6 +51,7 @@ from sglang.srt.utils.common import (
is_gfx95_supported,
is_hip,
is_mnnvl_fabric_device,
is_mps,
is_musa,
is_npu,
is_sm90_supported,
@@ -921,7 +922,8 @@ def _gpt_oss_overrides(server_args: Any, hf_config: Any) -> dict:
overrides["attention_backend"] = "intel_xpu"
elif is_hip():
overrides["attention_backend"] = "aiter"
else:
elif not is_mps():
# No triton on macOS; MPS keeps the platform default.
overrides["attention_backend"] = "triton"
if is_xpu():
# Check for bf16 dtype on Intel XPU. Reads the pristine dtype request,
+17 -1
View File
@@ -139,7 +139,23 @@ def _build_rope_kernel(inputs: MlxAOTKernelBuildInputs) -> MlxAOTRoPEKernel:
# AOT kernel currently requires rope_dim == head_dim.
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)
if num_qo_heads is None:
return MlxAOTRoPEKernel()
@@ -1,7 +1,10 @@
"""Cache components for the MLX backend."""
from sglang.srt.hardware_backend.mlx.kv_cache.attention_contract import (
get_attention_scale,
get_container_window_size,
get_head_dim,
get_layer_window_sizes,
get_num_heads,
get_num_kv_heads,
is_attention_module,
@@ -11,6 +14,7 @@ from sglang.srt.hardware_backend.mlx.kv_cache.attention_kv_cache import (
AttentionOffsetCache,
ContiguousAttentionKVCache,
PoolBackedAttentionKVCache,
make_attention_mask,
)
from sglang.srt.hardware_backend.mlx.kv_cache.attention_kv_pool import (
MlxAttentionKVPool,
@@ -40,12 +44,16 @@ __all__ = [
"AttentionOffsetCache",
"ContiguousAttentionKVCache",
"find_attention_layers",
"get_attention_scale",
"get_container_window_size",
"get_head_dim",
"get_context",
"get_layer_window_sizes",
"get_num_layers",
"get_num_heads",
"get_num_kv_heads",
"is_attention_module",
"make_attention_mask",
"MLXAttentionWrapper",
"MlxAttentionKVPool",
"MlxAuxiliaryStateComponent",
@@ -4,10 +4,12 @@ from __future__ import annotations
from typing import Any, Iterable
# ``rope`` and ``scale`` are required by MLXAttentionWrapper. Keeping them in
# the contract also prevents recurrent mixers such as DeltaNet from being
# mistaken for softmax attention just because they expose projection layers.
ATTENTION_API_ATTRS = ("q_proj", "k_proj", "v_proj", "o_proj", "rope", "scale")
# ``rope`` and a softmax scale are required by MLXAttentionWrapper. Keeping
# them in the contract also prevents recurrent mixers such as DeltaNet from
# being mistaken for softmax attention just because they expose projections.
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_KV_HEAD_ATTRS = ("n_kv_heads", "num_k_heads", "num_kv_heads", "num_key_value_heads")
SLIDING_ATTENTION_ATTRS = (
@@ -17,6 +19,9 @@ SLIDING_ATTENTION_ATTRS = (
"use_sliding_window",
"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:
@@ -50,14 +55,47 @@ def get_head_dim(module: Any) -> int | 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:
return (
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_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:
return any(
bool(getattr(module, attr, False))
@@ -5,6 +5,7 @@ from __future__ import annotations
from typing import TYPE_CHECKING
import mlx.core as mx
from mlx_lm.models.base import create_causal_mask
if TYPE_CHECKING:
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:
"""Data-free shim satisfying mlx-lm's cache protocol.
@@ -25,8 +43,10 @@ class AttentionOffsetCache:
def state(self):
return () # Empty — safe for mx.eval unpacking
def make_mask(self, N, **kwargs):
return None if N == 1 else "causal"
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 update_and_fetch(self, keys, values):
raise RuntimeError("AttentionOffsetCache should not store data")
@@ -60,6 +80,11 @@ class ContiguousAttentionKVCache:
self.offset = 0
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:
"""Allocate buffers matching the first key tensor's shape."""
B, n_kv_heads, _, head_dim = keys.shape
@@ -77,9 +102,6 @@ class ContiguousAttentionKVCache:
return ()
return (self.keys, self.values)
def make_mask(self, N, **kwargs):
return None if N == 1 else "causal"
def _grow(self, required: int) -> None:
"""Double the buffer until it can hold *required* tokens."""
new_max = self.max_seq_len
@@ -173,8 +195,10 @@ class PoolBackedAttentionKVCache:
return (self._full_keys, self._full_values)
return ()
def make_mask(self, N, **kwargs):
return None if N == 1 else "causal"
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 update_and_fetch(
self, keys: mx.array, values: mx.array
@@ -15,6 +15,7 @@ from sglang.srt.hardware_backend.mlx.aot import (
MlxAOTRoPEContext,
)
from sglang.srt.hardware_backend.mlx.kv_cache.attention_contract import (
get_attention_scale,
get_head_dim,
get_num_heads,
get_num_kv_heads,
@@ -120,12 +121,29 @@ class MLXAttentionWrapper(nn.Module):
When ``BatchedDecodeContext`` is set, performs per-request RoPE,
cache writes, and batched SDPA. Otherwise delegates to inner module.
``window_size`` marks a sliding-window layer: the pool keeps the full
KV history and the wrapper attends to the trailing window only, which
is numerically identical to a rotating cache.
"""
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__()
object.__setattr__(self, "_inner", inner)
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:
ctx = get_context()
@@ -200,7 +218,17 @@ class MLXAttentionWrapper(nn.Module):
keys = inner.rope(keys, offset=offsets)
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
# 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])
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]
if pad > 0:
@@ -226,17 +257,34 @@ class MLXAttentionWrapper(nn.Module):
keys_b = mx.concatenate(all_k, axis=0)
values_b = mx.concatenate(all_v, axis=0)
pad_mask = None
if window is None:
if ctx.needs_padding:
pad_mask = ctx.positions[None, :] >= ctx.valid_lens[:, None]
elif max(pad_sizes) > 0:
eff = mx.array(eff_lens, dtype=mx.int32)
pad_mask = mx.arange(max_eff)[None, :] >= eff[:, None]
attn_mask = None
if ctx.needs_padding:
mask_bool = ctx.positions[None, :] >= ctx.valid_lens[:, None]
if pad_mask is not None:
attn_mask = mx.where(
mask_bool[:, None, None, :],
pad_mask[:, None, None, :],
mx.array(mx.finfo(queries.dtype).min, dtype=queries.dtype),
mx.array(0.0, dtype=queries.dtype),
)
# Only pass sinks when the module has them: the kwarg requires a
# recent mlx and must not constrain models without sinks.
sink_kwargs = {}
if self._sinks is not None:
sink_kwargs["sinks"] = self._sinks
output = mx.fast.scaled_dot_product_attention(
queries, keys_b, values_b, scale=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)
@@ -1,16 +1,21 @@
"""Model introspection and attention patching."""
import logging
from typing import Any
import mlx.nn as nn
from sglang.srt.hardware_backend.mlx.kv_cache.attention_contract import (
get_container_window_size,
get_layer_window_sizes,
is_attention_module,
)
from sglang.srt.hardware_backend.mlx.kv_cache.attention_wrapper import (
MLXAttentionWrapper,
)
logger = logging.getLogger(__name__)
def _find_attention_attr(layer: Any) -> str | None:
"""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.
"""
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
for idx, (layer, attn_attr) in enumerate(zip(layer_list, attn_attrs)):
if attn_attr is None:
@@ -50,7 +68,11 @@ def patch_model_attention(model: Any) -> int:
attn = getattr(layer, attn_attr)
if isinstance(attn, MLXAttentionWrapper):
continue
setattr(layer, attn_attr, MLXAttentionWrapper(attn, idx))
setattr(
layer,
attn_attr,
MLXAttentionWrapper(attn, idx, window_size=window_sizes.get(idx)),
)
patched += 1
return patched
@@ -308,3 +308,12 @@ class MlxModelRunnerStub(ModelRunner):
def alloc_memory_pool(self, memory_pool_config=None):
"""No-op: MLX manages its own KV cache."""
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
+24 -22
View File
@@ -5359,28 +5359,30 @@ 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).
supported_backends = [
"triton",
"trtllm_mha",
"fa3",
"fa4",
"ascend",
"intel_amx",
"intel_xpu",
"aiter",
]
prefill_attn_backend, decode_attn_backend = (
self._resolved_attention_backends()
)
assert (
prefill_attn_backend in supported_backends
and decode_attn_backend in supported_backends
), (
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"
)
# None of these backends exist on MPS; attention_backend is still
# unset there at this point (the torch_native default fills later).
if not is_mps():
supported_backends = [
"triton",
"trtllm_mha",
"fa3",
"fa4",
"ascend",
"intel_amx",
"intel_xpu",
"aiter",
]
prefill_attn_backend, decode_attn_backend = (
self._resolved_attention_backends()
)
assert (
prefill_attn_backend in supported_backends
and decode_attn_backend in supported_backends
), (
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)
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
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
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 stub pulls in ``mlx.core``.
"""
@@ -75,6 +80,31 @@ class TestMlxRunnerPoolContract(unittest.TestCase):
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__":
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()