[MLX] Support Qwen3.5 (dense) Model (#25754)

Signed-off-by: Xiaodong Ye <yeahdongcn@gmail.com>
Co-authored-by: Alex Nails <alex.nails@radixark.ai>
Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
R0CKSTAR
2026-05-30 17:05:02 +08:00
committed by GitHub
co-authored by Alex Nails Claude Opus 4.6
parent 7c5708cba7
commit a952e9174f
23 changed files with 2943 additions and 291 deletions
+1
View File
@@ -1,3 +1,4 @@
[settings] [settings]
profile=black profile=black
known_first_party=sglang known_first_party=sglang
known_third_party=mlx,mlx_lm
+1
View File
@@ -19,6 +19,7 @@ dependencies = ["aiohttp", "requests", "tqdm", "numpy", "IPython", "setproctitle
runtime_common = [ runtime_common = [
"IPython", "IPython",
"aiohttp", "aiohttp",
"apache-tvm-ffi",
"anthropic>=0.20.0", "anthropic>=0.20.0",
"blobfile==3.0.0", "blobfile==3.0.0",
"av", "av",
+1 -1
View File
@@ -536,7 +536,7 @@ class _MlxBenchRunner:
if server_args.max_total_tokens is not None: if server_args.max_total_tokens is not None:
init_kwargs["pool_size"] = server_args.max_total_tokens init_kwargs["pool_size"] = server_args.max_total_tokens
self.mlx_runner = MlxModelRunner(**init_kwargs) self.mlx_runner = MlxModelRunner(**init_kwargs)
self.mlx_runner.init_kv_pool(req_to_token_pool=None) self.mlx_runner.init_cache_pools(req_to_token_pool=None)
self.fake_torch_runner = model_runner self.fake_torch_runner = model_runner
def clear(self): def clear(self):
+3 -2
View File
@@ -164,6 +164,7 @@ from sglang.srt.utils import (
add_prometheus_track_response_middleware, add_prometheus_track_response_middleware,
delete_directory, delete_directory,
get_bool_env_var, get_bool_env_var,
is_mps,
kill_process_tree, kill_process_tree,
set_uvicorn_logging_configs, set_uvicorn_logging_configs,
) )
@@ -1903,8 +1904,8 @@ def _execute_server_warmup(server_args: ServerArgs):
model_info = res.json() model_info = res.json()
# Construct a warmup request # Construct a warmup request (MLX: text warmup for VLM-advertising models; TODO: enable image warmup).
is_vlm = bool(model_info.get("has_image_understanding", False)) is_vlm = bool(model_info.get("has_image_understanding", False)) and not is_mps()
if model_info["is_generation"]: if model_info["is_generation"]:
if is_vlm and not server_args.skip_tokenizer_init: if is_vlm and not server_args.skip_tokenizer_init:
request_name = "/v1/chat/completions" request_name = "/v1/chat/completions"
+11 -4
View File
@@ -13,8 +13,8 @@ from sglang.srt.environ import envs
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.hardware_backend.mlx.kv_cache.contiguous_cache import ( from sglang.srt.hardware_backend.mlx.kv_cache.attention_kv_cache import (
ContiguousKVCache, ContiguousAttentionKVCache,
) )
@@ -123,6 +123,10 @@ class MlxAOTKernelRegistry:
def _build_rope_kernel(inputs: MlxAOTKernelBuildInputs) -> MlxAOTRoPEKernel: def _build_rope_kernel(inputs: MlxAOTKernelBuildInputs) -> MlxAOTRoPEKernel:
from sglang.srt.hardware_backend.mlx.kv_cache.attention_contract import (
get_num_heads,
)
sample_attn = getattr(inputs.sample_attn, "_inner", inputs.sample_attn) sample_attn = getattr(inputs.sample_attn, "_inner", inputs.sample_attn)
rope = getattr(sample_attn, "rope", None) rope = getattr(sample_attn, "rope", None)
if rope is None or getattr(rope, "traditional", False): if rope is None or getattr(rope, "traditional", False):
@@ -136,10 +140,13 @@ def _build_rope_kernel(inputs: MlxAOTKernelBuildInputs) -> MlxAOTRoPEKernel:
return MlxAOTRoPEKernel() return MlxAOTRoPEKernel()
base = float(getattr(rope, "base", 10000.0)) base = float(getattr(rope, "base", 10000.0))
num_qo_heads = get_num_heads(sample_attn)
if num_qo_heads is None:
return MlxAOTRoPEKernel()
config = { config = {
"head_dim": int(inputs.head_dim), "head_dim": int(inputs.head_dim),
"rope_dim": rope_dim, "rope_dim": rope_dim,
"num_qo_heads": int(sample_attn.n_heads), "num_qo_heads": int(num_qo_heads),
"num_kv_heads": int(inputs.n_kv_heads), "num_kv_heads": int(inputs.n_kv_heads),
} }
try: try:
@@ -196,7 +203,7 @@ class MlxAOTKernelContext:
req_ids: list[str], req_ids: list[str],
req_pool_idx: dict[str, int], req_pool_idx: dict[str, int],
req_to_token_pool: Any | None, req_to_token_pool: Any | None,
layer_caches: list[list[ContiguousKVCache]], layer_caches: list[list[ContiguousAttentionKVCache]],
) -> "MlxAOTKernelContext": ) -> "MlxAOTKernelContext":
"""Build optional AOT context for one batched decode step.""" """Build optional AOT context for one batched decode step."""
if not aot_kernels.rope.enabled or kv_pool is None: if not aot_kernels.rope.enabled or kv_pool is None:
@@ -1,5 +1,20 @@
"""KV cache components for the MLX backend.""" """Cache components for the MLX backend."""
from sglang.srt.hardware_backend.mlx.kv_cache.attention_contract import (
get_head_dim,
get_num_heads,
get_num_kv_heads,
is_attention_module,
uses_sliding_window_attention,
)
from sglang.srt.hardware_backend.mlx.kv_cache.attention_kv_cache import (
AttentionOffsetCache,
ContiguousAttentionKVCache,
PoolBackedAttentionKVCache,
)
from sglang.srt.hardware_backend.mlx.kv_cache.attention_kv_pool import (
MlxAttentionKVPool,
)
from sglang.srt.hardware_backend.mlx.kv_cache.attention_wrapper import ( from sglang.srt.hardware_backend.mlx.kv_cache.attention_wrapper import (
BatchedDecodeContext, BatchedDecodeContext,
MLXAttentionWrapper, MLXAttentionWrapper,
@@ -7,12 +22,12 @@ from sglang.srt.hardware_backend.mlx.kv_cache.attention_wrapper import (
get_context, get_context,
set_context, set_context,
) )
from sglang.srt.hardware_backend.mlx.kv_cache.contiguous_cache import ( from sglang.srt.hardware_backend.mlx.kv_cache.auxiliary_state import (
ContiguousKVCache, MlxAuxiliaryStateComponent,
OffsetCache, MlxAuxiliaryStatePool,
PoolBackedCache, MlxAuxiliaryStateReqToTokenPool,
) )
from sglang.srt.hardware_backend.mlx.kv_cache.kv_pool import MlxKVPool from sglang.srt.hardware_backend.mlx.kv_cache.layout import MlxModelCacheLayout
from sglang.srt.hardware_backend.mlx.kv_cache.model_patching import ( from sglang.srt.hardware_backend.mlx.kv_cache.model_patching import (
find_attention_layers, find_attention_layers,
get_num_layers, get_num_layers,
@@ -22,14 +37,23 @@ from sglang.srt.hardware_backend.mlx.kv_cache.model_patching import (
__all__ = [ __all__ = [
"BatchedDecodeContext", "BatchedDecodeContext",
"clear_context", "clear_context",
"ContiguousKVCache", "AttentionOffsetCache",
"ContiguousAttentionKVCache",
"find_attention_layers", "find_attention_layers",
"get_head_dim",
"get_context", "get_context",
"get_num_layers", "get_num_layers",
"get_num_heads",
"get_num_kv_heads",
"is_attention_module",
"MLXAttentionWrapper", "MLXAttentionWrapper",
"MlxKVPool", "MlxAttentionKVPool",
"OffsetCache", "MlxAuxiliaryStateComponent",
"MlxAuxiliaryStatePool",
"MlxAuxiliaryStateReqToTokenPool",
"MlxModelCacheLayout",
"patch_model_attention", "patch_model_attention",
"PoolBackedCache", "PoolBackedAttentionKVCache",
"set_context", "set_context",
"uses_sliding_window_attention",
] ]
@@ -0,0 +1,66 @@
"""Attention helpers based on duck typing for the MLX backend."""
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")
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 = (
"is_sliding",
"use_sliding",
"is_sliding_window",
"use_sliding_window",
"is_swa",
)
def first_present_attr(module: Any, names: Iterable[str]) -> Any | None:
"""Return the first present attribute value without treating 0 as absent."""
for name in names:
if hasattr(module, name):
return getattr(module, name)
return None
def get_num_heads(module: Any) -> int | None:
return first_present_attr(module, NUM_HEAD_ATTRS)
def get_num_kv_heads(module: Any) -> int | None:
return first_present_attr(module, NUM_KV_HEAD_ATTRS)
def get_head_dim(module: Any) -> int | None:
head_dim = first_present_attr(module, ("head_dim",))
if head_dim is not None:
return head_dim
n_kv_heads = get_num_kv_heads(module)
if n_kv_heads is None:
return None
if hasattr(module, "hidden_size") and hasattr(module, "num_k_heads"):
return module.hidden_size // module.num_k_heads
if hasattr(module, "k_proj") and hasattr(module.k_proj, "weight"):
return module.k_proj.weight.shape[0] // n_kv_heads
return None
def is_attention_module(module: Any) -> bool:
return (
all(hasattr(module, attr) for attr in ATTENTION_API_ATTRS)
and get_num_heads(module) is not None
and get_num_kv_heads(module) is not None
)
def uses_sliding_window_attention(*modules: Any) -> bool:
return any(
bool(getattr(module, attr, False))
for module in modules
for attr in SLIDING_ATTENTION_ATTRS
)
@@ -1,4 +1,4 @@
"""ContiguousKVCache, PoolBackedCache and OffsetCache for MLX backend.""" """Attention KV cache adapters for the MLX backend."""
from __future__ import annotations from __future__ import annotations
@@ -7,10 +7,12 @@ from typing import TYPE_CHECKING
import mlx.core as mx import mlx.core as mx
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.hardware_backend.mlx.kv_cache.kv_pool import MlxKVPool from sglang.srt.hardware_backend.mlx.kv_cache.attention_kv_pool import (
MlxAttentionKVPool,
)
class OffsetCache: class AttentionOffsetCache:
"""Data-free shim satisfying mlx-lm's cache protocol. """Data-free shim satisfying mlx-lm's cache protocol.
Provides ``make_mask`` and ``state`` without storing actual K/V. Provides ``make_mask`` and ``state`` without storing actual K/V.
@@ -27,14 +29,14 @@ class OffsetCache:
return None if N == 1 else "causal" return None if N == 1 else "causal"
def update_and_fetch(self, keys, values): def update_and_fetch(self, keys, values):
raise RuntimeError("OffsetCache should not store data") raise RuntimeError("AttentionOffsetCache should not store data")
_DEFAULT_MAX_SEQ_LEN = 4096 _DEFAULT_MAX_SEQ_LEN = 4096
class ContiguousKVCache: class ContiguousAttentionKVCache:
"""Pre-allocated KV buffer for one request × one layer. """Pre-allocated attention KV buffer for one request and one layer.
Shape ``(1, n_kv_heads, max_seq_len, head_dim)``. Slice assignment Shape ``(1, n_kv_heads, max_seq_len, head_dim)``. Slice assignment
instead of ``mx.concatenate``. Lazy-allocated on first write. instead of ``mx.concatenate``. Lazy-allocated on first write.
@@ -119,12 +121,12 @@ class ContiguousKVCache:
return self.keys[:, :, : self.offset, :], self.values[:, :, : self.offset, :] return self.keys[:, :, : self.offset, :], self.values[:, :, : self.offset, :]
class PoolBackedCache: class PoolBackedAttentionKVCache:
"""Lazily gathers cached KV from the shared pool during forward pass. """Lazily gathers cached attention KV from the shared pool during forward.
Each ``update_and_fetch`` gathers this layer's prefix from the pool Each ``update_and_fetch`` gathers this layer's prefix from the pool
on demand, keeping operations in the lazy compute graph. Convert to on demand, keeping operations in the lazy compute graph. Convert to
``ContiguousKVCache`` via ``to_contiguous`` after the forward pass. ``ContiguousAttentionKVCache`` via ``to_contiguous`` after the forward pass.
""" """
__slots__ = ( __slots__ = (
@@ -140,7 +142,7 @@ class PoolBackedCache:
def __init__( def __init__(
self, self,
pool: MlxKVPool, pool: MlxAttentionKVPool,
layer_idx: int, layer_idx: int,
slots: mx.array, slots: mx.array,
prefix_len: int, prefix_len: int,
@@ -197,9 +199,9 @@ class PoolBackedCache:
self._new_values = values self._new_values = values
return k_all, v_all return k_all, v_all
def to_contiguous(self, max_seq_len: int = 4096) -> ContiguousKVCache: def to_contiguous(self, max_seq_len: int = 4096) -> ContiguousAttentionKVCache:
"""Convert to ContiguousKVCache reusing forward-pass arrays.""" """Convert to contiguous attention KV reusing forward-pass arrays."""
cache = ContiguousKVCache(max_seq_len=max_seq_len) cache = ContiguousAttentionKVCache(max_seq_len=max_seq_len)
if self._full_keys is not None: if self._full_keys is not None:
cache.update_and_fetch(self._full_keys, self._full_values) cache.update_and_fetch(self._full_keys, self._full_values)
return cache return cache
@@ -1,4 +1,10 @@
"""Flat KV pool with per-layer buffers of shape (pool_size, n_kv_heads, head_dim). """Flat attention KV pool for the MLX backend.
Each layer buffer has shape ``(pool_size, n_kv_heads, head_dim)``.
This v1 pool is intentionally uniform: every wrapped softmax-attention
layer must share the same KV shape and full-context KV semantics.
Heterogeneous KV shapes and sliding-window KV need per-layer/window-aware
pools before they can use MLX radix reuse.
Slot 0 is reserved as padding (1-based indexing). Slot 0 is reserved as padding (1-based indexing).
""" """
@@ -10,8 +16,8 @@ import mlx.core as mx
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
class MlxKVPool: class MlxAttentionKVPool:
"""Pre-allocated KV pool indexed by integer slot IDs.""" """Pre-allocated attention KV pool indexed by integer slot IDs."""
def __init__( def __init__(
self, self,
@@ -27,7 +33,7 @@ class MlxKVPool:
self.head_dim = head_dim self.head_dim = head_dim
self.dtype = dtype self.dtype = dtype
# Per-layer buffers: (pool_size, n_kv_heads, head_dim) # Per-attention-layer buffers: (pool_size, n_kv_heads, head_dim)
self.k_buffer: list[mx.array] = [ self.k_buffer: list[mx.array] = [
mx.zeros((pool_size, n_kv_heads, head_dim), dtype=dtype) mx.zeros((pool_size, n_kv_heads, head_dim), dtype=dtype)
for _ in range(num_layers) for _ in range(num_layers)
@@ -41,8 +47,8 @@ class MlxKVPool:
1024 * 1024 1024 * 1024
) )
logger.info( logger.info(
f"MlxKVPool: {pool_size} slots × {num_layers} layers " f"MlxAttentionKVPool: {pool_size} slots x {num_layers} layers "
f"× {n_kv_heads} heads × {head_dim} dim, " f"x {n_kv_heads} heads x {head_dim} dim, "
f"dtype={dtype}, ~{mem_mb:.1f} MB" f"dtype={dtype}, ~{mem_mb:.1f} MB"
) )
@@ -14,7 +14,14 @@ from sglang.srt.hardware_backend.mlx.aot import (
MlxAOTKernelSet, MlxAOTKernelSet,
MlxAOTRoPEContext, MlxAOTRoPEContext,
) )
from sglang.srt.hardware_backend.mlx.kv_cache.contiguous_cache import ContiguousKVCache from sglang.srt.hardware_backend.mlx.kv_cache.attention_contract import (
get_head_dim,
get_num_heads,
get_num_kv_heads,
)
from sglang.srt.hardware_backend.mlx.kv_cache.attention_kv_cache import (
ContiguousAttentionKVCache,
)
_thread_local = threading.local() _thread_local = threading.local()
@@ -26,8 +33,9 @@ class BatchedDecodeContext:
batch_size: int batch_size: int
seq_lens: list[int] # per-request token count before the new token seq_lens: list[int] # per-request token count before the new token
# layer_caches[layer_idx][req_idx] = ContiguousKVCache # attention_layer_caches[attention_pool_idx][req_idx] = ContiguousAttentionKVCache
layer_caches: list[list[ContiguousKVCache]] attention_layer_caches: list[list[ContiguousAttentionKVCache]]
attention_pool_index_by_layer: dict[int, int] = field(default_factory=dict)
# Optional AOT kernel state. Keep kernel-specific fields out of the regular # Optional AOT kernel state. Keep kernel-specific fields out of the regular
# MLX decode path so future AOT kernels can be added without growing this # MLX decode path so future AOT kernels can be added without growing this
@@ -51,36 +59,46 @@ class BatchedDecodeContext:
self.needs_padding = min(seq_lens) < max_seq_len self.needs_padding = min(seq_lens) < max_seq_len
self.pad_sizes = [max_seq_len - s for s in seq_lens] self.pad_sizes = [max_seq_len - s for s in seq_lens]
self.positions = mx.arange(self.max_len) if self.needs_padding else None self.positions = mx.arange(self.max_len) if self.needs_padding else None
if not self.attention_pool_index_by_layer:
self.attention_pool_index_by_layer = {
idx: idx for idx in range(len(self.attention_layer_caches))
}
@classmethod @classmethod
def from_decode( def from_decode(
cls, cls,
*, *,
caches: list[list[ContiguousKVCache]], caches: list[list[Any]],
num_layers: int,
req_ids: list[str], req_ids: list[str],
aot_kernels: MlxAOTKernelSet, aot_kernels: MlxAOTKernelSet,
kv_pool: Any | None, kv_pool: Any | None,
req_pool_idx: dict[str, int], req_pool_idx: dict[str, int],
req_to_token_pool: Any | None, req_to_token_pool: Any | None,
attention_layer_indices: list[int] | None = None,
attention_pool_index_by_layer: dict[int, int] | None = None,
) -> "BatchedDecodeContext": ) -> "BatchedDecodeContext":
batch_size = len(req_ids) batch_size = len(req_ids)
seq_lens = [caches[i][0].offset for i in range(batch_size)] if attention_layer_indices is None:
layer_caches = [ attention_layer_indices = list(range(len(caches[0])))
seq_lens = [
caches[i][attention_layer_indices[0]].offset for i in range(batch_size)
]
attention_layer_caches = [
[caches[i][layer_idx] for i in range(batch_size)] [caches[i][layer_idx] for i in range(batch_size)]
for layer_idx in range(num_layers) for layer_idx in attention_layer_indices
] ]
return cls( return cls(
batch_size=batch_size, batch_size=batch_size,
seq_lens=seq_lens, seq_lens=seq_lens,
layer_caches=layer_caches, attention_layer_caches=attention_layer_caches,
attention_pool_index_by_layer=attention_pool_index_by_layer or {},
aot=MlxAOTKernelContext.from_decode( aot=MlxAOTKernelContext.from_decode(
aot_kernels=aot_kernels, aot_kernels=aot_kernels,
kv_pool=kv_pool, kv_pool=kv_pool,
req_ids=req_ids, req_ids=req_ids,
req_pool_idx=req_pool_idx, req_pool_idx=req_pool_idx,
req_to_token_pool=req_to_token_pool, req_to_token_pool=req_to_token_pool,
layer_caches=layer_caches, layer_caches=attention_layer_caches,
), ),
) )
@@ -119,15 +137,38 @@ class MLXAttentionWrapper(nn.Module):
inner = self._inner inner = self._inner
layer_idx = self._layer_idx layer_idx = self._layer_idx
B = ctx.batch_size B = ctx.batch_size
n_heads = get_num_heads(inner)
n_kv_heads = get_num_kv_heads(inner)
if n_heads is None or n_kv_heads is None:
raise RuntimeError(
f"Cannot determine attention head counts for {type(inner).__name__}"
)
queries = inner.q_proj(x) q_proj_output = inner.q_proj(x)
keys = inner.k_proj(x) keys = inner.k_proj(x)
values = inner.v_proj(x) values = inner.v_proj(x)
head_dim = queries.shape[-1] // inner.n_heads head_dim = get_head_dim(inner)
queries = queries.reshape(B, 1, inner.n_heads, head_dim) if head_dim is None:
keys = keys.reshape(B, 1, inner.n_kv_heads, head_dim) head_dim = keys.shape[-1] // n_kv_heads
values = values.reshape(B, 1, inner.n_kv_heads, head_dim)
q_width = n_heads * head_dim
gate = None
if q_proj_output.shape[-1] == q_width:
queries = q_proj_output.reshape(B, 1, n_heads, head_dim)
elif q_proj_output.shape[-1] == 2 * q_width:
queries, gate = mx.split(
q_proj_output.reshape(B, 1, n_heads, 2 * head_dim), 2, axis=-1
)
gate = gate.reshape(B, 1, q_width)
else:
raise RuntimeError(
f"Unexpected q_proj output shape {q_proj_output.shape} for "
f"{type(inner).__name__}"
)
keys = keys.reshape(B, 1, n_kv_heads, head_dim)
values = values.reshape(B, 1, n_kv_heads, head_dim)
if hasattr(inner, "q_norm"): if hasattr(inner, "q_norm"):
queries = inner.q_norm(queries) queries = inner.q_norm(queries)
@@ -140,6 +181,7 @@ class MLXAttentionWrapper(nn.Module):
# Vectorized RoPE with per-batch offsets (cached on the context). # Vectorized RoPE with per-batch offsets (cached on the context).
offsets = ctx.offsets offsets = ctx.offsets
attention_pool_idx = ctx.attention_pool_index_by_layer[layer_idx]
if ctx.aot.rope is not None: if ctx.aot.rope is not None:
# AOT path: real .metallib RoPE + fused KV pool scatter. # AOT path: real .metallib RoPE + fused KV pool scatter.
@@ -148,7 +190,7 @@ class MLXAttentionWrapper(nn.Module):
keys, keys,
values, values,
offsets, offsets,
layer_idx, attention_pool_idx,
ctx.aot.rope, ctx.aot.rope,
) )
else: else:
@@ -157,7 +199,7 @@ class MLXAttentionWrapper(nn.Module):
queries = inner.rope(queries, offset=offsets) queries = inner.rope(queries, offset=offsets)
keys = inner.rope(keys, offset=offsets) keys = inner.rope(keys, offset=offsets)
layer_caches = ctx.layer_caches[layer_idx] layer_caches = ctx.attention_layer_caches[attention_pool_idx]
pad_sizes = ctx.pad_sizes pad_sizes = ctx.pad_sizes
# TODO: replace per-request loop with native batched/ragged # TODO: replace per-request loop with native batched/ragged
@@ -173,12 +215,8 @@ class MLXAttentionWrapper(nn.Module):
pad = pad_sizes[i] pad = pad_sizes[i]
if pad > 0: if pad > 0:
k_pad = mx.zeros( k_pad = mx.zeros((1, n_kv_heads, pad, head_dim), dtype=k_all.dtype)
(1, inner.n_kv_heads, pad, head_dim), dtype=k_all.dtype v_pad = mx.zeros((1, n_kv_heads, pad, head_dim), dtype=v_all.dtype)
)
v_pad = mx.zeros(
(1, inner.n_kv_heads, pad, head_dim), dtype=v_all.dtype
)
k_all = mx.concatenate([k_all, k_pad], axis=2) k_all = mx.concatenate([k_all, k_pad], axis=2)
v_all = mx.concatenate([v_all, v_pad], axis=2) v_all = mx.concatenate([v_all, v_pad], axis=2)
@@ -202,6 +240,8 @@ class MLXAttentionWrapper(nn.Module):
) )
output = output.transpose(0, 2, 1, 3).reshape(B, 1, -1) output = output.transpose(0, 2, 1, 3).reshape(B, 1, -1)
if gate is not None:
output = output * mx.sigmoid(gate)
return inner.o_proj(output) return inner.o_proj(output)
@staticmethod @staticmethod
@@ -210,7 +250,7 @@ class MLXAttentionWrapper(nn.Module):
keys: mx.array, keys: mx.array,
values: mx.array, values: mx.array,
positions: mx.array, positions: mx.array,
layer_idx: int, attention_pool_idx: int,
rope_ctx: MlxAOTRoPEContext, rope_ctx: MlxAOTRoPEContext,
) -> tuple[mx.array, mx.array]: ) -> tuple[mx.array, mx.array]:
"""AOT path: rotate Q/K and scatter K/V into the shared pool. """AOT path: rotate Q/K and scatter K/V into the shared pool.
@@ -234,8 +274,8 @@ class MLXAttentionWrapper(nn.Module):
else: else:
slots = rope_ctx.new_token_slots.astype(mx.int32) slots = rope_ctx.new_token_slots.astype(mx.int32)
k_pool = rope_ctx.kv_pool.k_buffer[layer_idx] k_pool = rope_ctx.kv_pool.k_buffer[attention_pool_idx]
v_pool = rope_ctx.kv_pool.v_buffer[layer_idx] v_pool = rope_ctx.kv_pool.v_buffer[attention_pool_idx]
q_rot, k_rot, k_pool_new, v_pool_new = rope_ctx.kernel.rope_pool_fused( q_rot, k_rot, k_pool_new, v_pool_new = rope_ctx.kernel.rope_pool_fused(
q_flat, q_flat,
@@ -251,8 +291,8 @@ class MLXAttentionWrapper(nn.Module):
rope_base=rope_ctx.kernel.base, rope_base=rope_ctx.kernel.base,
) )
# Rebind pool buffers (zero-copy donation result). # Rebind pool buffers (zero-copy donation result).
rope_ctx.kv_pool.k_buffer[layer_idx] = k_pool_new rope_ctx.kv_pool.k_buffer[attention_pool_idx] = k_pool_new
rope_ctx.kv_pool.v_buffer[layer_idx] = v_pool_new rope_ctx.kv_pool.v_buffer[attention_pool_idx] = v_pool_new
# (B, n_heads, head_dim) -> (B, n_heads, 1, head_dim) for SDPA path # (B, n_heads, head_dim) -> (B, n_heads, 1, head_dim) for SDPA path
return q_rot[:, :, None, :], k_rot[:, :, None, :] return q_rot[:, :, None, :], k_rot[:, :, None, :]
@@ -0,0 +1,390 @@
"""MLX auxiliary-state snapshots for unified radix cache.
Hybrid MLX models may include non-softmax-attention layers whose native
``mlx-lm`` cache state cannot be reconstructed from the attention KV pool.
The global scheduler exposes that state through its existing MAMBA component
contract, so this MLX adapter keeps those scheduler-facing field names while
storing model-agnostic native cache snapshots.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Iterable, Optional
import mlx.core as mx
import torch
from sglang.srt.mem_cache.base_prefix_cache import EvictParams, InsertResult
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.mem_cache.unified_cache_components.mamba_component import (
MambaComponent,
)
from sglang.srt.mem_cache.unified_cache_components.tree_component import TreeComponent
_CACHE_ATTRS = ("offset", "lengths", "left_padding")
_MISSING = object()
def _clone_tree(value: Any) -> Any:
if isinstance(value, mx.array):
return mx.array(value)
if isinstance(value, list):
return [_clone_tree(item) for item in value]
if isinstance(value, tuple):
return tuple(_clone_tree(item) for item in value)
if isinstance(value, dict):
return {key: _clone_tree(item) for key, item in value.items()}
return value
def _arrays_in_tree(value: Any) -> list[mx.array]:
arrays: list[mx.array] = []
def collect(item: Any) -> None:
if isinstance(item, mx.array):
arrays.append(item)
elif isinstance(item, (list, tuple)):
for child in item:
collect(child)
elif isinstance(item, dict):
for child in item.values():
collect(child)
collect(value)
return arrays
@dataclass
class _CacheSnapshot:
state: Any
meta_state: Any = _MISSING
attrs: dict[str, Any] | None = None
def _snapshot_cache(cache: Any) -> _CacheSnapshot:
state = _clone_tree(getattr(cache, "state", ()))
meta_state = (
_clone_tree(cache.meta_state) if hasattr(cache, "meta_state") else _MISSING
)
attrs = {
name: _clone_tree(getattr(cache, name))
for name in _CACHE_ATTRS
if hasattr(cache, name)
}
arrays = _arrays_in_tree((state, meta_state, attrs))
if arrays:
mx.eval(*arrays)
return _CacheSnapshot(state=state, meta_state=meta_state, attrs=attrs)
def _restore_cache(cache: Any, snapshot: _CacheSnapshot) -> None:
cache.state = _clone_tree(snapshot.state)
if snapshot.meta_state is not _MISSING and hasattr(cache, "meta_state"):
cache.meta_state = _clone_tree(snapshot.meta_state)
for name, value in (snapshot.attrs or {}).items():
setattr(cache, name, _clone_tree(value))
class MlxAuxiliaryStatePool:
"""Index-addressable snapshots of native MLX auxiliary cache state."""
def __init__(self, size: int, device: str):
self.size = size
self.device = device
self.mamba_cache = None
self.mem_usage = 0
self._snapshots: dict[int, dict[int, _CacheSnapshot]] = {}
self.clear()
def _tensor(self, indices: Any) -> torch.Tensor:
return torch.as_tensor(indices, dtype=torch.int64, device=self.device).view(-1)
def _index(self, index: Any) -> int:
flat = self._tensor(index)
assert flat.numel() == 1
return int(flat.item())
def available_size(self) -> int:
return int(self.free_slots.numel())
def alloc(self, need_size: int) -> Optional[torch.Tensor]:
if need_size > self.available_size():
return None
slots = self.free_slots[:need_size].clone()
self.free_slots = self.free_slots[need_size:]
for slot in slots.tolist():
self._snapshots.pop(int(slot), None)
return slots
def free(self, indices: Any) -> None:
if indices is None:
return
indices = self._tensor(indices)
if indices.numel() == 0:
return
for slot in indices.tolist():
self._snapshots.pop(int(slot), None)
self.free_slots = torch.cat([self.free_slots, indices])
def clear(self) -> None:
self.free_slots = torch.arange(
1, self.size + 1, dtype=torch.int64, device=self.device
)
self._snapshots.clear()
def copy_from(self, src: Any, dst: Any) -> None:
src_indices = self._tensor(src)
dst_indices = self._tensor(dst)
assert src_indices.numel() == dst_indices.numel()
for src_idx, dst_idx in zip(src_indices.tolist(), dst_indices.tolist()):
snapshot = self._snapshots.get(int(src_idx))
if snapshot is None:
self._snapshots.pop(int(dst_idx), None)
else:
self._snapshots[int(dst_idx)] = {
layer_idx: _CacheSnapshot(
state=_clone_tree(cache_snapshot.state),
meta_state=_clone_tree(cache_snapshot.meta_state),
attrs=_clone_tree(cache_snapshot.attrs),
)
for layer_idx, cache_snapshot in snapshot.items()
}
def fork_from(self, src: Any) -> Optional[torch.Tensor]:
src_indices = self._tensor(src)
dst = self.alloc(src_indices.numel())
if dst is None:
return None
self.copy_from(src_indices, dst)
return dst
def store_cache(
self,
index: Any,
cache: list[Any],
layer_indices: Iterable[int],
) -> None:
self._snapshots[self._index(index)] = {
layer_idx: _snapshot_cache(cache[layer_idx]) for layer_idx in layer_indices
}
def restore_cache(
self,
index: Any,
cache: list[Any],
layer_indices: Iterable[int] | None = None,
) -> bool:
snapshot = self._snapshots.get(self._index(index))
if snapshot is None:
return False
selected_layers = set(layer_indices) if layer_indices is not None else None
for layer_idx, cache_snapshot in snapshot.items():
if selected_layers is not None and layer_idx not in selected_layers:
continue
_restore_cache(cache[layer_idx], cache_snapshot)
return True
def has_snapshot(self, index: Any) -> bool:
return self._index(index) in self._snapshots
class MlxAuxiliaryStateReqToTokenPool(ReqToTokenPool):
"""Req-to-token pool with MLX auxiliary-state slot bookkeeping."""
def __init__(
self,
*,
size: int,
max_context_len: int,
device: str,
enable_memory_saver: bool,
auxiliary_state_size: int,
):
super().__init__(
size=size,
max_context_len=max_context_len,
device=device,
enable_memory_saver=enable_memory_saver,
)
self.mamba_pool = MlxAuxiliaryStatePool(
size=auxiliary_state_size,
device=device,
)
# The unified radix base MAMBA component still reads ``mamba_pool``.
# Keep the MLX-owned name beside it so local code can avoid model-
# specific terminology.
self.auxiliary_state_pool = self.mamba_pool
self.enable_mamba_extra_buffer = False
self.req_index_to_auxiliary_state_index_mapping = torch.zeros(
self._alloc_size, dtype=torch.int32, device=device
)
def alloc(self, reqs):
select_index = super().alloc(reqs)
if select_index is None:
return None
auxiliary_state_indices = []
for req in reqs:
if getattr(req, "mamba_pool_idx", None) is not None:
mid = req.mamba_pool_idx
else:
allocated = self.auxiliary_state_pool.alloc(1)
assert allocated is not None, "Not enough MLX auxiliary state slots"
mid = allocated[0]
req.mamba_pool_idx = mid
auxiliary_state_indices.append(mid.to(dtype=torch.int32))
self.req_index_to_auxiliary_state_index_mapping[select_index] = torch.stack(
auxiliary_state_indices
)
return select_index
def get_auxiliary_state_indices(self, req_indices) -> torch.Tensor:
return self.req_index_to_auxiliary_state_index_mapping[req_indices]
def get_mamba_indices(self, req_indices) -> torch.Tensor:
return self.get_auxiliary_state_indices(req_indices)
def get_mamba_ping_pong_other_idx(self, mamba_next_track_idx: int) -> int:
return 0
def free_mamba_cache(self, req, mamba_ping_pong_track_buffer_to_keep=None):
if getattr(req, "mamba_pool_idx", None) is not None:
self.auxiliary_state_pool.free(req.mamba_pool_idx.unsqueeze(0))
req.mamba_pool_idx = None
track_buffer = getattr(req, "mamba_ping_pong_track_buffer", None)
if track_buffer is not None:
if mamba_ping_pong_track_buffer_to_keep is None:
self.auxiliary_state_pool.free(track_buffer)
req.mamba_ping_pong_track_buffer = None
req.mamba_next_track_idx = None
def free_auxiliary_state_cache(self, req, track_buffer_to_keep=None):
self.free_mamba_cache(
req,
mamba_ping_pong_track_buffer_to_keep=track_buffer_to_keep,
)
def free(self, req):
super().free(req)
def clear(self):
super().clear()
self.auxiliary_state_pool.clear()
self.req_index_to_auxiliary_state_index_mapping.zero_()
class MlxAuxiliaryStateComponent(MambaComponent):
"""Unified radix component for MLX native auxiliary-state snapshots."""
def __init__(self, cache, params):
if params.enable_mamba_extra_buffer:
raise NotImplementedError(
"MLX auxiliary-state radix cache does not support "
"enable_mamba_extra_buffer yet."
)
pool = getattr(cache.req_to_token_pool, "auxiliary_state_pool", None)
if not isinstance(pool, MlxAuxiliaryStatePool):
raise TypeError(
"MlxAuxiliaryStateComponent requires MlxAuxiliaryStatePool, "
f"got {type(pool)}"
)
TreeComponent.__init__(self, cache, params)
self.enable_mamba_extra_buffer = False
self._mamba_pool_host = None
@staticmethod
def _tracked_value(req) -> tuple[object | None, bool]:
track_buffer = getattr(req, "mamba_ping_pong_track_buffer", None)
track_len = getattr(req, "mamba_last_track_seqlen", None)
if track_buffer is not None and track_len is not None:
return track_buffer[0].unsqueeze(-1).clone(), True
if getattr(req, "mamba_pool_idx", None) is None:
return None, False
return req.mamba_pool_idx.unsqueeze(-1).clone(), False
def prepare_for_caching_req(
self,
req,
insert_params,
token_ids_len: int,
is_finished: bool,
) -> int | None:
cache_len = getattr(req, "mamba_last_track_seqlen", None)
auxiliary_value, uses_track_slot = self._tracked_value(req)
setattr(insert_params, "mlx_auxiliary_state_uses_track_slot", uses_track_slot)
if auxiliary_value is None:
return 0 if is_finished else None
if cache_len is None:
cache_len = token_ids_len
if is_finished:
insert_params.mamba_value = auxiliary_value
else:
source_value = auxiliary_value
forked_value = self.cache.req_to_token_pool.auxiliary_state_pool.fork_from(
source_value
)
if forked_value is None:
self.cache.evict(EvictParams(num_tokens=0, mamba_num=1))
forked_value = (
self.cache.req_to_token_pool.auxiliary_state_pool.fork_from(
source_value
)
)
assert forked_value is not None, "Can not alloc MLX auxiliary cache"
insert_params.mamba_value = forked_value
return cache_len
def cleanup_after_caching_req(
self,
req,
is_finished: bool,
insert_result: InsertResult | None = None,
insert_params=None,
) -> None:
if not is_finished:
if (
insert_params is not None
and insert_params.mamba_value is not None
and (insert_result is None or insert_result.mamba_exist)
):
self.cache.req_to_token_pool.auxiliary_state_pool.free(
insert_params.mamba_value
)
if bool(
getattr(insert_params, "mlx_auxiliary_state_uses_track_slot", False)
):
track_buffer = getattr(req, "mamba_ping_pong_track_buffer", None)
if track_buffer is not None:
self.cache.req_to_token_pool.auxiliary_state_pool.free(track_buffer)
req.mamba_ping_pong_track_buffer = None
req.mamba_next_track_idx = None
req.mamba_last_track_seqlen = None
return
auxiliary_value_exists = (
insert_result.mamba_exist if insert_result is not None else True
)
uses_track_slot = bool(
getattr(insert_params, "mlx_auxiliary_state_uses_track_slot", False)
)
if uses_track_slot:
keep_track_slot = not auxiliary_value_exists
self.cache.req_to_token_pool.free_auxiliary_state_cache(
req,
track_buffer_to_keep=0 if keep_track_slot else None,
)
elif auxiliary_value_exists:
self.cache.req_to_token_pool.free_auxiliary_state_cache(req)
else:
# The radix tree now owns the live auxiliary-state slot.
track_buffer = getattr(req, "mamba_ping_pong_track_buffer", None)
if track_buffer is not None:
self.cache.req_to_token_pool.auxiliary_state_pool.free(track_buffer)
req.mamba_ping_pong_track_buffer = None
req.mamba_next_track_idx = None
req.mamba_pool_idx = None
req.mamba_last_track_seqlen = None
@@ -0,0 +1,93 @@
"""Model cache layout helpers for the MLX backend."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Sequence
@dataclass(frozen=True)
class MlxModelCacheLayout:
"""Map model layers to MLX cache storage components.
Attention layers store softmax-attention KV in the MLX attention KV pool.
Auxiliary layers keep native ``mlx-lm`` cache state and are snapshotted by
the MLX auxiliary-state component.
"""
layers: tuple[Any, ...]
attention_attrs: tuple[str | None, ...]
attention_layer_indices: tuple[int, ...]
auxiliary_layer_indices: tuple[int, ...]
attention_pool_index_by_layer: dict[int, int]
@classmethod
def from_attention_discovery(
cls,
layers: Sequence[Any],
attention_attrs: Sequence[str | None],
) -> "MlxModelCacheLayout":
if len(layers) != len(attention_attrs):
raise ValueError(
"Layer count and attention attribute count differ: "
f"{len(layers)} != {len(attention_attrs)}"
)
attention_layer_indices = tuple(
idx for idx, attr in enumerate(attention_attrs) if attr is not None
)
auxiliary_layer_indices = tuple(
idx for idx, attr in enumerate(attention_attrs) if attr is None
)
attention_pool_index_by_layer = {
layer_idx: pool_idx
for pool_idx, layer_idx in enumerate(attention_layer_indices)
}
return cls(
layers=tuple(layers),
attention_attrs=tuple(attention_attrs),
attention_layer_indices=attention_layer_indices,
auxiliary_layer_indices=auxiliary_layer_indices,
attention_pool_index_by_layer=attention_pool_index_by_layer,
)
@property
def num_layers(self) -> int:
return len(self.layers)
@property
def num_attention_layers(self) -> int:
return len(self.attention_layer_indices)
@property
def has_auxiliary_state(self) -> bool:
return bool(self.auxiliary_layer_indices)
@property
def first_attention_layer_index(self) -> int:
if not self.attention_layer_indices:
raise RuntimeError("MLX model has no supported attention layers")
return self.attention_layer_indices[0]
def attention_pool_index(self, layer_idx: int) -> int:
try:
return self.attention_pool_index_by_layer[layer_idx]
except KeyError as exc:
raise KeyError(f"Layer {layer_idx} is not an attention layer") from exc
def attention_attr(self, layer_idx: int) -> str:
attr = self.attention_attrs[layer_idx]
if attr is None:
raise KeyError(f"Layer {layer_idx} is not an attention layer")
return attr
def attention_layer_caches(
self,
caches_by_request: list[list[Any]],
) -> list[list[Any]]:
"""Return layer-major attention caches for batched decode."""
return [
[request_cache[layer_idx] for request_cache in caches_by_request]
for layer_idx in self.attention_layer_indices
]
@@ -2,25 +2,38 @@
from typing import Any from typing import Any
import mlx.nn as nn
from sglang.srt.hardware_backend.mlx.kv_cache.attention_contract import (
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,
) )
def find_attention_layers(model: Any) -> tuple[list[Any], str]: def _find_attention_attr(layer: Any) -> str | None:
"""Find transformer layers and the attention attribute name.""" """Return the direct child name that satisfies the attention contract."""
if not isinstance(layer, nn.Module):
raise TypeError(f"Expected mlx.nn.Module layer, got {type(layer)}")
for name, module in layer.children().items():
if isinstance(module, MLXAttentionWrapper) or is_attention_module(module):
return name
return None
def find_attention_layers(model: Any) -> tuple[list[Any], list[str | None]]:
"""Find transformer layers and per-layer attention attribute names."""
root = getattr(model, "language_model", model) root = getattr(model, "language_model", model)
container = getattr(root, "model", root) container = getattr(root, "model", root)
layer_list = getattr(container, "layers", None) or getattr(root, "layers", []) layer_list = getattr(container, "layers", None) or getattr(root, "layers", [])
if layer_list: if layer_list:
sample = layer_list[0] attn_attrs = [_find_attention_attr(layer) for layer in layer_list]
if hasattr(sample, "self_attn"): if any(attr is not None for attr in attn_attrs):
return layer_list, "self_attn" return layer_list, attn_attrs
if hasattr(sample, "attention"): raise ValueError(f"No attention attribute in layer type {type(layer_list[0])}")
return layer_list, "attention" return layer_list, []
raise ValueError(f"No attention attribute in layer type {type(sample)}")
return layer_list, "self_attn"
def patch_model_attention(model: Any) -> int: def patch_model_attention(model: Any) -> int:
@@ -29,9 +42,11 @@ def patch_model_attention(model: Any) -> int:
The wrapper delegates to the inner module when no BatchedDecodeContext The wrapper delegates to the inner module when no BatchedDecodeContext
is set, so it is always installed and never removed. is set, so it is always installed and never removed.
""" """
layer_list, attn_attr = find_attention_layers(model) layer_list, attn_attrs = find_attention_layers(model)
patched = 0 patched = 0
for idx, layer in enumerate(layer_list): for idx, (layer, attn_attr) in enumerate(zip(layer_list, attn_attrs)):
if attn_attr is None:
continue
attn = getattr(layer, attn_attr) attn = getattr(layer, attn_attr)
if isinstance(attn, MLXAttentionWrapper): if isinstance(attn, MLXAttentionWrapper):
continue continue
File diff suppressed because it is too large Load Diff
@@ -9,6 +9,9 @@ from typing import Tuple
import torch import torch
from sglang.srt.hardware_backend.mlx.kv_cache.auxiliary_state import (
MlxAuxiliaryStateReqToTokenPool,
)
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
from sglang.srt.mem_cache.memory_pool import KVCache, ReqToTokenPool from sglang.srt.mem_cache.memory_pool import KVCache, ReqToTokenPool
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
@@ -17,11 +20,11 @@ logger = logging.getLogger(__name__)
class _DummyKVCache(KVCache): class _DummyKVCache(KVCache):
"""A KV cache that allocates no GPU memory. """Scheduler-facing KV cache that allocates no GPU memory.
Satisfies the KVCache interface so that TokenToKVPoolAllocator can be Satisfies the KVCache interface so that TokenToKVPoolAllocator can be
constructed, but every buffer access raises — the MLX backend manages constructed, but every buffer access raises. The MLX backend manages
its own KV cache internally. attention KV and auxiliary state internally.
""" """
def __init__(self, size: int, dtype: torch.dtype, device: str): def __init__(self, size: int, dtype: torch.dtype, device: str):
@@ -42,16 +45,16 @@ class _DummyKVCache(KVCache):
self.custom_mem_pool = None self.custom_mem_pool = None
def get_key_buffer(self, layer_id: int) -> torch.Tensor: def get_key_buffer(self, layer_id: int) -> torch.Tensor:
raise RuntimeError("_DummyKVCache has no key buffer (MLX manages KV cache)") raise RuntimeError("_DummyKVCache has no key buffer (MLX manages cache)")
def get_value_buffer(self, layer_id: int) -> torch.Tensor: def get_value_buffer(self, layer_id: int) -> torch.Tensor:
raise RuntimeError("_DummyKVCache has no value buffer (MLX manages KV cache)") raise RuntimeError("_DummyKVCache has no value buffer (MLX manages cache)")
def get_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]: def get_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
raise RuntimeError("_DummyKVCache has no kv buffer (MLX manages KV cache)") raise RuntimeError("_DummyKVCache has no kv buffer (MLX manages cache)")
def set_kv_buffer(self, layer, loc, cache_k, cache_v) -> None: def set_kv_buffer(self, layer, loc, cache_k, cache_v) -> None:
raise RuntimeError("_DummyKVCache cannot set kv buffer (MLX manages KV cache)") raise RuntimeError("_DummyKVCache cannot set kv buffer (MLX manages cache)")
def get_kv_size_bytes(self): def get_kv_size_bytes(self):
return 0, 0 return 0, 0
@@ -141,6 +144,18 @@ class MlxModelRunnerStub(ModelRunner):
self.is_hybrid_swa = False self.is_hybrid_swa = False
# Create minimal pools # Create minimal pools
if self.mambaish_config is not None:
auxiliary_state_size = self.server_args.max_mamba_cache_size
if auxiliary_state_size is None:
auxiliary_state_size = self.max_running_requests * 4
self.req_to_token_pool = MlxAuxiliaryStateReqToTokenPool(
size=self.max_running_requests,
max_context_len=self.model_config.context_len,
device="cpu",
enable_memory_saver=False,
auxiliary_state_size=auxiliary_state_size,
)
else:
self.req_to_token_pool = ReqToTokenPool( self.req_to_token_pool = ReqToTokenPool(
size=self.max_running_requests, size=self.max_running_requests,
max_context_len=self.model_config.context_len, max_context_len=self.model_config.context_len,
@@ -6,11 +6,11 @@ the scheduler runs its CPU-side bookkeeping on the tokens of the
older one. The lazy-graph primitives live in older one. The lazy-graph primitives live in
``hardware_backend/mlx/tp_worker.py`` and ``model_runner.py``. ``hardware_backend/mlx/tp_worker.py`` and ``model_runner.py``.
Each request's KV lives ina set of per-request, per-layer ``ContiguousKVCache`` Each request's attention KV lives in per-request, per-layer
objects that the ``MLXAttentionWrapper`` mutates in place during the forward pass. ``ContiguousAttentionKVCache`` objects that ``MLXAttentionWrapper`` mutates
Chained decodes reuse the same cache objects: step N+1's graph reads in place during the forward pass. Chained decodes reuse the same cache objects:
step N's lazy writes via MLX's dependency tracking, so the GPU runs step N+1's graph reads step N's lazy writes via MLX's dependency tracking, so
both steps back-to-back with no idle gap. the GPU runs both steps back-to-back with no idle gap.
""" """
from __future__ import annotations from __future__ import annotations
@@ -62,6 +62,9 @@ class MlxPendingJob:
time. Decoupled from the live batch so time. Decoupled from the live batch so
``process_batch_result`` can update request state without ``process_batch_result`` can update request state without
racing against the next scheduling decision. racing against the next scheduling decision.
schedule_batch: The full scheduler batch. Unlike ``batch_copy``,
this keeps allocator/cache fields needed when a prefill batch
becomes the next running decode batch.
reqs: Snapshot of ``batch.reqs`` at launch time. The overlap reqs: Snapshot of ``batch.reqs`` at launch time. The overlap
loop uses this to check ``req.finished()`` on the previous loop uses this to check ``req.finished()`` on the previous
step's request list without holding a reference to the step's request list without holding a reference to the
@@ -74,12 +77,27 @@ class MlxPendingJob:
decode: Optional["MlxPendingDecode"] decode: Optional["MlxPendingDecode"]
mode: str mode: str
batch_copy: "ScheduleBatch" batch_copy: "ScheduleBatch"
schedule_batch: "ScheduleBatch"
reqs: List[Req] reqs: List[Req]
class SchedulerMlxOverlapMixin: class SchedulerMlxOverlapMixin:
"""Mixin that adds MLX overlap scheduling to :class:`Scheduler`.""" """Mixin that adds MLX overlap scheduling to :class:`Scheduler`."""
def _finalize_mlx_pending_job(self: "Scheduler", pending: MlxPendingJob):
result = self.tp_worker.finalize_mlx_result(
pending.prefills,
pending.extends,
pending.decode,
pending.mode,
pending.reqs,
)
if result.next_token_ids is not None:
pending.batch_copy.input_ids = result.next_token_ids
pending.schedule_batch.input_ids = result.next_token_ids
self.last_batch = pending.schedule_batch
self.process_batch_result(pending.batch_copy, result)
@DynamicGradMode() @DynamicGradMode()
def event_loop_overlap_mlx(self: "Scheduler"): def event_loop_overlap_mlx(self: "Scheduler"):
"""MLX-specific overlap loop modelled on ``mlx_lm.generate.generate_step``. """MLX-specific overlap loop modelled on ``mlx_lm.generate.generate_step``.
@@ -123,18 +141,6 @@ class SchedulerMlxOverlapMixin:
pending_curr: Optional[MlxPendingJob] = None pending_curr: Optional[MlxPendingJob] = None
pending_next: Optional[MlxPendingJob] = None pending_next: Optional[MlxPendingJob] = None
def _finalize(pending: MlxPendingJob):
result = self.tp_worker.finalize_mlx_result(
pending.prefills,
pending.extends,
pending.decode,
pending.mode,
pending.reqs,
)
if result.next_token_ids is not None:
pending.batch_copy.input_ids = result.next_token_ids
self.process_batch_result(pending.batch_copy, result)
def _launch_fresh(batch: "ScheduleBatch") -> MlxPendingJob: def _launch_fresh(batch: "ScheduleBatch") -> MlxPendingJob:
lazy_tokens, prefills, extends, decode, mode = ( lazy_tokens, prefills, extends, decode, mode = (
self.tp_worker.async_forward_batch_generation_mlx(batch) self.tp_worker.async_forward_batch_generation_mlx(batch)
@@ -146,6 +152,7 @@ class SchedulerMlxOverlapMixin:
decode=decode, decode=decode,
mode=mode, mode=mode,
batch_copy=batch.copy(), batch_copy=batch.copy(),
schedule_batch=batch,
reqs=list(batch.reqs), reqs=list(batch.reqs),
) )
@@ -164,6 +171,7 @@ class SchedulerMlxOverlapMixin:
decode=decode, decode=decode,
mode=mode, mode=mode,
batch_copy=prev.batch_copy.copy(), batch_copy=prev.batch_copy.copy(),
schedule_batch=prev.schedule_batch,
reqs=prev.reqs, reqs=prev.reqs,
) )
@@ -191,7 +199,7 @@ class SchedulerMlxOverlapMixin:
# 2. Finalize/process on pending_curr's tokens. (GPU is already # 2. Finalize/process on pending_curr's tokens. (GPU is already
# executing pending_next at this point.) # executing pending_next at this point.)
if pending_curr is not None: if pending_curr is not None:
_finalize(pending_curr) self._finalize_mlx_pending_job(pending_curr)
self.result_queue.popleft() self.result_queue.popleft()
pending_curr = None pending_curr = None
@@ -208,8 +216,8 @@ class SchedulerMlxOverlapMixin:
): ):
pending_curr = pending_next pending_curr = pending_next
pending_next = None pending_next = None
self.cur_batch = pending_curr.batch_copy self.cur_batch = pending_curr.schedule_batch
self.last_batch = pending_curr.batch_copy self.last_batch = pending_curr.schedule_batch
if envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.get(): if envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.get():
self.invariant_checker.self_check_during_busy() self.invariant_checker.self_check_during_busy()
continue continue
@@ -217,7 +225,7 @@ class SchedulerMlxOverlapMixin:
# 4. Chain is broken. Finalise pending_next (if any), then # 4. Chain is broken. Finalise pending_next (if any), then
# schedule fresh. # schedule fresh.
if pending_next is not None: if pending_next is not None:
_finalize(pending_next) self._finalize_mlx_pending_job(pending_next)
self.result_queue.popleft() self.result_queue.popleft()
pending_next = None pending_next = None
next_batch = self.get_next_batch_to_run() next_batch = self.get_next_batch_to_run()
@@ -2,7 +2,7 @@
Routes forward passes through the MLX model runner, bypassing PyTorch Routes forward passes through the MLX model runner, bypassing PyTorch
MPS. A lightweight stub provides scheduler bookkeeping; the actual MPS. A lightweight stub provides scheduler bookkeeping; the actual
KV data lives in MlxKVPool. attention KV data lives in MlxAttentionKVPool.
The worker also exposes an async (lazy-eval) surface used by the MLX The worker also exposes an async (lazy-eval) surface used by the MLX
overlap scheduler: ``async_forward_batch_generation_mlx`` launches a overlap scheduler: ``async_forward_batch_generation_mlx`` launches a
@@ -87,9 +87,9 @@ class MlxTpModelWorker(TpModelWorker):
return None return None
def _ensure_mlx_pool_initialized(self): def _ensure_mlx_pool_initialized(self):
"""Lazily initialize the MlxKVPool after the stub's pools are ready.""" """Lazily initialize MLX cache pools after the stub pools are ready."""
if not self._mlx_pool_initialized: if not self._mlx_pool_initialized:
self._mlx_runner.init_kv_pool(self._model_runner.req_to_token_pool) self._mlx_runner.init_cache_pools(self._model_runner.req_to_token_pool)
self._mlx_pool_initialized = True self._mlx_pool_initialized = True
def forward_batch_generation( def forward_batch_generation(
@@ -124,6 +124,14 @@ class MlxTpModelWorker(TpModelWorker):
else: else:
self._mlx_active_rids |= current_rids self._mlx_active_rids |= current_rids
def prepare_for_kv_cache_release(self, req) -> None:
"""Snapshot MLX auxiliary state at the scheduler's radix insert point."""
if self._mlx_runner.has_request(req.rid):
self._mlx_runner.store_auxiliary_state_for_request(req.rid)
# Prefer the just-snapshotted live auxiliary state for the final
# insert. Any older tracked slot is released during component cleanup.
req.mamba_last_track_seqlen = None
def _forward_batch_generation_mlx( def _forward_batch_generation_mlx(
self, batch: ScheduleBatch self, batch: ScheduleBatch
) -> GenerationBatchResult: ) -> GenerationBatchResult:
@@ -144,7 +152,7 @@ class MlxTpModelWorker(TpModelWorker):
next_token_ids_list: list[int] = [] next_token_ids_list: list[int] = []
if forward_mode.is_extend(): if forward_mode.is_extend():
# Ensure pool is up-to-date before PoolBackedCache reads it # Ensure pool is up-to-date before pool-backed attention reads it
# for prefix-cached prefills. Only runs on extend batches. # for prefix-cached prefills. Only runs on extend batches.
self._mlx_runner.flush_all_decode_kv() self._mlx_runner.flush_all_decode_kv()
input_ids_cpu = batch.input_ids.cpu().tolist() input_ids_cpu = batch.input_ids.cpu().tolist()
@@ -185,6 +193,7 @@ class MlxTpModelWorker(TpModelWorker):
prefix_slot_ids=prefix_slot_ids, prefix_slot_ids=prefix_slot_ids,
new_slot_ids=req_new_slots, new_slot_ids=req_new_slots,
req_pool_idx=req.req_pool_idx, req_pool_idx=req.req_pool_idx,
req=req,
) )
prefill_rids.append((req.rid, next_token)) prefill_rids.append((req.rid, next_token))
@@ -270,7 +279,7 @@ class MlxTpModelWorker(TpModelWorker):
if forward_mode.is_extend(): if forward_mode.is_extend():
# TODO (changminbark): Implement per-batch flushing using prefix_slot_ids # TODO (changminbark): Implement per-batch flushing using prefix_slot_ids
# Ensure the pool is up-to-date before any PoolBackedCache # Ensure the pool is up-to-date before pool-backed attention
# reads it for prefix-cached prefills. Mirror the sync path. # reads it for prefix-cached prefills. Mirror the sync path.
self._mlx_runner.flush_all_decode_kv() self._mlx_runner.flush_all_decode_kv()
return self._async_extend_batch(batch) return self._async_extend_batch(batch)
@@ -330,6 +339,7 @@ class MlxTpModelWorker(TpModelWorker):
prefix_slot_ids=prefix_slot_ids, prefix_slot_ids=prefix_slot_ids,
new_slot_ids=req_new_slots, new_slot_ids=req_new_slots,
req_pool_idx=req.req_pool_idx, req_pool_idx=req.req_pool_idx,
req=req,
) )
) )
@@ -377,7 +387,23 @@ class MlxTpModelWorker(TpModelWorker):
@staticmethod @staticmethod
def _cache_state(cache_list) -> list[mx.array]: def _cache_state(cache_list) -> list[mx.array]:
"""Flatten a per-layer cache list to its ``state`` arrays.""" """Flatten a per-layer cache list to its ``state`` arrays."""
return [s for c in cache_list for s in c.state] arrays: list[mx.array] = []
def collect(value):
if isinstance(value, mx.array):
arrays.append(value)
elif value is None:
return
elif isinstance(value, (list, tuple)):
for item in value:
collect(item)
elif isinstance(value, dict):
for item in value.values():
collect(item)
for cache in cache_list:
collect(getattr(cache, "state", ()))
return arrays
def async_chained_decode_mlx( def async_chained_decode_mlx(
self, self,
@@ -828,6 +828,11 @@ class SchedulerBatchResultProcessor:
else: else:
if self.server_args.enable_hisparse: if self.server_args.enable_hisparse:
self.hisparse_coordinator.request_finished(req) self.hisparse_coordinator.request_finished(req)
prepare_release = getattr(
self.model_worker, "prepare_for_kv_cache_release", None
)
if callable(prepare_release):
prepare_release(req)
release_kv_cache(req, self.tree_cache) release_kv_cache(req, self.tree_cache)
req.time_stats.set_completion_time() req.time_stats.set_completion_time()
@@ -9,6 +9,9 @@ if TYPE_CHECKING:
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.mem_cache.unified_cache_components import ComponentType from sglang.srt.mem_cache.unified_cache_components import ComponentType
from sglang.srt.mem_cache.unified_cache_components.tree_component import (
TreeComponent,
)
@dataclasses.dataclass @dataclasses.dataclass
@@ -44,3 +47,6 @@ class CacheInitParams:
cache_ttl_seconds: Optional[float] = None cache_ttl_seconds: Optional[float] = None
tree_components: Optional[tuple[ComponentType, ...]] = None tree_components: Optional[tuple[ComponentType, ...]] = None
component_registry_override: Optional[dict[ComponentType, type[TreeComponent]]] = (
None
)
+10 -1
View File
@@ -17,6 +17,7 @@ from typing import TYPE_CHECKING, Any, Callable, Optional
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
from sglang.srt.mem_cache.cache_init_params import CacheInitParams from sglang.srt.mem_cache.cache_init_params import CacheInitParams
from sglang.srt.utils.tensor_bridge import use_mlx
if TYPE_CHECKING: if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig from sglang.srt.configs.model_config import ModelConfig
@@ -94,7 +95,7 @@ def default_radix_cache_factory(ctx: TreeCacheBuildContext) -> BasePrefixCache:
logger.info("Using experimental C++ radix tree implementation.") logger.info("Using experimental C++ radix tree implementation.")
return RadixCacheCpp(params=params, server_args=server_args) return RadixCacheCpp(params=params, server_args=server_args)
if envs.SGLANG_ENABLE_UNIFIED_RADIX_TREE.get(): if envs.SGLANG_ENABLE_UNIFIED_RADIX_TREE.get() or use_mlx():
from sglang.srt.mem_cache.unified_cache_components import ComponentType from sglang.srt.mem_cache.unified_cache_components import ComponentType
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
@@ -104,6 +105,14 @@ def default_radix_cache_factory(ctx: TreeCacheBuildContext) -> BasePrefixCache:
ComponentType.SWA if ctx.is_hybrid_swa else ComponentType.MAMBA ComponentType.SWA if ctx.is_hybrid_swa else ComponentType.MAMBA
) )
params.tree_components = tuple(tree_components) params.tree_components = tuple(tree_components)
if use_mlx() and ctx.is_hybrid_ssm:
from sglang.srt.hardware_backend.mlx.kv_cache.auxiliary_state import (
MlxAuxiliaryStateComponent,
)
params.component_registry_override = {
ComponentType.MAMBA: MlxAuxiliaryStateComponent,
}
cache = UnifiedRadixCache(params) cache = UnifiedRadixCache(params)
if ctx.enable_hierarchical_cache: if ctx.enable_hierarchical_cache:
cache.init_hicache(server_args, params) cache.init_hicache(server_args, params)
@@ -258,8 +258,14 @@ class UnifiedRadixCache(KVCacheEventMixin, BasePrefixCache):
assert params.tree_components is not None assert params.tree_components is not None
self.tree_components = tuple(params.tree_components) self.tree_components = tuple(params.tree_components)
component_registry = COMPONENT_REGISTRY
if params.component_registry_override:
component_registry = {
**COMPONENT_REGISTRY,
**params.component_registry_override,
}
self.components: dict[ComponentType, TreeComponent] = { self.components: dict[ComponentType, TreeComponent] = {
ct: COMPONENT_REGISTRY[ct](self, params) for ct in self.tree_components ct: component_registry[ct](self, params) for ct in self.tree_components
} }
self._components_tuple: tuple[TreeComponent, ...] = tuple( self._components_tuple: tuple[TreeComponent, ...] = tuple(
self.components.values() self.components.values()
File diff suppressed because it is too large Load Diff
@@ -41,8 +41,11 @@ from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool, SWATokenToKVPoolAllo
from sglang.srt.mem_cache.unified_cache_components.tree_component import ( from sglang.srt.mem_cache.unified_cache_components.tree_component import (
CacheTransferPhase, CacheTransferPhase,
ComponentType, ComponentType,
EvictLayer,
TreeComponent,
) )
from sglang.srt.mem_cache.unified_radix_cache import ( from sglang.srt.mem_cache.unified_radix_cache import (
COMPONENT_REGISTRY,
UnifiedRadixCache, UnifiedRadixCache,
UnifiedTreeNode, UnifiedTreeNode,
) )
@@ -119,6 +122,52 @@ class CacheConfig:
return "_".join(parts) return "_".join(parts)
class _FakeFullComponent(TreeComponent):
component_type = ComponentType.FULL
def create_match_validator(self, match_device_only: bool = False):
return lambda node: True
def redistribute_on_node_split(self, new_parent, child):
return None
def evict_component(
self, node, target: EvictLayer = EvictLayer.DEVICE
) -> tuple[int, int]:
return 0, 0
def drive_eviction(self, params: EvictParams, tracker: dict[ComponentType, int]):
return None
def acquire_component_lock(self, node, result):
return result
def release_component_lock(self, node, params):
return None
class TestUnifiedRadixComponentRegistryOverride(CustomTestCase):
def test_component_registry_override_is_instance_local(self):
params = CacheInitParams(
req_to_token_pool=ReqToTokenPool(
size=2,
max_context_len=8,
device="cpu",
enable_memory_saver=False,
),
token_to_kv_pool_allocator=None,
page_size=1,
disable=True,
tree_components=(ComponentType.FULL,),
component_registry_override={ComponentType.FULL: _FakeFullComponent},
)
tree = UnifiedRadixCache(params=params)
self.assertIsInstance(tree.components[ComponentType.FULL], _FakeFullComponent)
self.assertIsNot(COMPONENT_REGISTRY[ComponentType.FULL], _FakeFullComponent)
def build_fixture(cfg: CacheConfig, *, enable_kv_cache_events: bool = False): def build_fixture(cfg: CacheConfig, *, enable_kv_cache_events: bool = False):
"""Create (tree, allocator, req_to_token_pool) from a CacheConfig.""" """Create (tree, allocator, req_to_token_pool) from a CacheConfig."""
server_args = ServerArgs(model_path="dummy", page_size=cfg.page_size) server_args = ServerArgs(model_path="dummy", page_size=cfg.page_size)