[MLX] Add native MLX execution backend for Apple Silicon Mac (#20342)
Signed-off-by: Xiaodong Ye <yeahdongcn@gmail.com>
This commit is contained in:
@@ -142,6 +142,9 @@ diffusion_mps = [
|
||||
"cache-dit==1.2.3",
|
||||
"addict==2.4.0",
|
||||
"av==16.1.0",
|
||||
"scikit-image==0.25.2",
|
||||
"trimesh>=4.0.0",
|
||||
"xatlas",
|
||||
]
|
||||
|
||||
test = [
|
||||
|
||||
@@ -88,6 +88,7 @@ from sglang.srt.utils import (
|
||||
suppress_other_loggers,
|
||||
)
|
||||
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
|
||||
from sglang.srt.utils.tensor_bridge import use_mlx
|
||||
|
||||
|
||||
def start_profile(profile_activities, profile_record_shapes=False, rank_print=print):
|
||||
@@ -262,7 +263,7 @@ def load_model(server_args, port_args, gpu_id, tp_rank):
|
||||
moe_ep_rank = tp_rank // (server_args.tp_size // server_args.ep_size)
|
||||
|
||||
model_config = ModelConfig.from_server_args(server_args)
|
||||
model_runner = ModelRunner(
|
||||
runner_kwargs = dict(
|
||||
model_config=model_config,
|
||||
mem_fraction_static=server_args.mem_fraction_static,
|
||||
gpu_id=gpu_id,
|
||||
@@ -275,6 +276,16 @@ def load_model(server_args, port_args, gpu_id, tp_rank):
|
||||
nccl_port=port_args.nccl_port,
|
||||
server_args=server_args,
|
||||
)
|
||||
|
||||
_use_mlx = use_mlx()
|
||||
if _use_mlx:
|
||||
from sglang.srt.hardware_backend.mlx.model_runner_stub import (
|
||||
MlxModelRunnerStub,
|
||||
)
|
||||
|
||||
model_runner = MlxModelRunnerStub(**runner_kwargs)
|
||||
else:
|
||||
model_runner = ModelRunner(**runner_kwargs)
|
||||
rank_print(f"max_total_num_tokens={model_runner.max_total_num_tokens}")
|
||||
tokenizer = get_tokenizer(
|
||||
server_args.tokenizer_path,
|
||||
@@ -283,6 +294,12 @@ def load_model(server_args, port_args, gpu_id, tp_rank):
|
||||
)
|
||||
if server_args.tp_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
if _use_mlx:
|
||||
model_runner = _MlxBenchRunner(model_runner, server_args)
|
||||
else:
|
||||
model_runner = _TorchBenchRunner(model_runner)
|
||||
|
||||
return model_runner, tokenizer
|
||||
|
||||
|
||||
@@ -337,11 +354,12 @@ def prepare_extend_inputs_for_correctness_test(
|
||||
for i in range(len(reqs)):
|
||||
req: Req = reqs[i]
|
||||
req.fill_ids += input_ids[i][bench_args.cut_len :]
|
||||
req.prefix_indices = model_runner.req_to_token_pool.req_to_token[
|
||||
i, : bench_args.cut_len
|
||||
].to(req.prefix_indices.dtype)
|
||||
req.logprob_start_len = -1
|
||||
req.set_extend_input_len(len(req.fill_ids) - len(req.prefix_indices))
|
||||
if model_runner is not None:
|
||||
req.prefix_indices = model_runner.req_to_token_pool.req_to_token[
|
||||
i, : bench_args.cut_len
|
||||
].to(req.prefix_indices.dtype)
|
||||
req.logprob_start_len = -1
|
||||
req.set_extend_input_len(len(req.fill_ids) - len(req.prefix_indices))
|
||||
return reqs
|
||||
|
||||
|
||||
@@ -445,6 +463,69 @@ def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner):
|
||||
)
|
||||
|
||||
|
||||
class _TorchBenchRunner:
|
||||
"""Wraps ModelRunner for the standard PyTorch benchmark path."""
|
||||
|
||||
def __init__(self, model_runner):
|
||||
self.torch_runner = model_runner
|
||||
|
||||
def clear(self):
|
||||
self.torch_runner.req_to_token_pool.clear()
|
||||
self.torch_runner.token_to_kv_pool_allocator.clear()
|
||||
|
||||
def extend(self, reqs):
|
||||
return extend(reqs, self.torch_runner)
|
||||
|
||||
def decode(self, next_token_ids, batch):
|
||||
return decode(next_token_ids, batch, self.torch_runner)
|
||||
|
||||
def cleanup(self, batch):
|
||||
pass
|
||||
|
||||
def synchronize(self):
|
||||
synchronize(self.torch_runner.device)
|
||||
|
||||
def max_batch_size(self, input_len, output_len):
|
||||
return self.torch_runner.max_total_num_tokens // (input_len + output_len)
|
||||
|
||||
|
||||
class _MlxBenchRunner:
|
||||
"""Wraps MlxModelRunner for the MLX benchmark path."""
|
||||
|
||||
def __init__(self, model_runner, server_args):
|
||||
from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner
|
||||
|
||||
self.mlx_runner = MlxModelRunner(
|
||||
model_path=server_args.model_path,
|
||||
trust_remote_code=server_args.trust_remote_code,
|
||||
)
|
||||
self.fake_torch_runner = model_runner
|
||||
|
||||
def clear(self):
|
||||
self.mlx_runner.clear()
|
||||
|
||||
def extend(self, reqs):
|
||||
req_ids = [str(req.rid) for req in reqs]
|
||||
token_ids_list = [[int(t) for t in req.fill_ids] for req in reqs]
|
||||
next_token_ids = self.mlx_runner.prefill_batch(req_ids, token_ids_list)
|
||||
return torch.tensor(next_token_ids), None, req_ids
|
||||
|
||||
def decode(self, next_token_ids, req_ids):
|
||||
next_token_ids = self.mlx_runner.decode_batch(req_ids)
|
||||
return torch.tensor(next_token_ids), None
|
||||
|
||||
def cleanup(self, batch):
|
||||
if isinstance(batch, list):
|
||||
for req_id in batch:
|
||||
self.mlx_runner.remove_request(req_id)
|
||||
|
||||
def synchronize(self):
|
||||
pass
|
||||
|
||||
def max_batch_size(self, input_len, output_len):
|
||||
return self.fake_torch_runner.max_total_num_tokens // (input_len + output_len)
|
||||
|
||||
|
||||
def _read_prompts_from_file(prompt_file, rank_print):
|
||||
"""Read custom prompts from the file specified by `--prompt-filename`."""
|
||||
if not prompt_file:
|
||||
@@ -504,26 +585,30 @@ def correctness_test(
|
||||
|
||||
if bench_args.cut_len > 0:
|
||||
# Prefill
|
||||
next_token_ids, next_token_logits, batch = extend(reqs, model_runner)
|
||||
next_token_ids, next_token_logits, batch = model_runner.extend(reqs)
|
||||
rank_print(f"prefill logits (first half): {next_token_logits} \n")
|
||||
|
||||
# Prepare extend inputs
|
||||
reqs = prepare_extend_inputs_for_correctness_test(
|
||||
bench_args, input_ids, reqs, model_runner
|
||||
)
|
||||
# Prepare extend inputs
|
||||
torch_runner = getattr(model_runner, "torch_runner", None)
|
||||
reqs = prepare_extend_inputs_for_correctness_test(
|
||||
bench_args, input_ids, reqs, torch_runner
|
||||
)
|
||||
|
||||
# Extend (prefill w/ KV cache)
|
||||
next_token_ids, next_token_logits, batch = extend(reqs, model_runner)
|
||||
next_token_ids, next_token_logits, batch = model_runner.extend(reqs)
|
||||
rank_print(f"prefill logits (final): {next_token_logits} \n")
|
||||
|
||||
# Decode
|
||||
output_ids = [input_ids[i] + [next_token_ids[i]] for i in range(len(input_ids))]
|
||||
for _ in range(bench_args.output_len[0] - 1):
|
||||
next_token_ids, _ = decode(next_token_ids, batch, model_runner)
|
||||
next_token_ids, _ = model_runner.decode(next_token_ids, batch)
|
||||
next_token_ids_list = next_token_ids.tolist()
|
||||
for i in range(len(reqs)):
|
||||
output_ids[i].append(next_token_ids_list[i])
|
||||
|
||||
# Clean up
|
||||
model_runner.cleanup(batch)
|
||||
|
||||
# Print output texts
|
||||
for i in range(len(reqs)):
|
||||
rank_print(f"========== Prompt {i} ==========")
|
||||
@@ -542,7 +627,6 @@ def latency_test_run_once(
|
||||
batch_size,
|
||||
input_len,
|
||||
output_len,
|
||||
device,
|
||||
log_decode_step,
|
||||
profile,
|
||||
profile_record_shapes,
|
||||
@@ -553,15 +637,14 @@ def latency_test_run_once(
|
||||
profile_start_step=None,
|
||||
profile_steps=None,
|
||||
):
|
||||
max_batch_size = model_runner.max_total_num_tokens // (input_len + output_len)
|
||||
max_batch_size = model_runner.max_batch_size(input_len, output_len)
|
||||
if batch_size > max_batch_size:
|
||||
rank_print(
|
||||
f"skipping ({batch_size}, {input_len}, {output_len}) due to max batch size limit"
|
||||
)
|
||||
return
|
||||
|
||||
model_runner.req_to_token_pool.clear()
|
||||
model_runner.token_to_kv_pool_allocator.clear()
|
||||
model_runner.clear()
|
||||
|
||||
measurement_results = {
|
||||
"run_name": run_name,
|
||||
@@ -581,10 +664,10 @@ def latency_test_run_once(
|
||||
rank_print=rank_print,
|
||||
)
|
||||
|
||||
synchronize(device)
|
||||
model_runner.synchronize()
|
||||
tic = time.perf_counter()
|
||||
next_token_ids, _, batch = extend(reqs, model_runner)
|
||||
synchronize(device)
|
||||
next_token_ids, _, batch = model_runner.extend(reqs)
|
||||
model_runner.synchronize()
|
||||
prefill_latency = time.perf_counter() - tic
|
||||
|
||||
if enable_profile_prefill:
|
||||
@@ -617,7 +700,7 @@ def latency_test_run_once(
|
||||
enable_profile_decode = profile and profile_stage in ["all", "decode"]
|
||||
profiler = None
|
||||
for i in range(output_len - 1):
|
||||
synchronize(device)
|
||||
model_runner.synchronize()
|
||||
# Start profiler at the specified step
|
||||
if enable_profile_decode and i == profile_start:
|
||||
profiler = start_profile(
|
||||
@@ -627,8 +710,8 @@ def latency_test_run_once(
|
||||
)
|
||||
|
||||
tic = time.perf_counter()
|
||||
next_token_ids, _ = decode(next_token_ids, batch, model_runner)
|
||||
synchronize(device)
|
||||
next_token_ids, _ = model_runner.decode(next_token_ids, batch)
|
||||
model_runner.synchronize()
|
||||
latency = time.perf_counter() - tic
|
||||
|
||||
# Stop profiler after the specified number of steps
|
||||
@@ -670,6 +753,8 @@ def latency_test_run_once(
|
||||
)
|
||||
measurement_results["total_latency"] = tot_latency
|
||||
measurement_results["overall_throughput"] = throughput
|
||||
|
||||
model_runner.cleanup(batch)
|
||||
return measurement_results
|
||||
|
||||
|
||||
@@ -712,7 +797,6 @@ def latency_test(
|
||||
bench_args.batch_size[0],
|
||||
bench_args.input_len[0],
|
||||
min(32, bench_args.output_len[0]), # shorter decoding to speed up the warmup
|
||||
server_args.device,
|
||||
log_decode_step=0,
|
||||
profile=False,
|
||||
profile_record_shapes=False,
|
||||
@@ -764,7 +848,6 @@ def latency_test(
|
||||
bs,
|
||||
il,
|
||||
ol,
|
||||
server_args.device,
|
||||
bench_args.log_decode_step,
|
||||
bench_args.profile if tp_rank == 0 else None,
|
||||
bench_args.profile_record_shapes if tp_rank == 0 else None,
|
||||
|
||||
@@ -15,51 +15,13 @@ from typing import Optional
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.utils.tensor_bridge import mlx_to_torch, torch_to_mlx, use_mlx
|
||||
|
||||
# MLX acceleration – opt-in via SGLANG_USE_MLX=1
|
||||
_MLX_AVAILABLE = False
|
||||
try:
|
||||
_use_mlx = use_mlx()
|
||||
|
||||
if _use_mlx:
|
||||
import mlx.core as mx
|
||||
|
||||
_MLX_AVAILABLE = True
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
_USE_MLX = envs.SGLANG_USE_MLX.get() and _MLX_AVAILABLE
|
||||
|
||||
# Dtype mapping for torch <-> MLX tensor bridge
|
||||
_TORCH_TO_MLX_DTYPE = (
|
||||
{
|
||||
torch.float32: mx.float32,
|
||||
torch.float16: mx.float16,
|
||||
torch.bfloat16: mx.bfloat16,
|
||||
}
|
||||
if _MLX_AVAILABLE
|
||||
else {}
|
||||
)
|
||||
|
||||
_MLX_TO_TORCH_DTYPE = {v: k for k, v in _TORCH_TO_MLX_DTYPE.items()}
|
||||
|
||||
|
||||
def _torch_to_mlx(tensor: torch.Tensor) -> "mx.array":
|
||||
"""Convert a PyTorch tensor to an MLX array (via numpy on CPU)."""
|
||||
t = tensor.cpu().detach()
|
||||
if t.dtype == torch.bfloat16:
|
||||
return mx.array(t.float().numpy(), dtype=mx.bfloat16)
|
||||
return mx.array(t.numpy())
|
||||
|
||||
|
||||
def _mlx_to_torch(array: "mx.array", device: torch.device) -> torch.Tensor:
|
||||
"""Convert an MLX array to a PyTorch tensor (zero-copy via memoryview)."""
|
||||
torch_dtype = _MLX_TO_TORCH_DTYPE.get(array.dtype, torch.float32)
|
||||
array = mx.contiguous(array)
|
||||
mx.eval(array)
|
||||
tensor = torch.frombuffer(memoryview(array), dtype=torch_dtype).reshape(array.shape)
|
||||
if device.type == "mps":
|
||||
tensor = tensor.to(device)
|
||||
return tensor
|
||||
|
||||
|
||||
def fuse_scale_shift_kernel_native(
|
||||
x: torch.Tensor,
|
||||
@@ -197,7 +159,7 @@ def rms_norm_fn_native(
|
||||
# Uses mx.fast.rms_norm / mx.fast.layer_norm — single fused Metal kernels
|
||||
# instead of 7+ separate PyTorch MPS kernel launches.
|
||||
|
||||
if _USE_MLX:
|
||||
if _use_mlx:
|
||||
|
||||
def norm_infer_native( # noqa: F811
|
||||
x: Tensor,
|
||||
@@ -210,17 +172,17 @@ if _USE_MLX:
|
||||
"""MLX-accelerated norm_infer (layer norm / rms norm inference)."""
|
||||
device = x.device
|
||||
orig_dtype = x.dtype
|
||||
x_mx = _torch_to_mlx(x)
|
||||
x_mx = torch_to_mlx(x)
|
||||
if is_rms_norm:
|
||||
w_mx = (
|
||||
_torch_to_mlx(weight) if weight is not None else mx.ones(x_mx.shape[-1])
|
||||
torch_to_mlx(weight) if weight is not None else mx.ones(x_mx.shape[-1])
|
||||
)
|
||||
result_mx = mx.fast.rms_norm(x_mx, w_mx, eps)
|
||||
else:
|
||||
w_mx = _torch_to_mlx(weight) if weight is not None else None
|
||||
b_mx = _torch_to_mlx(bias) if bias is not None else None
|
||||
w_mx = torch_to_mlx(weight) if weight is not None else None
|
||||
b_mx = torch_to_mlx(bias) if bias is not None else None
|
||||
result_mx = mx.fast.layer_norm(x_mx, w_mx, b_mx, eps)
|
||||
result = _mlx_to_torch(result_mx, device).to(orig_dtype)
|
||||
result = mlx_to_torch(result_mx, device).to(orig_dtype)
|
||||
if out is not None:
|
||||
out.copy_(result)
|
||||
return out
|
||||
@@ -230,13 +192,12 @@ if _USE_MLX:
|
||||
x: torch.Tensor, w: torch.Tensor, eps: float = 1e-6
|
||||
) -> torch.Tensor:
|
||||
"""MLX-accelerated triton_one_pass_rms_norm."""
|
||||
shape = x.shape
|
||||
device = x.device
|
||||
orig_dtype = x.dtype
|
||||
x_mx = _torch_to_mlx(x.reshape(-1, x.shape[-1]))
|
||||
w_mx = _torch_to_mlx(w)
|
||||
x_mx = torch_to_mlx(x)
|
||||
w_mx = torch_to_mlx(w)
|
||||
result_mx = mx.fast.rms_norm(x_mx, w_mx, eps)
|
||||
return _mlx_to_torch(result_mx, device).to(orig_dtype).view(shape)
|
||||
return mlx_to_torch(result_mx, device).to(orig_dtype)
|
||||
|
||||
def rms_norm_fn_native( # noqa: F811
|
||||
x,
|
||||
@@ -258,30 +219,25 @@ if _USE_MLX:
|
||||
residual_out=None,
|
||||
):
|
||||
"""MLX-accelerated rms_norm_fn (inference only, no dropout/x1 support)."""
|
||||
x_shape_og = x.shape
|
||||
device = x.device
|
||||
orig_dtype = x.dtype
|
||||
x_flat = x.reshape(-1, x.shape[-1])
|
||||
if residual is not None:
|
||||
residual = residual.reshape(-1, residual.shape[-1]).float()
|
||||
x_flat = x_flat.float() + residual
|
||||
residual_out_val = x_flat.to(
|
||||
torch.float32 if residual_in_fp32 else orig_dtype
|
||||
)
|
||||
x = x.float() + residual.float()
|
||||
residual_out_val = x.to(torch.float32 if residual_in_fp32 else orig_dtype)
|
||||
else:
|
||||
residual_out_val = None
|
||||
if weight is not None and zero_centered_weight:
|
||||
w = weight.float() + 1.0
|
||||
else:
|
||||
w = weight
|
||||
x_mx = _torch_to_mlx(x_flat)
|
||||
w_mx = _torch_to_mlx(w) if w is not None else mx.ones(x_mx.shape[-1])
|
||||
x_mx = torch_to_mlx(x)
|
||||
w_mx = torch_to_mlx(w) if w is not None else mx.ones(x_mx.shape[-1])
|
||||
result_mx = mx.fast.rms_norm(x_mx, w_mx, eps)
|
||||
x_hat = _mlx_to_torch(result_mx, device)
|
||||
x_hat = mlx_to_torch(result_mx, device)
|
||||
if bias is not None:
|
||||
x_hat = x_hat + bias.to(x_hat.device, x_hat.dtype)
|
||||
final_dtype = out_dtype if out_dtype is not None else orig_dtype
|
||||
y = x_hat.to(final_dtype).reshape(x_shape_og)
|
||||
y = x_hat.to(final_dtype)
|
||||
if residual is not None and residual_out_val is not None:
|
||||
return y, residual_out_val.reshape(x_shape_og)
|
||||
return y, residual_out_val
|
||||
return y
|
||||
|
||||
@@ -0,0 +1,306 @@
|
||||
"""End-to-end MLX model runner for Apple Silicon.
|
||||
|
||||
Runs the entire model within MLX, bypassing PyTorch MPS entirely.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
|
||||
import mlx.core as mx
|
||||
from mlx_lm import load as mlx_lm_load
|
||||
from mlx_lm.models.cache import (
|
||||
BatchKVCache,
|
||||
BatchRotatingKVCache,
|
||||
KVCache,
|
||||
RotatingKVCache,
|
||||
make_prompt_cache,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class MlxRequestState:
|
||||
"""Per-request state for MLX inference."""
|
||||
|
||||
token_ids: list[int]
|
||||
cache: list # List of KVCache per layer
|
||||
generated_tokens: int = 0
|
||||
|
||||
|
||||
def _merge_kv_caches(
|
||||
caches_list: list[list],
|
||||
) -> list:
|
||||
"""Merge multiple per-request caches into batched caches."""
|
||||
if not caches_list:
|
||||
return []
|
||||
|
||||
num_layers = len(caches_list[0])
|
||||
merged = []
|
||||
|
||||
for layer_idx in range(num_layers):
|
||||
layer_caches = [caches[layer_idx] for caches in caches_list]
|
||||
if isinstance(layer_caches[0], KVCache):
|
||||
batch_cache = BatchKVCache.merge(layer_caches)
|
||||
elif isinstance(layer_caches[0], RotatingKVCache):
|
||||
batch_cache = BatchRotatingKVCache.merge(layer_caches)
|
||||
else:
|
||||
raise TypeError(f"Unsupported cache type: {type(layer_caches[0]).__name__}")
|
||||
merged.append(batch_cache)
|
||||
|
||||
return merged
|
||||
|
||||
|
||||
def _extract_kv_cache(batch_caches: list, idx: int) -> list:
|
||||
"""Extract a single request's cache from batched caches.
|
||||
|
||||
Works with both BatchKVCache (has .extract) and plain KVCache
|
||||
populated with batched data of shape (B, H, L, D).
|
||||
"""
|
||||
extracted = []
|
||||
for cache in batch_caches:
|
||||
if hasattr(cache, "extract"):
|
||||
extracted.append(cache.extract(idx))
|
||||
else:
|
||||
# Plain KVCache with batched data — slice along batch dim
|
||||
new_cache = KVCache()
|
||||
new_cache.keys = mx.contiguous(cache.keys[idx : idx + 1])
|
||||
new_cache.values = mx.contiguous(cache.values[idx : idx + 1])
|
||||
new_cache.offset = cache.offset
|
||||
extracted.append(new_cache)
|
||||
return extracted
|
||||
|
||||
|
||||
class MlxModelRunner:
|
||||
"""Model runner that executes the entire model in MLX.
|
||||
|
||||
This avoids the MPS<->MLX tensor bridge overhead by keeping all
|
||||
computation within MLX.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_path: str,
|
||||
trust_remote_code: bool = False,
|
||||
):
|
||||
self.model_path = model_path
|
||||
self.trust_remote_code = trust_remote_code
|
||||
self.model = None
|
||||
self._request_states: dict[str, MlxRequestState] = {}
|
||||
|
||||
self._load_model()
|
||||
|
||||
@staticmethod
|
||||
def _extract_logits(model_output):
|
||||
"""Extract logits from model output, handling both tuple and direct returns."""
|
||||
if isinstance(model_output, tuple):
|
||||
return model_output[0]
|
||||
return model_output
|
||||
|
||||
def _load_model(self):
|
||||
"""Load model using mlx_lm."""
|
||||
logger.info(f"Loading MLX model: {self.model_path}")
|
||||
start_time = time.time()
|
||||
|
||||
self.model, _ = mlx_lm_load(
|
||||
self.model_path,
|
||||
tokenizer_config={"trust_remote_code": self.trust_remote_code},
|
||||
)
|
||||
|
||||
load_time = time.time() - start_time
|
||||
logger.info(f"MLX model loaded in {load_time:.2f}s")
|
||||
|
||||
def prefill(
|
||||
self,
|
||||
req_id: str,
|
||||
token_ids: list[int],
|
||||
) -> int:
|
||||
"""Run prefill for a single request.
|
||||
|
||||
If a request with the same req_id already has state (e.g. from a
|
||||
previous partial prefill), the existing KV cache is reused and only
|
||||
the new tokens are fed through the model.
|
||||
|
||||
Args:
|
||||
req_id: Request identifier
|
||||
token_ids: Input token IDs (full sequence, including any
|
||||
previously prefilled tokens)
|
||||
|
||||
Returns:
|
||||
Next token ID (greedy sampled)
|
||||
"""
|
||||
existing_state = self._request_states.get(req_id)
|
||||
if existing_state is not None:
|
||||
# Continuation: reuse existing cache, feed only new tokens
|
||||
cached_input_len = (
|
||||
len(existing_state.token_ids) - existing_state.generated_tokens
|
||||
)
|
||||
new_tokens = token_ids[cached_input_len:]
|
||||
cache = existing_state.cache
|
||||
else:
|
||||
new_tokens = token_ids
|
||||
cache = make_prompt_cache(self.model)
|
||||
|
||||
input_ids = mx.array([new_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)
|
||||
|
||||
# Evaluate everything together
|
||||
mx.eval(next_token_mlx, *[c.state for c in cache])
|
||||
next_token = int(next_token_mlx.item())
|
||||
|
||||
# Store state for future decoding
|
||||
self._request_states[req_id] = MlxRequestState(
|
||||
token_ids=list(token_ids) + [next_token],
|
||||
cache=cache,
|
||||
generated_tokens=1,
|
||||
)
|
||||
|
||||
return next_token
|
||||
|
||||
def prefill_batch(
|
||||
self,
|
||||
req_ids: list[str],
|
||||
token_ids_list: list[list[int]],
|
||||
) -> list[int]:
|
||||
"""Run batched prefill for multiple requests in a single forward pass.
|
||||
|
||||
When all sequences have the same length, they are stacked into a single
|
||||
batch tensor for one forward pass. For variable-length sequences the
|
||||
method falls back to serial prefill.
|
||||
|
||||
Args:
|
||||
req_ids: List of request identifiers
|
||||
token_ids_list: List of token ID sequences, one per request
|
||||
|
||||
Returns:
|
||||
List of next token IDs (greedy sampled)
|
||||
"""
|
||||
if len(req_ids) == 1:
|
||||
return [self.prefill(req_ids[0], token_ids_list[0])]
|
||||
|
||||
# Check if all sequences have the same length (enables true batching)
|
||||
lengths = [len(tids) for tids in token_ids_list]
|
||||
if len(set(lengths)) != 1:
|
||||
# Variable lengths – fall back to serial prefill
|
||||
return [
|
||||
self.prefill(rid, tids) for rid, tids in zip(req_ids, token_ids_list)
|
||||
]
|
||||
|
||||
# All same length – use a single set of fresh caches;
|
||||
# they'll be populated with shape (batch_size, ...) on the first forward pass
|
||||
batch_cache = make_prompt_cache(self.model)
|
||||
|
||||
# Stack into (batch_size, seq_len)
|
||||
batched_input = mx.array(
|
||||
[list(tids) for tids in token_ids_list], dtype=mx.int32
|
||||
)
|
||||
|
||||
# Single forward pass
|
||||
model_output = self.model(batched_input, cache=batch_cache)
|
||||
logits = self._extract_logits(model_output)
|
||||
|
||||
last_logits = logits[:, -1, :]
|
||||
next_tokens_mlx = mx.argmax(last_logits, axis=-1)
|
||||
|
||||
# Evaluate everything together
|
||||
mx.eval(next_tokens_mlx, *[c.state for c in batch_cache])
|
||||
next_tokens = next_tokens_mlx.tolist()
|
||||
|
||||
# Extract individual caches and store per-request state
|
||||
for i, req_id in enumerate(req_ids):
|
||||
individual_cache = _extract_kv_cache(batch_cache, i)
|
||||
self._request_states[req_id] = MlxRequestState(
|
||||
token_ids=list(token_ids_list[i]) + [next_tokens[i]],
|
||||
cache=individual_cache,
|
||||
generated_tokens=1,
|
||||
)
|
||||
|
||||
return next_tokens
|
||||
|
||||
def decode_batch(
|
||||
self,
|
||||
req_ids: list[str],
|
||||
) -> list[int]:
|
||||
"""Run batched decode for multiple requests.
|
||||
|
||||
Args:
|
||||
req_ids: List of request IDs to decode
|
||||
|
||||
Returns:
|
||||
List of next token IDs
|
||||
"""
|
||||
if len(req_ids) == 1:
|
||||
return [self._decode_single(req_ids[0])]
|
||||
|
||||
decode_reqs = []
|
||||
for req_id in req_ids:
|
||||
state = self._request_states[req_id]
|
||||
decode_reqs.append((req_id, state))
|
||||
|
||||
return self._batched_decode(decode_reqs)
|
||||
|
||||
def _decode_single(self, req_id: str) -> int:
|
||||
"""Decode a single token for one request."""
|
||||
state = self._request_states[req_id]
|
||||
last_token = state.token_ids[-1]
|
||||
|
||||
input_ids = mx.array([[last_token]], dtype=mx.int32)
|
||||
model_output = self.model(input_ids, cache=state.cache)
|
||||
|
||||
logits = self._extract_logits(model_output)
|
||||
|
||||
last_logits = logits[:, -1, :]
|
||||
next_token_mlx = mx.argmax(last_logits, axis=-1)
|
||||
|
||||
mx.eval(next_token_mlx, *[c.state for c in state.cache])
|
||||
next_token = int(next_token_mlx.item())
|
||||
|
||||
state.token_ids.append(next_token)
|
||||
state.generated_tokens += 1
|
||||
|
||||
return next_token
|
||||
|
||||
def _batched_decode(
|
||||
self, decode_reqs: list[tuple[str, MlxRequestState]]
|
||||
) -> list[int]:
|
||||
"""Run a single batched forward pass for multiple decode requests."""
|
||||
last_tokens = [state.token_ids[-1] for _, state in decode_reqs]
|
||||
|
||||
# Merge individual KV caches into batched cache
|
||||
caches_list = [state.cache for _, state in decode_reqs]
|
||||
batch_cache = _merge_kv_caches(caches_list)
|
||||
|
||||
# Create batched input: shape (batch_size, 1)
|
||||
batched_input = mx.array(last_tokens, dtype=mx.int32)[:, None]
|
||||
|
||||
# Single forward pass
|
||||
model_output = self.model(batched_input, cache=batch_cache)
|
||||
logits = self._extract_logits(model_output)
|
||||
|
||||
next_token_logits = logits[:, -1, :]
|
||||
next_tokens_mlx = mx.argmax(next_token_logits, axis=-1)
|
||||
|
||||
mx.eval(next_tokens_mlx, *[c.state for c in batch_cache])
|
||||
next_tokens = next_tokens_mlx.tolist()
|
||||
|
||||
# Extract updated caches back to individual requests
|
||||
for i, (_, state) in enumerate(decode_reqs):
|
||||
state.cache = _extract_kv_cache(batch_cache, i)
|
||||
state.token_ids.append(next_tokens[i])
|
||||
state.generated_tokens += 1
|
||||
|
||||
return next_tokens
|
||||
|
||||
def remove_request(self, req_id: str):
|
||||
"""Clean up state for a completed request."""
|
||||
self._request_states.pop(req_id, None)
|
||||
|
||||
def clear(self):
|
||||
"""Clear all request states."""
|
||||
self._request_states.clear()
|
||||
@@ -0,0 +1,172 @@
|
||||
"""Lightweight ModelRunner stub for MLX on Apple Silicon.
|
||||
|
||||
Subclasses ModelRunner but overrides both load_model() and initialize()
|
||||
to skip PyTorch weight loading entirely. No GPU memory is consumed:
|
||||
the KV cache pool uses a zero-allocation _DummyKVCache, and only
|
||||
CPU-side bookkeeping structures (req_to_token_pool,
|
||||
token_to_kv_pool_allocator) are created so the SGLang scheduler can
|
||||
function. The actual KV cache is managed by the MLX model runner.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.memory_pool import KVCache, ReqToTokenPool
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class _DummyKVCache(KVCache):
|
||||
"""A KV cache that allocates no GPU memory.
|
||||
|
||||
Satisfies the KVCache interface so that TokenToKVPoolAllocator can be
|
||||
constructed, but every buffer access raises — the MLX backend manages
|
||||
its own KV cache internally.
|
||||
"""
|
||||
|
||||
def __init__(self, size: int, dtype: torch.dtype, device: str):
|
||||
# Bypass KVCache.__init__ to avoid custom_mem_pool / memory_saver
|
||||
# initialization that may touch CUDA APIs.
|
||||
self.size = size
|
||||
self.page_size = 1
|
||||
self.dtype = dtype
|
||||
self.store_dtype = dtype
|
||||
self.device = device
|
||||
self.layer_num = 0
|
||||
self.start_layer = 0
|
||||
self.end_layer = 0
|
||||
self.mem_usage = 0
|
||||
self.cpu_offloading_chunk_size = 8192
|
||||
self.layer_transfer_counter = None
|
||||
self.enable_custom_mem_pool = False
|
||||
self.custom_mem_pool = None
|
||||
|
||||
def get_key_buffer(self, layer_id: int) -> torch.Tensor:
|
||||
raise RuntimeError("_DummyKVCache has no key buffer (MLX manages KV cache)")
|
||||
|
||||
def get_value_buffer(self, layer_id: int) -> torch.Tensor:
|
||||
raise RuntimeError("_DummyKVCache has no value buffer (MLX manages KV cache)")
|
||||
|
||||
def get_kv_buffer(self, layer_id: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
raise RuntimeError("_DummyKVCache has no kv buffer (MLX manages KV cache)")
|
||||
|
||||
def set_kv_buffer(self, layer, loc, cache_k, cache_v) -> None:
|
||||
raise RuntimeError("_DummyKVCache cannot set kv buffer (MLX manages KV cache)")
|
||||
|
||||
def get_kv_size_bytes(self):
|
||||
return 0, 0
|
||||
|
||||
|
||||
class _DummyModel:
|
||||
"""Minimal stand-in so that `inspect.signature(model.forward)` and
|
||||
`getattr(model, ...)` calls in ModelRunner.__init__ don't crash."""
|
||||
|
||||
@staticmethod
|
||||
def forward():
|
||||
pass
|
||||
|
||||
|
||||
class MlxModelRunnerStub(ModelRunner):
|
||||
"""ModelRunner that skips PyTorch weight loading and KV cache allocation.
|
||||
|
||||
Overrides both load_model() and initialize() so that no PyTorch model
|
||||
weights are loaded and no large KV cache tensors are allocated. Only
|
||||
the minimal bookkeeping pools needed by the scheduler are created.
|
||||
"""
|
||||
|
||||
def load_model(self):
|
||||
"""Set only the metadata that downstream code needs, without
|
||||
loading any PyTorch model weights."""
|
||||
logger.info(
|
||||
"MLX stub: skipping PyTorch model weight loading "
|
||||
"(inference runs through MLX)"
|
||||
)
|
||||
|
||||
self.model = _DummyModel()
|
||||
|
||||
self.sliding_window_size = None
|
||||
if (
|
||||
self.model_config.is_hybrid_swa
|
||||
and self.model_config.sliding_window_size is not None
|
||||
):
|
||||
self.sliding_window_size = self.model_config.sliding_window_size
|
||||
elif self.model_config.attention_chunk_size is not None:
|
||||
self.sliding_window_size = self.model_config.attention_chunk_size
|
||||
|
||||
self.dtype = self.model_config.dtype
|
||||
self.weight_load_mem_usage = 0
|
||||
|
||||
def initialize(self, pre_model_load_memory: float):
|
||||
"""Lightweight initialize that skips heavy PyTorch setup.
|
||||
|
||||
Creates minimal req_to_token_pool and token_to_kv_pool_allocator
|
||||
with a dummy KV cache (zero GPU memory) so the scheduler works.
|
||||
"""
|
||||
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
||||
|
||||
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
|
||||
enable=self.server_args.enable_memory_saver
|
||||
)
|
||||
|
||||
# Load model (sets metadata only)
|
||||
self.sampler = None
|
||||
self.load_model()
|
||||
|
||||
# Layer metadata
|
||||
model_num_layers = max(
|
||||
self.model_config.num_hidden_layers,
|
||||
self.model_config.num_attention_layers,
|
||||
)
|
||||
self.start_layer = 0
|
||||
self.end_layer = model_num_layers
|
||||
self.num_effective_layers = model_num_layers
|
||||
|
||||
# KV cache dtype
|
||||
self.kv_cache_dtype = self.dtype
|
||||
|
||||
# Pool sizing — use context_len as the capacity.
|
||||
# No actual GPU memory is consumed because _DummyKVCache is empty.
|
||||
self.max_total_num_tokens = self.model_config.context_len
|
||||
self.max_running_requests = min(
|
||||
self.max_total_num_tokens // 2,
|
||||
4096,
|
||||
)
|
||||
self.is_hybrid_swa = False
|
||||
|
||||
# Create minimal pools
|
||||
self.req_to_token_pool = ReqToTokenPool(
|
||||
size=self.max_running_requests,
|
||||
max_context_len=self.model_config.context_len,
|
||||
device="cpu",
|
||||
enable_memory_saver=False,
|
||||
)
|
||||
|
||||
dummy_kv = _DummyKVCache(
|
||||
size=self.max_total_num_tokens,
|
||||
dtype=self.kv_cache_dtype,
|
||||
device="cpu",
|
||||
)
|
||||
self.token_to_kv_pool = dummy_kv
|
||||
self.token_to_kv_pool_allocator = TokenToKVPoolAllocator(
|
||||
size=self.max_total_num_tokens,
|
||||
dtype=self.kv_cache_dtype,
|
||||
device="cpu",
|
||||
kvcache=dummy_kv,
|
||||
need_sort=False,
|
||||
)
|
||||
|
||||
# No CUDA graphs, no attention backend
|
||||
self.graph_runner = None
|
||||
self.graph_mem_usage = 0
|
||||
self.attn_backend = None
|
||||
|
||||
logger.info(
|
||||
f"MLX stub: initialized minimal pools "
|
||||
f"(max_total_num_tokens={self.max_total_num_tokens}, "
|
||||
f"max_running_requests={self.max_running_requests}, "
|
||||
f"zero GPU KV cache allocation)"
|
||||
)
|
||||
@@ -0,0 +1,176 @@
|
||||
"""MLX-specific TpModelWorker subclass for Apple Silicon.
|
||||
|
||||
Overrides the standard TpModelWorker to route forward passes through
|
||||
the native MLX model runner, avoiding PyTorch MPS entirely for inference.
|
||||
|
||||
PyTorch model weights are never loaded. A lightweight ModelRunner stub
|
||||
(MlxModelRunnerStub) provides only the minimal bookkeeping structures
|
||||
(req_to_token_pool, token_to_kv_pool_allocator with a zero-memory
|
||||
dummy KV cache) that the SGLang scheduler expects. The actual KV cache
|
||||
is managed internally by the MLX model runner.
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.managers.schedule_batch import ModelWorkerBatch
|
||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||
from sglang.srt.managers.utils import GenerationBatchResult
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class MlxTpModelWorker(TpModelWorker):
|
||||
"""A tensor parallel model worker that routes inference through MLX.
|
||||
|
||||
Inherits from TpModelWorker for scheduler integration, but replaces
|
||||
the standard ModelRunner with MlxModelRunnerStub (no PyTorch weights,
|
||||
zero-memory KV cache) and delegates all forward passes to a native
|
||||
MlxModelRunner.
|
||||
"""
|
||||
|
||||
def _init_model_runner(self):
|
||||
"""Override to use a lightweight ModelRunner that skips weight loading."""
|
||||
from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner
|
||||
from sglang.srt.hardware_backend.mlx.model_runner_stub import (
|
||||
MlxModelRunnerStub,
|
||||
)
|
||||
|
||||
self._model_runner = MlxModelRunnerStub(
|
||||
model_config=self.model_config,
|
||||
mem_fraction_static=self.server_args.mem_fraction_static,
|
||||
gpu_id=self.gpu_id,
|
||||
tp_rank=self.tp_rank,
|
||||
tp_size=self.tp_size,
|
||||
moe_ep_rank=self.moe_ep_rank,
|
||||
moe_ep_size=self.ep_size,
|
||||
pp_rank=self.pp_rank,
|
||||
pp_size=self.pp_size,
|
||||
nccl_port=self.nccl_port,
|
||||
dp_rank=self.dp_rank,
|
||||
server_args=self.server_args,
|
||||
is_draft_worker=self.is_draft_worker,
|
||||
req_to_token_pool=self.req_to_token_pool,
|
||||
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
|
||||
memory_pool_config=self.memory_pool_config,
|
||||
)
|
||||
|
||||
# Initialize the MLX model runner (loads weights via MLX, not PyTorch)
|
||||
logger.info("Initializing MlxModelRunner for end-to-end MLX inference")
|
||||
self._mlx_runner = MlxModelRunner(
|
||||
model_path=self.server_args.model_path,
|
||||
trust_remote_code=self.server_args.trust_remote_code,
|
||||
)
|
||||
self._mlx_active_rids: set[str] = set()
|
||||
|
||||
def get_pad_input_ids_func(self):
|
||||
"""Override since the stub ModelRunner has no real model."""
|
||||
return None
|
||||
|
||||
def forward_batch_generation(
|
||||
self,
|
||||
model_worker_batch: ModelWorkerBatch,
|
||||
forward_batch: Optional[ForwardBatch] = None,
|
||||
pp_proxy_tensors: Optional[PPProxyTensors] = None,
|
||||
is_verify: bool = False,
|
||||
skip_attn_backend_init=False,
|
||||
) -> GenerationBatchResult:
|
||||
"""Override to route through MLX model runner."""
|
||||
if model_worker_batch is not None:
|
||||
return self._forward_batch_generation_mlx(model_worker_batch)
|
||||
|
||||
# Fallback to standard path for None batches
|
||||
return super().forward_batch_generation(
|
||||
model_worker_batch,
|
||||
forward_batch,
|
||||
pp_proxy_tensors,
|
||||
is_verify,
|
||||
skip_attn_backend_init,
|
||||
)
|
||||
|
||||
def _forward_batch_generation_mlx(
|
||||
self,
|
||||
model_worker_batch: ModelWorkerBatch,
|
||||
) -> GenerationBatchResult:
|
||||
"""Run forward pass through the MLX model runner.
|
||||
|
||||
Bypasses the standard ModelRunner forward+sample and uses native MLX
|
||||
inference for the entire model. Only supports greedy sampling.
|
||||
"""
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||
|
||||
forward_mode = model_worker_batch.forward_mode
|
||||
reqs = model_worker_batch.reqs
|
||||
|
||||
if forward_mode.is_idle():
|
||||
return GenerationBatchResult(
|
||||
logits_output=LogitsProcessorOutput(next_token_logits=None),
|
||||
can_run_cuda_graph=False,
|
||||
)
|
||||
|
||||
# Auto-cleanup: remove MLX state for requests no longer in the batch
|
||||
current_rids = {req.rid for req in reqs}
|
||||
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
|
||||
|
||||
next_token_ids_list = []
|
||||
|
||||
if forward_mode.is_extend():
|
||||
# Prefill (or MIXED): extract per-request tokens from concatenated input_ids
|
||||
input_ids_cpu = model_worker_batch.input_ids.cpu().tolist()
|
||||
extend_seq_lens = model_worker_batch.extend_seq_lens
|
||||
offset = 0
|
||||
prefill_rids = []
|
||||
decode_rids = []
|
||||
for i, req in enumerate(reqs):
|
||||
seq_len = extend_seq_lens[i]
|
||||
req_token_ids = input_ids_cpu[offset : offset + seq_len]
|
||||
offset += seq_len
|
||||
if req.rid in self._mlx_runner._request_states:
|
||||
# MIXED mode: this request already has MLX state, decode it
|
||||
decode_rids.append(req.rid)
|
||||
else:
|
||||
# Prefill: new request
|
||||
next_token = self._mlx_runner.prefill(req.rid, req_token_ids)
|
||||
prefill_rids.append((req.rid, next_token))
|
||||
|
||||
# Batch decode all existing requests at once
|
||||
if decode_rids:
|
||||
decode_results = self._mlx_runner.decode_batch(decode_rids)
|
||||
decode_map = dict(zip(decode_rids, decode_results))
|
||||
else:
|
||||
decode_map = {}
|
||||
|
||||
prefill_map = dict(prefill_rids)
|
||||
|
||||
# Reassemble in original request order
|
||||
for req in reqs:
|
||||
if req.rid in decode_map:
|
||||
next_token_ids_list.append(decode_map[req.rid])
|
||||
else:
|
||||
next_token_ids_list.append(prefill_map[req.rid])
|
||||
|
||||
elif forward_mode.is_decode():
|
||||
# Decode: batch decode all requests
|
||||
req_ids = [req.rid for req in reqs]
|
||||
next_token_ids_list = self._mlx_runner.decode_batch(req_ids)
|
||||
|
||||
else:
|
||||
raise ValueError(
|
||||
f"MLX runner does not support forward mode: {forward_mode}"
|
||||
)
|
||||
|
||||
next_token_ids = torch.tensor(
|
||||
next_token_ids_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,
|
||||
)
|
||||
@@ -224,6 +224,7 @@ from sglang.srt.utils.hf_transformers_utils import (
|
||||
get_tokenizer_from_processor,
|
||||
)
|
||||
from sglang.srt.utils.network import get_zmq_socket
|
||||
from sglang.srt.utils.tensor_bridge import use_mlx
|
||||
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
||||
from sglang.utils import TypeBasedDispatcher, get_exception_traceback
|
||||
|
||||
@@ -569,9 +570,8 @@ class Scheduler(
|
||||
self.require_mlp_sync = require_mlp_sync(self.server_args)
|
||||
|
||||
def init_tp_model_worker(self):
|
||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||
|
||||
self.tp_worker = TpModelWorker(
|
||||
worker_kwargs = dict(
|
||||
server_args=self.server_args,
|
||||
gpu_id=self.gpu_id,
|
||||
tp_rank=self.tp_rank,
|
||||
@@ -583,6 +583,16 @@ class Scheduler(
|
||||
nccl_port=self.nccl_port,
|
||||
)
|
||||
|
||||
# FIXME: move tp worker's init logic outside of the scheduler.
|
||||
if use_mlx():
|
||||
from sglang.srt.hardware_backend.mlx.tp_worker import MlxTpModelWorker
|
||||
|
||||
self.tp_worker = MlxTpModelWorker(**worker_kwargs)
|
||||
else:
|
||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||
|
||||
self.tp_worker = TpModelWorker(**worker_kwargs)
|
||||
|
||||
def maybe_init_draft_worker(self):
|
||||
if self.spec_algorithm.is_none():
|
||||
self.draft_worker = None
|
||||
|
||||
@@ -764,6 +764,7 @@ class ServerArgs:
|
||||
self._handle_hpu_backends()
|
||||
self._handle_cpu_backends()
|
||||
self._handle_npu_backends()
|
||||
self._handle_mps_backends()
|
||||
self._handle_xpu_backends()
|
||||
|
||||
# Handle piecewise CUDA graph.
|
||||
@@ -1043,6 +1044,10 @@ class ServerArgs:
|
||||
)
|
||||
self.piecewise_cuda_graph_compiler = "eager"
|
||||
|
||||
def _handle_mps_backends(self):
|
||||
if self.device == "mps":
|
||||
self.disable_overlap_schedule = True
|
||||
|
||||
def _handle_xpu_backends(self):
|
||||
if self.device == "xpu":
|
||||
if not self.disable_piecewise_cuda_graph:
|
||||
|
||||
@@ -0,0 +1,228 @@
|
||||
# Copied and adapted from: https://github.com/vllm-project/vllm-metal
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tensor bridge between MLX and PyTorch.
|
||||
|
||||
Provides zero-copy conversion when possible using Apple Silicon's unified memory.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from functools import lru_cache
|
||||
from typing import TYPE_CHECKING, Literal
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import mlx.core as mx
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_MLX_AVAILABLE: bool = False
|
||||
try:
|
||||
import mlx.core as mx # noqa: F811
|
||||
|
||||
_MLX_AVAILABLE = True
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
|
||||
def is_mlx_available() -> bool:
|
||||
"""Return True when the ``mlx`` package can be imported."""
|
||||
return _MLX_AVAILABLE
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def use_mlx() -> bool:
|
||||
"""Return True when the user opted-in via ``SGLANG_USE_MLX=1`` **and** MLX is importable."""
|
||||
return bool(envs.SGLANG_USE_MLX.get()) and _MLX_AVAILABLE
|
||||
|
||||
|
||||
# MPS has a 4GB (2^32 bytes) limit for MPSTemporaryNDArray allocations.
|
||||
# Metal may allocate multiple temporary buffers internally, so we use a
|
||||
# conservative threshold of 1GB to avoid hitting the limit.
|
||||
# See: https://github.com/anthropics/vllm-metal/issues/43
|
||||
_MPS_SAFE_SIZE_BYTES = 1 << 30 # 1GB
|
||||
|
||||
# MLX to PyTorch dtype mapping
|
||||
# TODO(perf): float64 is CPU-only in MLX (see ml-explore/mlx#1843).
|
||||
# When the target device is GPU/MPS we should auto-downcast float64 → float32
|
||||
# to avoid a runtime error; when the target is CPU we can keep float64.
|
||||
# For now float64 is omitted from the mapping so it hits the ValueError
|
||||
# fallback in mlx_to_torch().
|
||||
MLX_TO_TORCH_DTYPE = (
|
||||
{
|
||||
mx.float32: torch.float32,
|
||||
mx.float16: torch.float16,
|
||||
mx.bfloat16: torch.bfloat16,
|
||||
mx.int32: torch.int32,
|
||||
mx.int64: torch.int64,
|
||||
mx.int16: torch.int16,
|
||||
mx.int8: torch.int8,
|
||||
mx.uint8: torch.uint8,
|
||||
mx.bool_: torch.bool,
|
||||
}
|
||||
if _MLX_AVAILABLE
|
||||
else {}
|
||||
)
|
||||
|
||||
# PyTorch to MLX dtype mapping
|
||||
TORCH_TO_MLX_DTYPE = {v: k for k, v in MLX_TO_TORCH_DTYPE.items()}
|
||||
|
||||
|
||||
def get_torch_device() -> torch.device:
|
||||
"""Get the PyTorch device for Metal/MPS.
|
||||
|
||||
Returns:
|
||||
torch.device for MPS if available, else CPU
|
||||
"""
|
||||
if torch.backends.mps.is_available():
|
||||
return torch.device("mps")
|
||||
return torch.device("cpu")
|
||||
|
||||
|
||||
def _get_tensor_size_bytes(array: mx.array) -> int:
|
||||
"""Calculate the size of an MLX array in bytes.
|
||||
|
||||
Args:
|
||||
array: MLX array
|
||||
|
||||
Returns:
|
||||
Size in bytes
|
||||
"""
|
||||
return array.size * array.dtype.size
|
||||
|
||||
|
||||
def _is_safe_for_mps(array: mx.array) -> bool:
|
||||
"""Check if an array is safe to transfer to MPS without hitting size limits.
|
||||
|
||||
MPS has a 4GB limit for MPSTemporaryNDArray, but Metal may allocate
|
||||
multiple temporary buffers internally. We use a conservative threshold.
|
||||
|
||||
Args:
|
||||
array: MLX array to check
|
||||
|
||||
Returns:
|
||||
True if safe to transfer to MPS, False if should stay on CPU
|
||||
"""
|
||||
return _get_tensor_size_bytes(array) < _MPS_SAFE_SIZE_BYTES
|
||||
|
||||
|
||||
def torch_to_mlx(tensor: torch.Tensor) -> mx.array:
|
||||
"""Convert PyTorch tensor to MLX array.
|
||||
|
||||
Uses numpy as an intermediate to enable zero-copy on unified memory.
|
||||
|
||||
Args:
|
||||
tensor: PyTorch tensor (can be on any device)
|
||||
|
||||
Returns:
|
||||
MLX array with the same data
|
||||
"""
|
||||
# Move to CPU if on MPS for numpy conversion
|
||||
if tensor.device.type != "cpu":
|
||||
tensor = tensor.cpu()
|
||||
|
||||
tensor = tensor.detach()
|
||||
|
||||
# Note: numpy does not support bfloat16.
|
||||
if tensor.dtype == torch.bfloat16:
|
||||
return mx.array(tensor)
|
||||
|
||||
return mx.array(tensor.numpy())
|
||||
|
||||
|
||||
# TODO(perf): accept a list/batch of arrays and convert them in one pass
|
||||
# to reduce the Python ↔ MLX round-trip overhead.
|
||||
def mlx_to_torch(
|
||||
array: mx.array,
|
||||
device: torch.device | Literal["mps", "cpu"] | None = None,
|
||||
already_contiguous: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Convert MLX array to PyTorch tensor.
|
||||
|
||||
Uses numpy as an intermediate to enable zero-copy on unified memory.
|
||||
|
||||
Args:
|
||||
array: MLX array
|
||||
device: Target PyTorch device (default: MPS if available)
|
||||
already_contiguous: Skip contiguity check if array is known contiguous
|
||||
|
||||
Returns:
|
||||
PyTorch tensor with the same data
|
||||
"""
|
||||
if device is None:
|
||||
device = get_torch_device()
|
||||
elif isinstance(device, str):
|
||||
device = torch.device(device)
|
||||
|
||||
# Use memoryview for zero-copy conversion (bypasses numpy for bfloat16)
|
||||
# reference: https://github.com/ml-explore/mlx/issues/403
|
||||
torch_dtype = MLX_TO_TORCH_DTYPE.get(array.dtype)
|
||||
if torch_dtype is not None:
|
||||
if already_contiguous:
|
||||
# Fast path: skip contiguity check, single eval
|
||||
mx.eval(array)
|
||||
buffer = memoryview(array)
|
||||
else:
|
||||
# MLX views / non-contiguous arrays expose a non-contiguous buffer (or
|
||||
# sometimes no usable buffer), which `torch.frombuffer` can't consume.
|
||||
# Make contiguous first, then eval once
|
||||
array = mx.contiguous(array)
|
||||
mx.eval(array)
|
||||
buffer = memoryview(array)
|
||||
|
||||
tensor = torch.frombuffer(buffer, dtype=torch_dtype).reshape(array.shape)
|
||||
else:
|
||||
# Fallback to numpy path for unsupported dtypes
|
||||
raise ValueError(f"Unsupported MLX dtype: {array.dtype}")
|
||||
|
||||
# Move to target device, but check for MPS size limits first
|
||||
if device.type == "mps":
|
||||
if _is_safe_for_mps(array):
|
||||
tensor = tensor.to(device)
|
||||
else:
|
||||
# Large tensor - keep on CPU to avoid MPS 4GB limit crash
|
||||
# See: https://github.com/anthropics/vllm-metal/issues/43
|
||||
logger.debug(
|
||||
"Tensor too large for MPS (%d bytes > %d limit), keeping on CPU",
|
||||
_get_tensor_size_bytes(array),
|
||||
_MPS_SAFE_SIZE_BYTES,
|
||||
)
|
||||
elif device.type != "cpu":
|
||||
tensor = tensor.to(device)
|
||||
|
||||
return tensor
|
||||
|
||||
|
||||
def sync_mlx() -> None:
|
||||
"""Synchronize MLX operations.
|
||||
|
||||
Call this before converting MLX arrays to ensure all operations complete.
|
||||
"""
|
||||
# Prefer an explicit MLX barrier when available; otherwise force evaluation.
|
||||
# `mx.eval([])` is a no-op, so we evaluate a tiny scalar as a safe fallback.
|
||||
try:
|
||||
mx.synchronize()
|
||||
except (AttributeError, TypeError):
|
||||
mx.eval(mx.array(0, dtype=mx.int32))
|
||||
|
||||
|
||||
def sync_torch() -> None:
|
||||
"""Synchronize PyTorch MPS operations.
|
||||
|
||||
Call this before converting PyTorch tensors to ensure all operations complete.
|
||||
"""
|
||||
if torch.backends.mps.is_available():
|
||||
torch.mps.synchronize()
|
||||
|
||||
|
||||
__all__ = [
|
||||
"is_mlx_available",
|
||||
"use_mlx",
|
||||
"mlx_to_torch",
|
||||
"torch_to_mlx",
|
||||
"get_torch_device",
|
||||
]
|
||||
Reference in New Issue
Block a user