[Apple Silicon] [MLX] MLX decode partial overlap scheduling for generation (async eval) (#22416)
Co-authored-by: R0CKSTAR <yeahdongcn@gmail.com> Co-authored-by: Alex Nails <alex.nails@radixark.ai>
This commit is contained in:
co-authored by
R0CKSTAR
Alex Nails
parent
d4040e7010
commit
3272af2f00
@@ -1,6 +1,6 @@
|
|||||||
# Apple Silicon with Metal
|
# Apple Silicon with Metal (MLX)
|
||||||
|
|
||||||
This document describes how run SGLang on Apple Silicon using [Metal](https://developer.apple.com/metal/). If you encounter issues or have questions, please [open an issue](https://github.com/sgl-project/sglang/issues).
|
This document describes how run SGLang on Apple Silicon using [Metal (MLX)](https://opensource.apple.com/projects/mlx/). If you encounter issues or have questions, please [open an issue](https://github.com/sgl-project/sglang/issues).
|
||||||
|
|
||||||
## Install SGLang
|
## Install SGLang
|
||||||
|
|
||||||
@@ -18,3 +18,57 @@ pip install --upgrade pip
|
|||||||
rm -f python/pyproject.toml && mv python/pyproject_other.toml python/pyproject.toml
|
rm -f python/pyproject.toml && mv python/pyproject_other.toml python/pyproject.toml
|
||||||
uv pip install -e "python[all_mps]"
|
uv pip install -e "python[all_mps]"
|
||||||
```
|
```
|
||||||
|
|
||||||
|
## Launch of the Serving Engine
|
||||||
|
|
||||||
|
Launch the server with:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
SGLANG_USE_MLX=1 python -m sglang.launch_server \
|
||||||
|
--model <MODEL_ID_OR_PATH> \
|
||||||
|
--disable-cuda-graph \
|
||||||
|
--host 0.0.0.0
|
||||||
|
```
|
||||||
|
|
||||||
|
**Key Parameters Explained:**
|
||||||
|
|
||||||
|
1. `SGLANG_USE_MLX=1` - Enables the use of MLX as the SGLang runtime backend (if disabled, SGLang will fall back to `torch.mps`, which has less support)
|
||||||
|
2. `--disable-cuda-graph` - Disables usage of CUDA graph, which is not relevant for Apple Metal.
|
||||||
|
3. `--disable-overlap-schedule` - Disables overlap scheduling (enabled/not present by default) achieved using MLX's `async_eval()`
|
||||||
|
|
||||||
|
|
||||||
|
## Benchmarking with Requests
|
||||||
|
|
||||||
|
`sglang.benchmark_one_batch` calls the synchronous prefill/decode methods directly without going through the scheduler and the overlap code path.
|
||||||
|
|
||||||
|
`sglang.benchmark_offline_throughput` can toggle overlap scheduling as it uses the scheduler and the overlap code path by using the flag `--disable-overlap-schedule`.
|
||||||
|
|
||||||
|
### Throughput Testing
|
||||||
|
|
||||||
|
Basic synchronous one batch throughput:
|
||||||
|
```bash
|
||||||
|
SGLANG_USE_MLX=1 python -m sglang.bench_one_batch \
|
||||||
|
--model-path <MODEL_ID_OR_PATH> \
|
||||||
|
--disable-cuda-graph \
|
||||||
|
--tp-size 1 \
|
||||||
|
--batch-size 1 \
|
||||||
|
--input-len 60 \
|
||||||
|
--output-len 10
|
||||||
|
```
|
||||||
|
|
||||||
|
Synchronous offline throughput:
|
||||||
|
```bash
|
||||||
|
SGLANG_USE_MLX=1 python -m sglang.bench_offline_throughput \
|
||||||
|
--model-path <MODEL_ID_OR_PATH> \
|
||||||
|
--disable-cuda-graph \
|
||||||
|
--num-prompts 1 \
|
||||||
|
--disable-overlap-schedule
|
||||||
|
```
|
||||||
|
|
||||||
|
Asynchronous offline throughput:
|
||||||
|
```bash
|
||||||
|
SGLANG_USE_MLX=1 python -m sglang.bench_offline_throughput \
|
||||||
|
--model-path <MODEL_ID_OR_PATH> \
|
||||||
|
--disable-cuda-graph \
|
||||||
|
--num-prompts 1
|
||||||
|
```
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ metatags:
|
|||||||
description: "Run SGLang on Apple Silicon using the Metal backend."
|
description: "Run SGLang on Apple Silicon using the Metal backend."
|
||||||
---
|
---
|
||||||
|
|
||||||
This document describes how run SGLang on Apple Silicon using [Metal](https://developer.apple.com/metal/). If you encounter issues or have questions, please [open an issue](https://github.com/sgl-project/sglang/issues).
|
This document describes how run SGLang on Apple Silicon using [Metal (MLX)](https://opensource.apple.com/projects/mlx/). If you encounter issues or have questions, please [open an issue](https://github.com/sgl-project/sglang/issues).
|
||||||
|
|
||||||
## Install SGLang
|
## Install SGLang
|
||||||
|
|
||||||
@@ -22,3 +22,57 @@ pip install --upgrade pip
|
|||||||
rm -f python/pyproject.toml && mv python/pyproject_other.toml python/pyproject.toml
|
rm -f python/pyproject.toml && mv python/pyproject_other.toml python/pyproject.toml
|
||||||
uv pip install -e "python[all_mps]"
|
uv pip install -e "python[all_mps]"
|
||||||
```
|
```
|
||||||
|
|
||||||
|
## Launch of the Serving Engine
|
||||||
|
|
||||||
|
Launch the server with:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
SGLANG_USE_MLX=1 python -m sglang.launch_server \
|
||||||
|
--model <MODEL_ID_OR_PATH> \
|
||||||
|
--disable-cuda-graph \
|
||||||
|
--host 0.0.0.0
|
||||||
|
```
|
||||||
|
|
||||||
|
**Key Parameters Explained:**
|
||||||
|
|
||||||
|
1. `SGLANG_USE_MLX=1` - Enables the use of MLX as the SGLang runtime backend (if disabled, SGLang will fall back to `torch.mps`, which has less support)
|
||||||
|
2. `--disable-cuda-graph` - Disables usage of CUDA graph, which is not relevant for Apple Metal.
|
||||||
|
3. `--disable-overlap-schedule` - Disables overlap scheduling (enabled/not present by default) achieved using MLX's `async_eval()`
|
||||||
|
|
||||||
|
|
||||||
|
## Benchmarking with Requests
|
||||||
|
|
||||||
|
`sglang.benchmark_one_batch` calls the synchronous prefill/decode methods directly without going through the scheduler and the overlap code path.
|
||||||
|
|
||||||
|
`sglang.benchmark_offline_throughput` can toggle overlap scheduling as it uses the scheduler and the overlap code path by using the flag `--disable-overlap-schedule`.
|
||||||
|
|
||||||
|
### Throughput Testing
|
||||||
|
|
||||||
|
Basic synchronous one batch throughput:
|
||||||
|
```bash
|
||||||
|
SGLANG_USE_MLX=1 python -m sglang.bench_one_batch \
|
||||||
|
--model-path <MODEL_ID_OR_PATH> \
|
||||||
|
--disable-cuda-graph \
|
||||||
|
--tp-size 1 \
|
||||||
|
--batch-size 1 \
|
||||||
|
--input-len 60 \
|
||||||
|
--output-len 10
|
||||||
|
```
|
||||||
|
|
||||||
|
Synchronous offline throughput:
|
||||||
|
```bash
|
||||||
|
SGLANG_USE_MLX=1 python -m sglang.bench_offline_throughput \
|
||||||
|
--model-path <MODEL_ID_OR_PATH> \
|
||||||
|
--disable-cuda-graph \
|
||||||
|
--num-prompts 1 \
|
||||||
|
--disable-overlap-schedule
|
||||||
|
```
|
||||||
|
|
||||||
|
Asynchronous offline throughput:
|
||||||
|
```bash
|
||||||
|
SGLANG_USE_MLX=1 python -m sglang.bench_offline_throughput \
|
||||||
|
--model-path <MODEL_ID_OR_PATH> \
|
||||||
|
--disable-cuda-graph \
|
||||||
|
--num-prompts 1
|
||||||
|
```
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from sglang.srt.hardware_backend.mlx.kv_cache.contiguous_cache import Contiguous
|
|||||||
_thread_local = threading.local()
|
_thread_local = threading.local()
|
||||||
|
|
||||||
|
|
||||||
|
# TODO: Move from threading to multiprocessing or asyncio
|
||||||
@dataclass
|
@dataclass
|
||||||
class BatchedDecodeContext:
|
class BatchedDecodeContext:
|
||||||
"""Context set before batched decode, read by attention wrappers."""
|
"""Context set before batched decode, read by attention wrappers."""
|
||||||
|
|||||||
@@ -5,10 +5,18 @@ scheduler (``TokenToKVPoolAllocator`` / ``RadixCache``). This runner
|
|||||||
reads cached KV from ``MlxKVPool``, runs the forward pass, and writes
|
reads cached KV from ``MlxKVPool``, runs the forward pass, and writes
|
||||||
new KV back. Each request also keeps a ``ContiguousKVCache`` for
|
new KV back. Each request also keeps a ``ContiguousKVCache`` for
|
||||||
decode-time attention.
|
decode-time attention.
|
||||||
|
|
||||||
|
The module also exposes a lazy-eval (`*_start` / `*_finalize`) surface
|
||||||
|
used by the MLX overlap scheduler to pipeline CPU bookkeeping with
|
||||||
|
GPU execution. The lazy API is a thin split of the synchronous API:
|
||||||
|
``*_start`` builds the compute graph without materialising outputs,
|
||||||
|
``*_finalize`` blocks on the lazy token(s) and commits per-request
|
||||||
|
state.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
import time
|
import time
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
import psutil
|
import psutil
|
||||||
@@ -32,6 +40,56 @@ from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class MlxPendingPrefill:
|
||||||
|
"""Lazy prefill state, finalised after ``mx.eval``/``async_eval``.
|
||||||
|
|
||||||
|
``cache`` is the per-layer list of ``ContiguousKVCache`` that will
|
||||||
|
become ``_req_caches[req_id]`` once the request is committed. It
|
||||||
|
may have been converted from a transient ``PoolBackedCache`` list
|
||||||
|
already (so its ``state`` arrays are safe to hand to ``async_eval``).
|
||||||
|
"""
|
||||||
|
|
||||||
|
lazy_token: mx.array
|
||||||
|
cache: list # list[ContiguousKVCache]
|
||||||
|
req_id: str
|
||||||
|
full_token_ids: list[int]
|
||||||
|
req_pool_idx: int
|
||||||
|
synced_offset: int
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class MlxPendingExtend:
|
||||||
|
"""Lazy chunked-prefill-continuation state for an existing request.
|
||||||
|
|
||||||
|
Mirrors :meth:`MlxModelRunner.extend` split into launch/finalize
|
||||||
|
halves. ``cache`` is the request's existing per-layer cache (not a
|
||||||
|
fresh one) so the graph writes extend onto the already-materialised
|
||||||
|
prefix.
|
||||||
|
"""
|
||||||
|
|
||||||
|
lazy_token: mx.array
|
||||||
|
req_id: str
|
||||||
|
new_token_ids: list[int]
|
||||||
|
new_synced_offset: int
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class MlxPendingDecode:
|
||||||
|
"""Lazy decode state, finalised after ``mx.eval``/``async_eval``.
|
||||||
|
|
||||||
|
``caches`` is a per-request list of per-layer ``ContiguousKVCache``
|
||||||
|
references (``caches[req_idx][layer_idx]``). These are the same
|
||||||
|
objects the attention wrapper writes into during the forward pass,
|
||||||
|
so :meth:`decode_batch_start_chained` can launch the next step on
|
||||||
|
top of the same caches without materialising this step first.
|
||||||
|
"""
|
||||||
|
|
||||||
|
lazy_tokens: mx.array
|
||||||
|
req_ids: list[str]
|
||||||
|
caches: list # list[list[ContiguousKVCache]]
|
||||||
|
|
||||||
|
|
||||||
class MlxModelRunner:
|
class MlxModelRunner:
|
||||||
"""MLX model runner with radix-cache prefix sharing."""
|
"""MLX model runner with radix-cache prefix sharing."""
|
||||||
|
|
||||||
@@ -48,6 +106,8 @@ class MlxModelRunner:
|
|||||||
self.model = None
|
self.model = None
|
||||||
self.disable_radix_cache = disable_radix_cache
|
self.disable_radix_cache = disable_radix_cache
|
||||||
self._mem_fraction_static = mem_fraction_static
|
self._mem_fraction_static = mem_fraction_static
|
||||||
|
# Counter used to trigger periodic mx.clear_cache() calls.
|
||||||
|
self._decode_step_ct: int = 0
|
||||||
|
|
||||||
self._load_model()
|
self._load_model()
|
||||||
|
|
||||||
@@ -104,6 +164,21 @@ class MlxModelRunner:
|
|||||||
"""Evaluate token result and all cache buffers in one mx.eval call."""
|
"""Evaluate token result and all cache buffers in one mx.eval call."""
|
||||||
mx.eval(token_result, *[s for c in cache for s in c.state])
|
mx.eval(token_result, *[s for c in cache for s in c.state])
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _cache_state_arrays(
|
||||||
|
pending_caches: list[list[ContiguousKVCache | PoolBackedCache]],
|
||||||
|
) -> list[mx.array]:
|
||||||
|
"""Flatten pending decode cache state list into an array list.
|
||||||
|
|
||||||
|
Safe to hand to ``mx.async_eval``.
|
||||||
|
"""
|
||||||
|
return [
|
||||||
|
s
|
||||||
|
for cache_list in pending_caches
|
||||||
|
for cache in cache_list
|
||||||
|
for s in cache.state
|
||||||
|
]
|
||||||
|
|
||||||
def _load_model(self):
|
def _load_model(self):
|
||||||
"""Load model using mlx_lm."""
|
"""Load model using mlx_lm."""
|
||||||
logger.info(f"Loading MLX model: {self.model_path}")
|
logger.info(f"Loading MLX model: {self.model_path}")
|
||||||
@@ -203,74 +278,16 @@ class MlxModelRunner:
|
|||||||
req_pool_idx: int,
|
req_pool_idx: int,
|
||||||
) -> int:
|
) -> int:
|
||||||
"""Prefill a request. Returns next_token_id."""
|
"""Prefill a request. Returns next_token_id."""
|
||||||
num_layers = self._num_layers
|
pending = self.prefill_start(
|
||||||
prefix_len = len(prefix_slot_ids)
|
req_id=req_id,
|
||||||
|
new_token_ids=new_token_ids,
|
||||||
if self.disable_radix_cache:
|
full_token_ids=full_token_ids,
|
||||||
cache = self._acquire_cache()
|
prefix_slot_ids=prefix_slot_ids,
|
||||||
input_ids = mx.array([new_token_ids], dtype=mx.int32)
|
new_slot_ids=new_slot_ids,
|
||||||
model_output = self.model(input_ids, cache=cache)
|
req_pool_idx=req_pool_idx,
|
||||||
logits = self._extract_logits(model_output)
|
|
||||||
next_token_mlx = mx.argmax(logits[:, -1, :], axis=-1)
|
|
||||||
self._eval_with_cache(next_token_mlx, cache)
|
|
||||||
next_token = int(next_token_mlx.item())
|
|
||||||
|
|
||||||
self._req_token_ids[req_id] = list(full_token_ids) + [next_token]
|
|
||||||
self._req_caches[req_id] = cache
|
|
||||||
self._req_pool_idx[req_id] = req_pool_idx
|
|
||||||
self._req_synced_offset[req_id] = 0
|
|
||||||
return next_token
|
|
||||||
|
|
||||||
assert self._kv_pool is not None
|
|
||||||
|
|
||||||
new_token_count = len(new_token_ids)
|
|
||||||
|
|
||||||
if prefix_len > 0:
|
|
||||||
slot_ids_mx = mx.array(prefix_slot_ids, dtype=mx.int32)
|
|
||||||
cache = [
|
|
||||||
PoolBackedCache(self._kv_pool, i, slot_ids_mx, prefix_len)
|
|
||||||
for i in range(num_layers)
|
|
||||||
]
|
|
||||||
else:
|
|
||||||
cache = self._acquire_cache()
|
|
||||||
|
|
||||||
if new_token_count > 0:
|
|
||||||
extend_tokens = new_token_ids
|
|
||||||
else:
|
|
||||||
# Full cache hit — rerun last token to get next-token logits
|
|
||||||
extend_tokens = full_token_ids[-1:]
|
|
||||||
for c in cache:
|
|
||||||
c.offset = max(c.offset - 1, 0)
|
|
||||||
|
|
||||||
input_ids = mx.array([extend_tokens], dtype=mx.int32)
|
|
||||||
model_output = self.model(input_ids, cache=cache)
|
|
||||||
logits = self._extract_logits(model_output)
|
|
||||||
|
|
||||||
last_logits = logits[:, -1, :]
|
|
||||||
next_token_mlx = mx.argmax(last_logits, axis=-1)
|
|
||||||
|
|
||||||
# Convert PoolBackedCache → ContiguousKVCache for decode
|
|
||||||
if prefix_len > 0:
|
|
||||||
contiguous_cache = self._acquire_cache()
|
|
||||||
for layer_idx in range(num_layers):
|
|
||||||
pbc = cache[layer_idx]
|
|
||||||
contiguous_cache[layer_idx].update_and_fetch(
|
|
||||||
pbc._full_keys, pbc._full_values
|
|
||||||
)
|
)
|
||||||
cache = contiguous_cache
|
self._eval_with_cache(pending.lazy_token, pending.cache)
|
||||||
|
return self.prefill_finalize(pending)
|
||||||
self._eval_with_cache(next_token_mlx, cache)
|
|
||||||
next_token = int(next_token_mlx.item())
|
|
||||||
|
|
||||||
if new_slot_ids:
|
|
||||||
self._sync_new_kv_to_pool(cache, prefix_len, new_slot_ids)
|
|
||||||
|
|
||||||
self._req_token_ids[req_id] = list(full_token_ids) + [next_token]
|
|
||||||
self._req_caches[req_id] = cache
|
|
||||||
self._req_pool_idx[req_id] = req_pool_idx
|
|
||||||
self._req_synced_offset[req_id] = prefix_len + len(new_slot_ids)
|
|
||||||
|
|
||||||
return next_token
|
|
||||||
|
|
||||||
def extend(
|
def extend(
|
||||||
self,
|
self,
|
||||||
@@ -279,32 +296,9 @@ class MlxModelRunner:
|
|||||||
new_slot_ids: list[int],
|
new_slot_ids: list[int],
|
||||||
) -> int:
|
) -> int:
|
||||||
"""Continue prefill for a chunked request. Returns next_token_id."""
|
"""Continue prefill for a chunked request. Returns next_token_id."""
|
||||||
assert req_id in self._req_caches, f"extend called for unknown request {req_id}"
|
pending = self.extend_start(req_id, new_token_ids, new_slot_ids)
|
||||||
|
self._eval_with_cache(pending.lazy_token, self._req_caches[req_id])
|
||||||
cache = self._req_caches[req_id]
|
return self.extend_finalize(pending)
|
||||||
|
|
||||||
input_ids = mx.array([new_token_ids], dtype=mx.int32)
|
|
||||||
model_output = self.model(input_ids, cache=cache)
|
|
||||||
logits = self._extract_logits(model_output)
|
|
||||||
|
|
||||||
last_logits = logits[:, -1, :]
|
|
||||||
next_token_mlx = mx.argmax(last_logits, axis=-1)
|
|
||||||
self._eval_with_cache(next_token_mlx, cache)
|
|
||||||
next_token = int(next_token_mlx.item())
|
|
||||||
|
|
||||||
prev_tokens = self._req_token_ids[req_id]
|
|
||||||
if prev_tokens:
|
|
||||||
prev_tokens.pop() # remove stale intermediate token
|
|
||||||
prev_tokens.extend(new_token_ids)
|
|
||||||
prev_tokens.append(next_token)
|
|
||||||
|
|
||||||
# Sync new chunk KV to pool immediately
|
|
||||||
if not self.disable_radix_cache and new_slot_ids:
|
|
||||||
synced = self._req_synced_offset[req_id]
|
|
||||||
self._sync_new_kv_to_pool(cache, synced, new_slot_ids)
|
|
||||||
self._req_synced_offset[req_id] = synced + len(new_slot_ids)
|
|
||||||
|
|
||||||
return next_token
|
|
||||||
|
|
||||||
def _sync_new_kv_to_pool(
|
def _sync_new_kv_to_pool(
|
||||||
self,
|
self,
|
||||||
@@ -318,6 +312,7 @@ class MlxModelRunner:
|
|||||||
num_layers = len(cache)
|
num_layers = len(cache)
|
||||||
end = cache_start + len(slot_ids)
|
end = cache_start + len(slot_ids)
|
||||||
slot_ids_mx = mx.array(slot_ids, dtype=mx.int32)
|
slot_ids_mx = mx.array(slot_ids, dtype=mx.int32)
|
||||||
|
# TODO: Standardize ContiguousKVCache size to avoid transpose
|
||||||
# Transpose cache (1, n_kv_heads, S, head_dim) → pool (S, n_kv_heads, head_dim)
|
# Transpose cache (1, n_kv_heads, S, head_dim) → pool (S, n_kv_heads, head_dim)
|
||||||
k_all = mx.stack(
|
k_all = mx.stack(
|
||||||
[
|
[
|
||||||
@@ -370,11 +365,172 @@ class MlxModelRunner:
|
|||||||
req_ids: list[str],
|
req_ids: list[str],
|
||||||
) -> list[int]:
|
) -> list[int]:
|
||||||
"""Decode one token per request."""
|
"""Decode one token per request."""
|
||||||
|
pending = self.decode_batch_start(req_ids)
|
||||||
|
# Evaluate lazy_tokens together with every affected cache buffer so
|
||||||
|
# the attention write-then-read ordering is materialised in one
|
||||||
|
# kernel submission.
|
||||||
|
cache_arrays = self._cache_state_arrays(pending.caches)
|
||||||
|
mx.eval(pending.lazy_tokens, *cache_arrays)
|
||||||
|
return self.decode_batch_finalize(pending)
|
||||||
|
|
||||||
|
def prefill_start(
|
||||||
|
self,
|
||||||
|
req_id: str,
|
||||||
|
new_token_ids: list[int],
|
||||||
|
full_token_ids: list[int],
|
||||||
|
prefix_slot_ids: list[int],
|
||||||
|
new_slot_ids: list[int],
|
||||||
|
req_pool_idx: int,
|
||||||
|
) -> MlxPendingPrefill:
|
||||||
|
"""Queue a prefill forward pass without evaluating.
|
||||||
|
|
||||||
|
Returns an :class:`MlxPendingPrefill` containing the lazy
|
||||||
|
next-token ``mx.array`` plus everything needed to commit the
|
||||||
|
request in :meth:`prefill_finalize`. The caller drives the GPU
|
||||||
|
by handing ``lazy_token`` (and cache state) to ``mx.async_eval``.
|
||||||
|
"""
|
||||||
|
num_layers = self._num_layers
|
||||||
|
prefix_len = len(prefix_slot_ids)
|
||||||
|
|
||||||
|
if self.disable_radix_cache:
|
||||||
|
cache = self._acquire_cache()
|
||||||
|
input_ids = mx.array([new_token_ids], dtype=mx.int32)
|
||||||
|
model_output = self.model(input_ids, cache=cache)
|
||||||
|
logits = self._extract_logits(model_output)
|
||||||
|
lazy_token = mx.argmax(logits[:, -1, :], axis=-1)
|
||||||
|
return MlxPendingPrefill(
|
||||||
|
lazy_token=lazy_token,
|
||||||
|
cache=cache,
|
||||||
|
req_id=req_id,
|
||||||
|
full_token_ids=list(full_token_ids),
|
||||||
|
req_pool_idx=req_pool_idx,
|
||||||
|
synced_offset=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert self._kv_pool is not None
|
||||||
|
|
||||||
|
new_token_count = len(new_token_ids)
|
||||||
|
|
||||||
|
if prefix_len > 0:
|
||||||
|
slot_ids_mx = mx.array(prefix_slot_ids, dtype=mx.int32)
|
||||||
|
cache = [
|
||||||
|
PoolBackedCache(self._kv_pool, i, slot_ids_mx, prefix_len)
|
||||||
|
for i in range(num_layers)
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
cache = self._acquire_cache()
|
||||||
|
|
||||||
|
if new_token_count > 0:
|
||||||
|
extend_tokens = new_token_ids
|
||||||
|
else:
|
||||||
|
# Full cache hit — rerun last token to get next-token logits
|
||||||
|
extend_tokens = full_token_ids[-1:]
|
||||||
|
for c in cache:
|
||||||
|
c.offset = max(c.offset - 1, 0)
|
||||||
|
|
||||||
|
input_ids = mx.array([extend_tokens], dtype=mx.int32)
|
||||||
|
model_output = self.model(input_ids, cache=cache)
|
||||||
|
logits = self._extract_logits(model_output)
|
||||||
|
|
||||||
|
last_logits = logits[:, -1, :]
|
||||||
|
lazy_token = mx.argmax(last_logits, axis=-1)
|
||||||
|
|
||||||
|
# Convert PoolBackedCache → ContiguousKVCache for decode.
|
||||||
|
# This appends a lazy slice-assign onto the forward graph; the
|
||||||
|
# arrays get materialised when the caller evaluates lazy_token.
|
||||||
|
if prefix_len > 0:
|
||||||
|
contiguous_cache = self._acquire_cache()
|
||||||
|
for layer_idx in range(num_layers):
|
||||||
|
pbc = cache[layer_idx]
|
||||||
|
contiguous_cache[layer_idx].update_and_fetch(
|
||||||
|
pbc._full_keys, pbc._full_values
|
||||||
|
)
|
||||||
|
cache = contiguous_cache
|
||||||
|
|
||||||
|
if new_slot_ids:
|
||||||
|
self._sync_new_kv_to_pool(cache, prefix_len, new_slot_ids)
|
||||||
|
|
||||||
|
return MlxPendingPrefill(
|
||||||
|
lazy_token=lazy_token,
|
||||||
|
cache=cache,
|
||||||
|
req_id=req_id,
|
||||||
|
full_token_ids=list(full_token_ids),
|
||||||
|
req_pool_idx=req_pool_idx,
|
||||||
|
synced_offset=prefix_len + len(new_slot_ids),
|
||||||
|
)
|
||||||
|
|
||||||
|
def prefill_finalize(self, pending: MlxPendingPrefill) -> int:
|
||||||
|
"""Materialise a pending prefill and commit per-request state.
|
||||||
|
|
||||||
|
Must be called *after* ``pending.lazy_token`` has been handed to
|
||||||
|
``mx.async_eval`` / ``mx.eval``. ``.item()`` here is blocking on
|
||||||
|
that specific lazy scalar.
|
||||||
|
"""
|
||||||
|
next_token = int(pending.lazy_token.item())
|
||||||
|
self._req_token_ids[pending.req_id] = list(pending.full_token_ids) + [
|
||||||
|
next_token
|
||||||
|
]
|
||||||
|
self._req_caches[pending.req_id] = pending.cache
|
||||||
|
self._req_pool_idx[pending.req_id] = pending.req_pool_idx
|
||||||
|
self._req_synced_offset[pending.req_id] = pending.synced_offset
|
||||||
|
return next_token
|
||||||
|
|
||||||
|
def extend_start(
|
||||||
|
self,
|
||||||
|
req_id: str,
|
||||||
|
new_token_ids: list[int],
|
||||||
|
new_slot_ids: list[int],
|
||||||
|
) -> MlxPendingExtend:
|
||||||
|
"""Queue chunked-prefill continuation without evaluating."""
|
||||||
|
assert (
|
||||||
|
req_id in self._req_caches
|
||||||
|
), f"extend_start called for unknown request {req_id}"
|
||||||
|
|
||||||
|
cache = self._req_caches[req_id]
|
||||||
|
|
||||||
|
input_ids = mx.array([new_token_ids], dtype=mx.int32)
|
||||||
|
model_output = self.model(input_ids, cache=cache)
|
||||||
|
logits = self._extract_logits(model_output)
|
||||||
|
lazy_token = mx.argmax(logits[:, -1, :], axis=-1)
|
||||||
|
|
||||||
|
if not self.disable_radix_cache and new_slot_ids:
|
||||||
|
synced = self._req_synced_offset[req_id]
|
||||||
|
self._sync_new_kv_to_pool(cache, synced, new_slot_ids)
|
||||||
|
new_synced_offset = synced + len(new_slot_ids)
|
||||||
|
else:
|
||||||
|
new_synced_offset = self._req_synced_offset.get(req_id, 0)
|
||||||
|
|
||||||
|
return MlxPendingExtend(
|
||||||
|
lazy_token=lazy_token,
|
||||||
|
req_id=req_id,
|
||||||
|
new_token_ids=list(new_token_ids),
|
||||||
|
new_synced_offset=new_synced_offset,
|
||||||
|
)
|
||||||
|
|
||||||
|
def extend_finalize(self, pending: MlxPendingExtend) -> int:
|
||||||
|
"""Materialise a pending extend and commit per-request state."""
|
||||||
|
next_token = int(pending.lazy_token.item())
|
||||||
|
|
||||||
|
prev_tokens = self._req_token_ids[pending.req_id]
|
||||||
|
if prev_tokens:
|
||||||
|
prev_tokens.pop() # remove stale intermediate token
|
||||||
|
prev_tokens.extend(pending.new_token_ids)
|
||||||
|
prev_tokens.append(next_token)
|
||||||
|
|
||||||
|
self._req_synced_offset[pending.req_id] = pending.new_synced_offset
|
||||||
|
return next_token
|
||||||
|
|
||||||
|
def decode_batch_start(self, req_ids: list[str]) -> MlxPendingDecode:
|
||||||
|
"""Queue a decode forward pass without evaluating.
|
||||||
|
|
||||||
|
The caller is responsible for calling ``mx.async_eval`` on the
|
||||||
|
returned ``lazy_tokens`` (and optionally per-cache state arrays)
|
||||||
|
to kick off GPU work before :meth:`decode_batch_finalize`.
|
||||||
|
"""
|
||||||
batch_size = len(req_ids)
|
batch_size = len(req_ids)
|
||||||
num_layers = self._num_layers
|
num_layers = self._num_layers
|
||||||
|
|
||||||
caches = [self._req_caches[rid] for rid in req_ids]
|
caches = [self._req_caches[rid] for rid in req_ids]
|
||||||
seq_lens = [caches[i][0].offset for i in range(batch_size)]
|
|
||||||
|
|
||||||
if batch_size == 1:
|
if batch_size == 1:
|
||||||
cache = caches[0]
|
cache = caches[0]
|
||||||
@@ -382,9 +538,14 @@ class MlxModelRunner:
|
|||||||
input_ids = mx.array([[last_token]], dtype=mx.int32)
|
input_ids = mx.array([[last_token]], dtype=mx.int32)
|
||||||
model_output = self.model(input_ids, cache=cache)
|
model_output = self.model(input_ids, cache=cache)
|
||||||
logits = self._extract_logits(model_output)
|
logits = self._extract_logits(model_output)
|
||||||
next_tokens_mlx = mx.argmax(logits[:, -1, :], axis=-1)
|
lazy_tokens = mx.argmax(logits[:, -1, :], axis=-1)
|
||||||
self._eval_with_cache(next_tokens_mlx, cache)
|
return MlxPendingDecode(
|
||||||
else:
|
lazy_tokens=lazy_tokens,
|
||||||
|
req_ids=list(req_ids),
|
||||||
|
caches=caches,
|
||||||
|
)
|
||||||
|
|
||||||
|
seq_lens = [caches[i][0].offset for i in range(batch_size)]
|
||||||
layer_caches = [
|
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 range(num_layers)
|
||||||
@@ -402,22 +563,116 @@ class MlxModelRunner:
|
|||||||
batched_input = mx.array(last_tokens, dtype=mx.int32)[:, None]
|
batched_input = mx.array(last_tokens, dtype=mx.int32)[:, None]
|
||||||
model_output = self.model(batched_input, cache=shim_cache)
|
model_output = self.model(batched_input, cache=shim_cache)
|
||||||
logits = self._extract_logits(model_output)
|
logits = self._extract_logits(model_output)
|
||||||
next_tokens_mlx = mx.argmax(logits[:, -1, :], axis=-1)
|
lazy_tokens = mx.argmax(logits[:, -1, :], axis=-1)
|
||||||
|
|
||||||
eval_targets = [next_tokens_mlx]
|
|
||||||
for c_list in caches:
|
|
||||||
for c in c_list:
|
|
||||||
eval_targets.append(c.keys)
|
|
||||||
eval_targets.append(c.values)
|
|
||||||
mx.eval(*eval_targets)
|
|
||||||
finally:
|
finally:
|
||||||
clear_context()
|
clear_context()
|
||||||
|
|
||||||
next_tokens = next_tokens_mlx.tolist()
|
return MlxPendingDecode(
|
||||||
|
lazy_tokens=lazy_tokens,
|
||||||
|
req_ids=list(req_ids),
|
||||||
|
caches=caches,
|
||||||
|
)
|
||||||
|
|
||||||
for i, rid in enumerate(req_ids):
|
def decode_batch_start_chained(
|
||||||
|
self,
|
||||||
|
prev: MlxPendingDecode,
|
||||||
|
) -> MlxPendingDecode:
|
||||||
|
"""Build the next decode step on top of a still-lazy previous decode.
|
||||||
|
|
||||||
|
Feeds ``prev.lazy_tokens`` (an unevaluated ``mx.array`` of shape
|
||||||
|
``(B,)``) as the next step's input ids, reusing
|
||||||
|
``prev.caches`` in-place so that the per-layer ``ContiguousKVCache``
|
||||||
|
writes from step N and step N+1 land in the same buffers. MLX
|
||||||
|
tracks the full dependency graph, so once ``mx.async_eval`` is
|
||||||
|
called the GPU executes N+1 immediately after N with no gap.
|
||||||
|
|
||||||
|
Caller contract:
|
||||||
|
|
||||||
|
* ``prev`` MUST refer to the same set of requests (same order) as
|
||||||
|
the batch the caller intends to run next. Composition changes
|
||||||
|
(finished reqs, new prefills) must break the chain instead.
|
||||||
|
* After calling this, finalise ``prev`` BEFORE finalising the
|
||||||
|
returned pending: state bookkeeping for step N has to happen
|
||||||
|
before step N+1's bookkeeping.
|
||||||
|
"""
|
||||||
|
batch_size = len(prev.req_ids)
|
||||||
|
num_layers = self._num_layers
|
||||||
|
caches = prev.caches
|
||||||
|
|
||||||
|
# TODO (changminbark): Need to fix ContiguousKVCache.write_token
|
||||||
|
# to accommodate dynamic growing like ContiguousKVCache.update_and_fetch.
|
||||||
|
|
||||||
|
# After prev's graph ran, each ContiguousKVCache.offset was
|
||||||
|
# bumped by one per layer — attention wrapper's `write_token`
|
||||||
|
# mutates the Python offset synchronously at graph-build time.
|
||||||
|
# So layer-0 offsets reflect the position the NEW token will
|
||||||
|
# be written at in step N+1 (and equivalently the RoPE offset).
|
||||||
|
seq_lens = [caches[i][0].offset for i in range(batch_size)]
|
||||||
|
|
||||||
|
if batch_size == 1:
|
||||||
|
cache = caches[0]
|
||||||
|
batched_input = prev.lazy_tokens[:, None]
|
||||||
|
model_output = self.model(batched_input, cache=cache)
|
||||||
|
logits = self._extract_logits(model_output)
|
||||||
|
lazy_tokens = mx.argmax(logits[:, -1, :], axis=-1)
|
||||||
|
return MlxPendingDecode(
|
||||||
|
lazy_tokens=lazy_tokens,
|
||||||
|
req_ids=prev.req_ids,
|
||||||
|
caches=caches,
|
||||||
|
)
|
||||||
|
|
||||||
|
layer_caches = [
|
||||||
|
[caches[i][layer_idx] for i in range(batch_size)]
|
||||||
|
for layer_idx in range(num_layers)
|
||||||
|
]
|
||||||
|
ctx = BatchedDecodeContext(
|
||||||
|
batch_size=batch_size,
|
||||||
|
seq_lens=seq_lens,
|
||||||
|
layer_caches=layer_caches,
|
||||||
|
)
|
||||||
|
set_context(ctx)
|
||||||
|
try:
|
||||||
|
max_offset = max(seq_lens)
|
||||||
|
shim_cache = [OffsetCache(offset=max_offset) for _ in range(num_layers)]
|
||||||
|
batched_input = prev.lazy_tokens[:, None]
|
||||||
|
model_output = self.model(batched_input, cache=shim_cache)
|
||||||
|
logits = self._extract_logits(model_output)
|
||||||
|
lazy_tokens = mx.argmax(logits[:, -1, :], axis=-1)
|
||||||
|
finally:
|
||||||
|
clear_context()
|
||||||
|
|
||||||
|
return MlxPendingDecode(
|
||||||
|
lazy_tokens=lazy_tokens,
|
||||||
|
req_ids=prev.req_ids,
|
||||||
|
caches=caches,
|
||||||
|
)
|
||||||
|
|
||||||
|
def decode_batch_finalize(
|
||||||
|
self,
|
||||||
|
pending: MlxPendingDecode,
|
||||||
|
) -> list[int]:
|
||||||
|
"""Materialise a pending decode and update per-request token lists.
|
||||||
|
|
||||||
|
``pending.lazy_tokens.tolist()`` implicitly blocks until that
|
||||||
|
specific lazy array (and its graph ancestors, including the
|
||||||
|
per-request cache writes for this step) is evaluated. The
|
||||||
|
caller should have previously handed this pending's lazy_tokens
|
||||||
|
to ``mx.async_eval`` (or to a subsequent chained step that will
|
||||||
|
be async_eval'd).
|
||||||
|
"""
|
||||||
|
raw = pending.lazy_tokens.tolist()
|
||||||
|
if not isinstance(raw, list):
|
||||||
|
raw = [raw]
|
||||||
|
next_tokens = [int(t) for t in raw]
|
||||||
|
|
||||||
|
for i, rid in enumerate(pending.req_ids):
|
||||||
self._req_token_ids[rid].append(next_tokens[i])
|
self._req_token_ids[rid].append(next_tokens[i])
|
||||||
|
|
||||||
|
self._decode_step_ct += 1
|
||||||
|
# TODO (changminbark): allow for flag configuration for clearing mx cache
|
||||||
|
if self._decode_step_ct % 256 == 0:
|
||||||
|
mx.clear_cache()
|
||||||
|
|
||||||
return next_tokens
|
return next_tokens
|
||||||
|
|
||||||
def has_request(self, req_id: str) -> bool:
|
def has_request(self, req_id: str) -> bool:
|
||||||
|
|||||||
@@ -0,0 +1,234 @@
|
|||||||
|
"""MLX overlap scheduling mixin for the SGLang scheduler.
|
||||||
|
|
||||||
|
Provides ``event_loop_overlap_mlx``, which pipelines MLX forward
|
||||||
|
passes by keeping two in-flight lazy graphs queued on the GPU while
|
||||||
|
the scheduler runs its CPU-side bookkeeping on the tokens of the
|
||||||
|
older one. The lazy-graph primitives live in
|
||||||
|
``hardware_backend/mlx/tp_worker.py`` and ``model_runner.py``.
|
||||||
|
|
||||||
|
Each request's KV lives ina set of per-request, per-layer ``ContiguousKVCache``
|
||||||
|
objects that the ``MLXAttentionWrapper`` mutates in place during the forward pass.
|
||||||
|
Chained decodes reuse the same cache objects: step N+1's graph reads
|
||||||
|
step N's lazy writes via MLX's dependency tracking, so the GPU runs
|
||||||
|
both steps back-to-back with no idle gap.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import TYPE_CHECKING, List, Optional
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
|
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.utils import DynamicGradMode
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.hardware_backend.mlx.model_runner import (
|
||||||
|
MlxPendingDecode,
|
||||||
|
MlxPendingExtend,
|
||||||
|
MlxPendingPrefill,
|
||||||
|
)
|
||||||
|
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
||||||
|
from sglang.srt.managers.scheduler import Scheduler
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class MlxPendingJob:
|
||||||
|
"""Unfinished MLX work and graphs queued on the GPU.
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
lazy_tokens: Lazily evaluated token IDs produced by the forward
|
||||||
|
pass. Unevaluated; calling ``.tolist()`` / ``.item()`` /
|
||||||
|
``mx.eval`` on it will block until the Metal kernel finishes.
|
||||||
|
``None`` for idle batches.
|
||||||
|
prefills: MLX prefill state returned by the model worker — one
|
||||||
|
entry per new request in an extend batch. Used by
|
||||||
|
``finalize_mlx_result`` to commit per-request caches. Empty
|
||||||
|
list for pure-decode steps.
|
||||||
|
extends: Chunked-prefill-continuation state, one entry per
|
||||||
|
already-active request whose extend seq_len > 1. Also empty
|
||||||
|
for pure-decode steps.
|
||||||
|
decode: Decode state covering full-decode mode AND mixed
|
||||||
|
single-token decodes inside an extend batch. Used as the
|
||||||
|
chaining root by :meth:`async_chained_decode_mlx`.
|
||||||
|
mode: One of ``"decode"``, ``"extend"``, ``"idle"`` describing
|
||||||
|
which forward pass produced this job. Drives finalise
|
||||||
|
dispatch and whether chaining is safe.
|
||||||
|
batch_copy: Snapshot of the :class:`ScheduleBatch` at launch
|
||||||
|
time. Decoupled from the live batch so
|
||||||
|
``process_batch_result`` can update request state without
|
||||||
|
racing against the next scheduling decision.
|
||||||
|
reqs: Snapshot of ``batch.reqs`` at launch time. The overlap
|
||||||
|
loop uses this to check ``req.finished()`` on the previous
|
||||||
|
step's request list without holding a reference to the
|
||||||
|
mutable batch object.
|
||||||
|
"""
|
||||||
|
|
||||||
|
lazy_tokens: Optional[mx.array]
|
||||||
|
prefills: list["MlxPendingPrefill"]
|
||||||
|
extends: list["MlxPendingExtend"]
|
||||||
|
decode: Optional["MlxPendingDecode"]
|
||||||
|
mode: str
|
||||||
|
batch_copy: "ScheduleBatch"
|
||||||
|
reqs: List[Req]
|
||||||
|
|
||||||
|
|
||||||
|
class SchedulerMlxOverlapMixin:
|
||||||
|
"""Mixin that adds MLX overlap scheduling to :class:`Scheduler`."""
|
||||||
|
|
||||||
|
@DynamicGradMode()
|
||||||
|
def event_loop_overlap_mlx(self: "Scheduler"):
|
||||||
|
"""MLX-specific overlap loop modelled on ``mlx_lm.generate.generate_step``.
|
||||||
|
|
||||||
|
At steady state we keep TWO in-flight MLX graphs queued on the
|
||||||
|
GPU:
|
||||||
|
|
||||||
|
* ``pending_curr`` — the step whose tokens we are about to block
|
||||||
|
on and feed into the scheduler's bookkeeping.
|
||||||
|
* ``pending_next`` — the step that was built on top of
|
||||||
|
``pending_curr``'s still-lazy output tokens via
|
||||||
|
``async_chained_decode_mlx`` and has already been handed to
|
||||||
|
``mx.async_eval``. Because MLX tracks the full dependency
|
||||||
|
graph, the GPU will execute ``pending_next`` back-to-back
|
||||||
|
with ``pending_curr`` — there is no scheduling gap on the
|
||||||
|
device.
|
||||||
|
|
||||||
|
Bookkeeping timeline for a steady-state decode loop:
|
||||||
|
|
||||||
|
iter k:
|
||||||
|
build pending_next (CPU graph build + mx.async_eval; cheap)
|
||||||
|
block on pending_curr via .tolist() (wait only on curr's tokens)
|
||||||
|
process_batch_result(pending_curr) <-- GPU is running pending_next
|
||||||
|
pending_curr = pending_next
|
||||||
|
|
||||||
|
The chain is broken (we fall back to a "schedule + launch" step)
|
||||||
|
whenever any of the following holds:
|
||||||
|
|
||||||
|
* ``pending_curr`` is not a pure decode (e.g. prefill/extend).
|
||||||
|
* The waiting queue has new requests that need prefill.
|
||||||
|
* Any req in ``pending_curr`` just finished this iteration, so
|
||||||
|
the composition for ``pending_next`` would need to shrink.
|
||||||
|
|
||||||
|
When the chain breaks mid-flight we still finalise the
|
||||||
|
already-launched ``pending_next`` normally (its tokens are
|
||||||
|
valid for all surviving reqs). With RadixCache-backed caches
|
||||||
|
(#21509) there is no ``extract_cache`` step: per-request caches
|
||||||
|
are the source of truth and are never merged into a shared
|
||||||
|
batched buffer.
|
||||||
|
"""
|
||||||
|
pending_curr: 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.output_ids = result.next_token_ids
|
||||||
|
self.process_batch_result(pending.batch_copy, result)
|
||||||
|
|
||||||
|
def _launch_fresh(batch: "ScheduleBatch") -> MlxPendingJob:
|
||||||
|
mwb = batch.get_model_worker_batch()
|
||||||
|
lazy_tokens, prefills, extends, decode, mode = (
|
||||||
|
self.tp_worker.async_forward_batch_generation_mlx(mwb)
|
||||||
|
)
|
||||||
|
return MlxPendingJob(
|
||||||
|
lazy_tokens=lazy_tokens,
|
||||||
|
prefills=prefills,
|
||||||
|
extends=extends,
|
||||||
|
decode=decode,
|
||||||
|
mode=mode,
|
||||||
|
batch_copy=batch.copy(),
|
||||||
|
reqs=list(batch.reqs),
|
||||||
|
)
|
||||||
|
|
||||||
|
def _launch_chained(prev: MlxPendingJob) -> MlxPendingJob:
|
||||||
|
assert prev.decode is not None
|
||||||
|
lazy_tokens, prefills, extends, decode, mode = (
|
||||||
|
self.tp_worker.async_chained_decode_mlx(prev.decode)
|
||||||
|
)
|
||||||
|
# Composition is identical to prev: reuse a fresh batch copy
|
||||||
|
# of the same underlying ScheduleBatch so process_batch_result
|
||||||
|
# updates the same req objects with the new token.
|
||||||
|
return MlxPendingJob(
|
||||||
|
lazy_tokens=lazy_tokens,
|
||||||
|
prefills=prefills,
|
||||||
|
extends=extends,
|
||||||
|
decode=decode,
|
||||||
|
mode=mode,
|
||||||
|
batch_copy=prev.batch_copy.copy(),
|
||||||
|
reqs=prev.reqs,
|
||||||
|
)
|
||||||
|
|
||||||
|
while True:
|
||||||
|
recv_reqs = self.recv_requests()
|
||||||
|
self.process_input_requests(recv_reqs)
|
||||||
|
if self._engine_paused:
|
||||||
|
continue
|
||||||
|
|
||||||
|
# 1. If pending_curr is a pure decode AND no new prefill is waiting,
|
||||||
|
# build pending_next on top of it NOW — before we block on curr.
|
||||||
|
can_chain = (
|
||||||
|
pending_curr is not None
|
||||||
|
and pending_curr.mode == "decode"
|
||||||
|
and pending_curr.decode is not None
|
||||||
|
and not self.waiting_queue
|
||||||
|
)
|
||||||
|
if can_chain and pending_next is None:
|
||||||
|
# Build + launch the chained step BEFORE we block on
|
||||||
|
# pending_curr — this is the "no idle gap" trick.
|
||||||
|
# GPU now has 2 steps queued.
|
||||||
|
pending_next = _launch_chained(pending_curr)
|
||||||
|
self.result_queue.append(pending_next)
|
||||||
|
|
||||||
|
# 2. Finalize/process on pending_curr's tokens. (GPU is already
|
||||||
|
# executing pending_next at this point.)
|
||||||
|
if pending_curr is not None:
|
||||||
|
_finalize(pending_curr)
|
||||||
|
self.result_queue.popleft()
|
||||||
|
pending_curr = None
|
||||||
|
|
||||||
|
# 3. Decide whether pending_next is still valid (if no reqs finished)
|
||||||
|
# and promote it.
|
||||||
|
finished_any = any(
|
||||||
|
req.finished() for req in (pending_next.reqs if pending_next else [])
|
||||||
|
)
|
||||||
|
new_prefill_waiting = bool(self.waiting_queue)
|
||||||
|
if (
|
||||||
|
pending_next is not None
|
||||||
|
and not finished_any
|
||||||
|
and not new_prefill_waiting
|
||||||
|
):
|
||||||
|
pending_curr = pending_next
|
||||||
|
pending_next = None
|
||||||
|
self.cur_batch = pending_curr.batch_copy
|
||||||
|
self.last_batch = pending_curr.batch_copy
|
||||||
|
if envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.get():
|
||||||
|
self.self_check_during_busy()
|
||||||
|
continue
|
||||||
|
|
||||||
|
# 4. Chain is broken. Finalise pending_next (if any), then
|
||||||
|
# schedule fresh.
|
||||||
|
if pending_next is not None:
|
||||||
|
_finalize(pending_next)
|
||||||
|
self.result_queue.popleft()
|
||||||
|
pending_next = None
|
||||||
|
next_batch = self.get_next_batch_to_run()
|
||||||
|
self.cur_batch = next_batch
|
||||||
|
if next_batch:
|
||||||
|
pending_curr = _launch_fresh(next_batch)
|
||||||
|
self.result_queue.append(pending_curr)
|
||||||
|
else:
|
||||||
|
self.on_idle()
|
||||||
|
|
||||||
|
self.last_batch = next_batch
|
||||||
|
if envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.get():
|
||||||
|
self.self_check_during_busy()
|
||||||
@@ -3,13 +3,26 @@
|
|||||||
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.
|
KV data lives in MlxKVPool.
|
||||||
|
|
||||||
|
The worker also exposes an async (lazy-eval) surface used by the MLX
|
||||||
|
overlap scheduler: ``async_forward_batch_generation_mlx`` launches a
|
||||||
|
batch without blocking on the GPU, ``async_chained_decode_mlx`` builds
|
||||||
|
the next decode step on top of a still-lazy previous decode, and
|
||||||
|
``finalize_mlx_result`` blocks on the lazy outputs and produces a
|
||||||
|
normal ``GenerationBatchResult``.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
from typing import Optional
|
from typing import Optional, Union
|
||||||
|
|
||||||
|
import mlx.core as mx
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.hardware_backend.mlx.model_runner import (
|
||||||
|
MlxPendingDecode,
|
||||||
|
MlxPendingExtend,
|
||||||
|
MlxPendingPrefill,
|
||||||
|
)
|
||||||
from sglang.srt.managers.schedule_batch import ModelWorkerBatch
|
from sglang.srt.managers.schedule_batch import ModelWorkerBatch
|
||||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||||
from sglang.srt.managers.utils import GenerationBatchResult
|
from sglang.srt.managers.utils import GenerationBatchResult
|
||||||
@@ -100,6 +113,16 @@ class MlxTpModelWorker(TpModelWorker):
|
|||||||
skip_attn_backend_init,
|
skip_attn_backend_init,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _cleanup_stale_rids(self, forward_mode, current_rids: set[str]) -> None:
|
||||||
|
"""Remove MLX state for decode-mode requests that dropped out of the batch."""
|
||||||
|
if forward_mode.is_decode():
|
||||||
|
stale_rids = self._mlx_active_rids - current_rids
|
||||||
|
for rid in stale_rids:
|
||||||
|
self._mlx_runner.remove_request(rid)
|
||||||
|
self._mlx_active_rids = current_rids
|
||||||
|
else:
|
||||||
|
self._mlx_active_rids |= current_rids
|
||||||
|
|
||||||
def _forward_batch_generation_mlx(
|
def _forward_batch_generation_mlx(
|
||||||
self,
|
self,
|
||||||
model_worker_batch: ModelWorkerBatch,
|
model_worker_batch: ModelWorkerBatch,
|
||||||
@@ -116,17 +139,9 @@ class MlxTpModelWorker(TpModelWorker):
|
|||||||
can_run_cuda_graph=False,
|
can_run_cuda_graph=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Auto-cleanup: remove MLX state for requests no longer in the batch.
|
self._cleanup_stale_rids(forward_mode, {req.rid for req in reqs})
|
||||||
current_rids = {req.rid for req in reqs}
|
|
||||||
if forward_mode.is_decode():
|
|
||||||
stale_rids = self._mlx_active_rids - current_rids
|
|
||||||
for rid in stale_rids:
|
|
||||||
self._mlx_runner.remove_request(rid)
|
|
||||||
self._mlx_active_rids = current_rids
|
|
||||||
else:
|
|
||||||
self._mlx_active_rids |= current_rids
|
|
||||||
|
|
||||||
next_token_ids_list = []
|
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 PoolBackedCache reads it
|
||||||
@@ -138,9 +153,9 @@ class MlxTpModelWorker(TpModelWorker):
|
|||||||
|
|
||||||
offset = 0 # into input_ids_cpu
|
offset = 0 # into input_ids_cpu
|
||||||
slot_offset = 0 # into out_cache_loc_cpu
|
slot_offset = 0 # into out_cache_loc_cpu
|
||||||
prefill_rids = []
|
prefill_rids: list[tuple[str, int]] = []
|
||||||
extend_rids = []
|
extend_rids: list[tuple[str, int]] = []
|
||||||
decode_rids = []
|
decode_rids: list[str] = []
|
||||||
|
|
||||||
for i, req in enumerate(reqs):
|
for i, req in enumerate(reqs):
|
||||||
seq_len = extend_seq_lens[i]
|
seq_len = extend_seq_lens[i]
|
||||||
@@ -209,3 +224,262 @@ class MlxTpModelWorker(TpModelWorker):
|
|||||||
next_token_ids=next_token_ids,
|
next_token_ids=next_token_ids,
|
||||||
can_run_cuda_graph=False,
|
can_run_cuda_graph=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def async_forward_batch_generation_mlx(
|
||||||
|
self,
|
||||||
|
model_worker_batch: ModelWorkerBatch,
|
||||||
|
) -> tuple[
|
||||||
|
Union[mx.array, None],
|
||||||
|
list[MlxPendingPrefill],
|
||||||
|
list[MlxPendingExtend],
|
||||||
|
Optional[MlxPendingDecode],
|
||||||
|
str,
|
||||||
|
]:
|
||||||
|
"""Start an async (lazy) forward pass through the MLX model runner.
|
||||||
|
|
||||||
|
Returns ``(lazy_result, prefills, extends, decode, mode)``:
|
||||||
|
|
||||||
|
* ``lazy_result`` — an ``mx.array`` that, when evaluated, forces
|
||||||
|
materialisation of the whole batch's outputs. ``None`` for
|
||||||
|
idle batches.
|
||||||
|
* ``prefills`` — list of :class:`MlxPendingPrefill` for new
|
||||||
|
requests in an extend batch.
|
||||||
|
* ``extends`` — list of :class:`MlxPendingExtend` for chunked
|
||||||
|
prefill continuations in an extend batch.
|
||||||
|
* ``decode`` — :class:`MlxPendingDecode` for the decode
|
||||||
|
sub-batch (covers full decode mode AND mixed decodes inside
|
||||||
|
an extend batch).
|
||||||
|
* ``mode`` — one of ``"idle"``, ``"decode"``, ``"extend"``.
|
||||||
|
|
||||||
|
The caller must make sure the returned pendings are fed into a
|
||||||
|
subsequent ``mx.async_eval`` or ``.item()`` / ``.tolist()`` call
|
||||||
|
— :meth:`finalize_mlx_result` does that.
|
||||||
|
"""
|
||||||
|
self._ensure_mlx_pool_initialized()
|
||||||
|
|
||||||
|
forward_mode = model_worker_batch.forward_mode
|
||||||
|
reqs = model_worker_batch.reqs
|
||||||
|
|
||||||
|
if forward_mode.is_idle():
|
||||||
|
return None, [], [], None, "idle"
|
||||||
|
|
||||||
|
self._cleanup_stale_rids(forward_mode, {req.rid for req in reqs})
|
||||||
|
|
||||||
|
if forward_mode.is_decode():
|
||||||
|
req_ids = [req.rid for req in reqs]
|
||||||
|
pending_decode = self._mlx_runner.decode_batch_start(req_ids)
|
||||||
|
mx.async_eval(pending_decode.lazy_tokens)
|
||||||
|
return pending_decode.lazy_tokens, [], [], pending_decode, "decode"
|
||||||
|
|
||||||
|
if forward_mode.is_extend():
|
||||||
|
# TODO (changminbark): Implement per-batch flushing using prefix_slot_ids
|
||||||
|
# Ensure the pool is up-to-date before any PoolBackedCache
|
||||||
|
# reads it for prefix-cached prefills. Mirror the sync path.
|
||||||
|
self._mlx_runner.flush_all_decode_kv()
|
||||||
|
return self._async_extend_batch(model_worker_batch)
|
||||||
|
|
||||||
|
raise ValueError(
|
||||||
|
f"MLX async runner does not support forward mode: {forward_mode}"
|
||||||
|
)
|
||||||
|
|
||||||
|
def _async_extend_batch(
|
||||||
|
self,
|
||||||
|
model_worker_batch: ModelWorkerBatch,
|
||||||
|
) -> tuple[
|
||||||
|
Union[mx.array, None],
|
||||||
|
list[MlxPendingPrefill],
|
||||||
|
list[MlxPendingExtend],
|
||||||
|
Optional[MlxPendingDecode],
|
||||||
|
str,
|
||||||
|
]:
|
||||||
|
"""Launch each request in an EXTEND batch lazily and kick GPU work."""
|
||||||
|
reqs = model_worker_batch.reqs
|
||||||
|
input_ids_cpu = model_worker_batch.input_ids.cpu().tolist()
|
||||||
|
out_cache_loc_cpu = model_worker_batch.out_cache_loc.cpu().tolist()
|
||||||
|
extend_seq_lens = model_worker_batch.extend_seq_lens
|
||||||
|
|
||||||
|
offset = 0
|
||||||
|
slot_offset = 0
|
||||||
|
pending_prefills: list[MlxPendingPrefill] = []
|
||||||
|
pending_extends: list[MlxPendingExtend] = []
|
||||||
|
mixed_decode_rids: list[str] = []
|
||||||
|
|
||||||
|
for i, req in enumerate(reqs):
|
||||||
|
seq_len = extend_seq_lens[i]
|
||||||
|
req_token_ids = input_ids_cpu[offset : offset + seq_len]
|
||||||
|
req_new_slots = out_cache_loc_cpu[slot_offset : slot_offset + seq_len]
|
||||||
|
offset += seq_len
|
||||||
|
slot_offset += seq_len
|
||||||
|
|
||||||
|
if self._mlx_runner.has_request(req.rid):
|
||||||
|
if seq_len > 1:
|
||||||
|
# Chunked prefill continuation
|
||||||
|
pending_extends.append(
|
||||||
|
self._mlx_runner.extend_start(
|
||||||
|
req_id=req.rid,
|
||||||
|
new_token_ids=req_token_ids,
|
||||||
|
new_slot_ids=req_new_slots,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# MIXED mode: single-token decode
|
||||||
|
mixed_decode_rids.append(req.rid)
|
||||||
|
else:
|
||||||
|
# New prefill
|
||||||
|
prefix_slot_ids = req.prefix_indices.tolist()
|
||||||
|
full_token_ids = list(req.fill_ids)
|
||||||
|
pending_prefills.append(
|
||||||
|
self._mlx_runner.prefill_start(
|
||||||
|
req_id=req.rid,
|
||||||
|
new_token_ids=req_token_ids,
|
||||||
|
full_token_ids=full_token_ids,
|
||||||
|
prefix_slot_ids=prefix_slot_ids,
|
||||||
|
new_slot_ids=req_new_slots,
|
||||||
|
req_pool_idx=req.req_pool_idx,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
pending_mixed_decode: Optional[MlxPendingDecode] = None
|
||||||
|
if mixed_decode_rids:
|
||||||
|
pending_mixed_decode = self._mlx_runner.decode_batch_start(
|
||||||
|
mixed_decode_rids
|
||||||
|
)
|
||||||
|
|
||||||
|
# Stack lazy tokens so the caller has a single handle to evaluate
|
||||||
|
# after CPU scheduling work. We also hand every cache buffer
|
||||||
|
# (and the decode cache arrays) to mx.async_eval so the GPU
|
||||||
|
# kernel-launch stream sees everything the next step depends on
|
||||||
|
# before we actually block on anything.
|
||||||
|
prefill_ext_tokens: list[mx.array] = [p.lazy_token for p in pending_prefills]
|
||||||
|
prefill_ext_tokens.extend(e.lazy_token for e in pending_extends)
|
||||||
|
|
||||||
|
async_args: list[mx.array] = []
|
||||||
|
if prefill_ext_tokens:
|
||||||
|
lazy_stacked = mx.stack(prefill_ext_tokens, axis=0)
|
||||||
|
async_args.append(lazy_stacked)
|
||||||
|
else:
|
||||||
|
lazy_stacked = None
|
||||||
|
|
||||||
|
for p in pending_prefills:
|
||||||
|
async_args.extend(self._cache_state(p.cache))
|
||||||
|
for e in pending_extends:
|
||||||
|
async_args.extend(self._cache_state(self._mlx_runner._req_caches[e.req_id]))
|
||||||
|
if pending_mixed_decode is not None:
|
||||||
|
async_args.append(pending_mixed_decode.lazy_tokens)
|
||||||
|
for c_list in pending_mixed_decode.caches:
|
||||||
|
async_args.extend(self._cache_state(c_list))
|
||||||
|
|
||||||
|
if async_args:
|
||||||
|
mx.async_eval(*async_args)
|
||||||
|
|
||||||
|
return (
|
||||||
|
lazy_stacked,
|
||||||
|
pending_prefills,
|
||||||
|
pending_extends,
|
||||||
|
pending_mixed_decode,
|
||||||
|
"extend",
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _cache_state(cache_list) -> list[mx.array]:
|
||||||
|
"""Flatten a per-layer cache list to its ``state`` arrays."""
|
||||||
|
return [s for c in cache_list for s in c.state]
|
||||||
|
|
||||||
|
def async_chained_decode_mlx(
|
||||||
|
self,
|
||||||
|
prev_pending: MlxPendingDecode,
|
||||||
|
) -> tuple[mx.array, list, list, MlxPendingDecode, str]:
|
||||||
|
"""Launch a decode step that chains off a still-lazy previous decode.
|
||||||
|
|
||||||
|
This is the "no idle gap" pipelining primitive: build the next
|
||||||
|
decode's compute graph using ``prev_pending.lazy_tokens`` (still
|
||||||
|
unevaluated) as its input ids, hand the combined graph to
|
||||||
|
``mx.async_eval``, and return. The GPU runs the new step
|
||||||
|
immediately after ``prev_pending`` with no scheduling gap, while
|
||||||
|
the caller is free to block on ``prev_pending`` and run CPU-side
|
||||||
|
bookkeeping.
|
||||||
|
|
||||||
|
Preconditions (caller must ensure):
|
||||||
|
|
||||||
|
* ``prev_pending`` was produced by a previous decode start
|
||||||
|
(either :meth:`async_forward_batch_generation_mlx` in decode
|
||||||
|
mode or a previous :meth:`async_chained_decode_mlx`).
|
||||||
|
* The batch composition for this step is identical to
|
||||||
|
``prev_pending`` — same requests, same order. Composition
|
||||||
|
changes (finished reqs, new prefills) must break the chain.
|
||||||
|
* ``prev_pending`` should be finalised BEFORE the returned
|
||||||
|
pending, so per-request token lists are appended in order.
|
||||||
|
|
||||||
|
Returns a 5-tuple matching
|
||||||
|
:meth:`async_forward_batch_generation_mlx` for the decode case:
|
||||||
|
``(lazy_tokens, [], [], pending_decode, "decode")``. The empty
|
||||||
|
prefill/extend lists are always absent for chained decodes.
|
||||||
|
"""
|
||||||
|
pending = self._mlx_runner.decode_batch_start_chained(prev_pending)
|
||||||
|
mx.async_eval(pending.lazy_tokens)
|
||||||
|
return pending.lazy_tokens, [], [], pending, "decode"
|
||||||
|
|
||||||
|
def finalize_mlx_result(
|
||||||
|
self,
|
||||||
|
prefills: list[MlxPendingPrefill],
|
||||||
|
extends: list[MlxPendingExtend],
|
||||||
|
decode: Optional[MlxPendingDecode],
|
||||||
|
mode: str,
|
||||||
|
reqs: list,
|
||||||
|
) -> GenerationBatchResult:
|
||||||
|
"""Materialise a lazy MLX result into a :class:`GenerationBatchResult`.
|
||||||
|
|
||||||
|
The blocking wait happens inside ``decode_batch_finalize`` /
|
||||||
|
``prefill_finalize`` / ``extend_finalize`` via ``.tolist()`` /
|
||||||
|
``.item()`` on the specific lazy outputs.
|
||||||
|
"""
|
||||||
|
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||||
|
|
||||||
|
if mode == "idle":
|
||||||
|
return GenerationBatchResult(
|
||||||
|
logits_output=LogitsProcessorOutput(next_token_logits=None),
|
||||||
|
can_run_cuda_graph=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
if mode == "decode":
|
||||||
|
assert decode is not None
|
||||||
|
next_tokens_list = self._mlx_runner.decode_batch_finalize(decode)
|
||||||
|
|
||||||
|
elif mode == "extend":
|
||||||
|
prefill_map: dict[str, int] = {}
|
||||||
|
for pending_p in prefills:
|
||||||
|
prefill_map[pending_p.req_id] = self._mlx_runner.prefill_finalize(
|
||||||
|
pending_p
|
||||||
|
)
|
||||||
|
|
||||||
|
extend_map: dict[str, int] = {}
|
||||||
|
for pending_e in extends:
|
||||||
|
extend_map[pending_e.req_id] = self._mlx_runner.extend_finalize(
|
||||||
|
pending_e
|
||||||
|
)
|
||||||
|
|
||||||
|
decode_map: dict[str, int] = {}
|
||||||
|
if decode is not None:
|
||||||
|
mixed_tokens = self._mlx_runner.decode_batch_finalize(decode)
|
||||||
|
decode_map = {
|
||||||
|
rid: tok for rid, tok in zip(decode.req_ids, mixed_tokens)
|
||||||
|
}
|
||||||
|
|
||||||
|
next_tokens_list = []
|
||||||
|
for req in reqs:
|
||||||
|
if req.rid in decode_map:
|
||||||
|
next_tokens_list.append(decode_map[req.rid])
|
||||||
|
elif req.rid in extend_map:
|
||||||
|
next_tokens_list.append(extend_map[req.rid])
|
||||||
|
else:
|
||||||
|
next_tokens_list.append(prefill_map[req.rid])
|
||||||
|
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unknown MLX async mode: {mode}")
|
||||||
|
|
||||||
|
next_token_ids = torch.tensor(next_tokens_list, dtype=torch.long, device="cpu")
|
||||||
|
return GenerationBatchResult(
|
||||||
|
logits_output=LogitsProcessorOutput(next_token_logits=None),
|
||||||
|
next_token_ids=next_token_ids,
|
||||||
|
can_run_cuda_graph=False,
|
||||||
|
)
|
||||||
|
|||||||
@@ -239,9 +239,14 @@ from sglang.utils import TypeBasedDispatcher, get_exception_traceback
|
|||||||
|
|
||||||
if is_mps():
|
if is_mps():
|
||||||
CudaStreamContext = nullcontext
|
CudaStreamContext = nullcontext
|
||||||
|
from sglang.srt.hardware_backend.mlx.scheduler_mixin import SchedulerMlxOverlapMixin
|
||||||
else:
|
else:
|
||||||
from torch.cuda import StreamContext as CudaStreamContext
|
from torch.cuda import StreamContext as CudaStreamContext
|
||||||
|
|
||||||
|
class SchedulerMlxOverlapMixin:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
# Test retract decode for debugging purposes
|
# Test retract decode for debugging purposes
|
||||||
@@ -326,6 +331,7 @@ class Scheduler(
|
|||||||
SchedulerPPMixin,
|
SchedulerPPMixin,
|
||||||
SchedulerDPAttnMixin,
|
SchedulerDPAttnMixin,
|
||||||
SchedulerDllmMixin,
|
SchedulerDllmMixin,
|
||||||
|
SchedulerMlxOverlapMixin,
|
||||||
):
|
):
|
||||||
"""A scheduler that manages a tensor parallel GPU worker."""
|
"""A scheduler that manages a tensor parallel GPU worker."""
|
||||||
|
|
||||||
@@ -373,7 +379,8 @@ class Scheduler(
|
|||||||
self.enable_lora = server_args.enable_lora
|
self.enable_lora = server_args.enable_lora
|
||||||
self.enable_lora_overlap_loading = server_args.enable_lora_overlap_loading
|
self.enable_lora_overlap_loading = server_args.enable_lora_overlap_loading
|
||||||
self.max_loras_per_batch = server_args.max_loras_per_batch
|
self.max_loras_per_batch = server_args.max_loras_per_batch
|
||||||
self.enable_overlap = not server_args.disable_overlap_schedule
|
self.enable_overlap = not server_args.disable_overlap_schedule and not use_mlx()
|
||||||
|
self.enable_overlap_mlx = not server_args.disable_overlap_schedule and use_mlx()
|
||||||
self.enable_pdmux = server_args.enable_pdmux
|
self.enable_pdmux = server_args.enable_pdmux
|
||||||
self.skip_tokenizer_init = server_args.skip_tokenizer_init
|
self.skip_tokenizer_init = server_args.skip_tokenizer_init
|
||||||
self.stream_interval = server_args.stream_interval
|
self.stream_interval = server_args.stream_interval
|
||||||
@@ -1247,6 +1254,15 @@ class Scheduler(
|
|||||||
def init_overlap(self):
|
def init_overlap(self):
|
||||||
self.device_module = torch.get_device_module(self.device)
|
self.device_module = torch.get_device_module(self.device)
|
||||||
|
|
||||||
|
if use_mlx():
|
||||||
|
# MLX overlap scheduling uses mx.async_eval / mx.eval for
|
||||||
|
# synchronisation so no CUDA/MPS streams or FutureMap needed.
|
||||||
|
self.future_map = None
|
||||||
|
# Empty result_queue is needed because idle-check references it
|
||||||
|
# when enable_overlap is True.
|
||||||
|
self.result_queue: Deque = deque()
|
||||||
|
return
|
||||||
|
|
||||||
self.forward_stream_ctx: CudaStreamContext = self.device_module.stream(
|
self.forward_stream_ctx: CudaStreamContext = self.device_module.stream(
|
||||||
self.forward_stream
|
self.forward_stream
|
||||||
)
|
)
|
||||||
@@ -1437,6 +1453,12 @@ class Scheduler(
|
|||||||
Sets up the schedule stream and dispatches to the appropriate event loop.
|
Sets up the schedule stream and dispatches to the appropriate event loop.
|
||||||
The event loop blocks until shutdown.
|
The event loop blocks until shutdown.
|
||||||
"""
|
"""
|
||||||
|
if use_mlx():
|
||||||
|
# MLX overlap uses mx.async_eval for CPU/GPU overlap,
|
||||||
|
# not PyTorch MPS streams.
|
||||||
|
dispatch_event_loop(self)
|
||||||
|
return
|
||||||
|
|
||||||
self.schedule_stream = self.device_module.Stream(priority=0)
|
self.schedule_stream = self.device_module.Stream(priority=0)
|
||||||
if self.device == "cpu":
|
if self.device == "cpu":
|
||||||
self.schedule_stream.synchronize = lambda: None # No-op for CPU
|
self.schedule_stream.synchronize = lambda: None # No-op for CPU
|
||||||
@@ -1523,6 +1545,7 @@ class Scheduler(
|
|||||||
|
|
||||||
# Update last_batch
|
# Update last_batch
|
||||||
self.last_batch = batch
|
self.last_batch = batch
|
||||||
|
|
||||||
if envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.get():
|
if envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.get():
|
||||||
self.self_check_during_busy()
|
self.self_check_during_busy()
|
||||||
|
|
||||||
@@ -2874,7 +2897,6 @@ class Scheduler(
|
|||||||
model_worker_batch.sampling_info = (
|
model_worker_batch.sampling_info = (
|
||||||
model_worker_batch.sampling_info.copy_for_forward()
|
model_worker_batch.sampling_info.copy_for_forward()
|
||||||
)
|
)
|
||||||
|
|
||||||
bs = len(model_worker_batch.seq_lens)
|
bs = len(model_worker_batch.seq_lens)
|
||||||
future_indices = self.future_map.alloc_future_indices(bs)
|
future_indices = self.future_map.alloc_future_indices(bs)
|
||||||
|
|
||||||
@@ -3741,6 +3763,8 @@ def dispatch_event_loop(scheduler: Scheduler):
|
|||||||
scheduler.event_loop_pdmux()
|
scheduler.event_loop_pdmux()
|
||||||
elif server_args.pp_size > 1:
|
elif server_args.pp_size > 1:
|
||||||
scheduler.event_loop_pp()
|
scheduler.event_loop_pp()
|
||||||
|
elif scheduler.enable_overlap_mlx:
|
||||||
|
scheduler.event_loop_overlap_mlx()
|
||||||
elif scheduler.enable_overlap:
|
elif scheduler.enable_overlap:
|
||||||
scheduler.event_loop_overlap()
|
scheduler.event_loop_overlap()
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -409,6 +409,8 @@ class SchedulerOutputProcessorMixin:
|
|||||||
if batch.spec_algorithm.is_none() or batch.is_spec_v2:
|
if batch.spec_algorithm.is_none() or batch.is_spec_v2:
|
||||||
if batch.is_spec_v2:
|
if batch.is_spec_v2:
|
||||||
next_token_ids = self._resolve_spec_overlap_token_ids(result, batch)
|
next_token_ids = self._resolve_spec_overlap_token_ids(result, batch)
|
||||||
|
elif isinstance(next_token_ids, list):
|
||||||
|
pass # MLX path: already a list[int], skip torch round-trip
|
||||||
else:
|
else:
|
||||||
next_token_ids = next_token_ids.tolist()
|
next_token_ids = next_token_ids.tolist()
|
||||||
|
|
||||||
@@ -447,7 +449,9 @@ class SchedulerOutputProcessorMixin:
|
|||||||
for i, req in enumerate(batch.reqs):
|
for i, req in enumerate(batch.reqs):
|
||||||
req: Req
|
req: Req
|
||||||
|
|
||||||
if self.enable_overlap and (req.finished() or req.is_retracted):
|
if (self.enable_overlap or self.enable_overlap_mlx) and (
|
||||||
|
req.finished() or req.is_retracted
|
||||||
|
):
|
||||||
# NOTE: This (req.finished() or req.is_retracted) should only happen when overlap scheduling is enabled.
|
# NOTE: This (req.finished() or req.is_retracted) should only happen when overlap scheduling is enabled.
|
||||||
# And all the over-allocated tokens will be freed in `release_kv_cache`.
|
# And all the over-allocated tokens will be freed in `release_kv_cache`.
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -70,6 +70,7 @@ from sglang.srt.utils.common import (
|
|||||||
from sglang.srt.utils.hf_transformers_utils import check_gguf_file
|
from sglang.srt.utils.hf_transformers_utils import check_gguf_file
|
||||||
from sglang.srt.utils.network import NetworkAddress, get_free_port, wait_port_available
|
from sglang.srt.utils.network import NetworkAddress, get_free_port, wait_port_available
|
||||||
from sglang.srt.utils.runai_utils import ObjectStorageModel, is_runai_obj_uri
|
from sglang.srt.utils.runai_utils import ObjectStorageModel, is_runai_obj_uri
|
||||||
|
from sglang.srt.utils.tensor_bridge import use_mlx
|
||||||
from sglang.utils import is_in_ci
|
from sglang.utils import is_in_ci
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -1164,6 +1165,7 @@ class ServerArgs:
|
|||||||
|
|
||||||
def _handle_mps_backends(self):
|
def _handle_mps_backends(self):
|
||||||
if self.device == "mps":
|
if self.device == "mps":
|
||||||
|
if not use_mlx():
|
||||||
self.disable_overlap_schedule = True
|
self.disable_overlap_schedule = True
|
||||||
|
|
||||||
def _handle_xpu_backends(self):
|
def _handle_xpu_backends(self):
|
||||||
|
|||||||
Reference in New Issue
Block a user