Skip to content

Commit 7a27935

Browse files
committed
Set wired_limit to fix low GPU utilization
Without wired_limit, model weights cannot stay pinned in GPU-accessible memory, causing frequent memory paging and GPU stalls. This results in ~30% GPU utilization instead of ~100%. The fix sets max_recommended_working_set_size as the wired limit before model loading, matching the fix in ml-explore/mlx-lm#652. Changes: - Set wired_limit in MetalWorker.init_device() for v1 engine - Set wired_limit in MetalModelRunner.load_model() for legacy path - Use new mx.set_wired_limit() API with fallback to deprecated API - Add error handling for graceful degradation Signed-off-by: otarkhan <osama.taha1994@gmail.com>
1 parent ae6a1cb commit 7a27935

3 files changed

Lines changed: 27 additions & 1 deletion

File tree

vllm_metal/model_runner.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111

1212
from vllm_metal.config import get_config
1313
from vllm_metal.mlx_backend.cache import PagedKVCache
14+
from vllm_metal.platform import set_wired_limit
1415

1516
if TYPE_CHECKING:
1617
from vllm.config import VllmConfig
@@ -44,6 +45,7 @@ def load_model(self) -> None:
4445
model_name = model_config.model
4546

4647
logger.info(f"Loading model: {model_name}")
48+
set_wired_limit()
4749

4850
# Load model and tokenizer using mlx_lm
4951
self.model, self.tokenizer = mlx_load(

vllm_metal/platform.py

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,29 @@
1919
logger = logging.getLogger(__name__)
2020

2121

22+
def set_wired_limit() -> None:
23+
"""Set Metal wired memory limit for optimal GPU performance.
24+
25+
Pins model weights in GPU-accessible memory to prevent memory paging
26+
and GPU stalls during inference.
27+
28+
See: https://github.com/ml-explore/mlx-lm/pull/652
29+
"""
30+
try:
31+
import mlx.core as mx
32+
33+
device_info = mx.metal.device_info()
34+
max_wired = device_info.get("max_recommended_working_set_size", 0)
35+
if max_wired > 0:
36+
if hasattr(mx, "set_wired_limit"):
37+
mx.set_wired_limit(max_wired)
38+
elif hasattr(mx.metal, "set_wired_limit"):
39+
mx.metal.set_wired_limit(max_wired)
40+
logger.info(f"Set Metal wired_limit to {max_wired / (1024**3):.1f} GB")
41+
except Exception as e:
42+
logger.warning(f"Failed to set wired_limit: {e}")
43+
44+
2245
class MetalPlatform(Platform):
2346
"""Platform implementation for Apple Silicon Metal/MLX.
2447

vllm_metal/v1/worker.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@
2222
from vllm.v1.worker.worker_base import WorkerBase
2323

2424
from vllm_metal.config import get_config
25-
from vllm_metal.platform import MetalPlatform
25+
from vllm_metal.platform import MetalPlatform, set_wired_limit
2626

2727
if TYPE_CHECKING:
2828
from vllm_metal.v1.model_runner import MetalModelRunner
@@ -95,6 +95,7 @@ def init_device(self) -> None:
9595
)
9696
mx.set_default_device(mx.Device(device_type))
9797
logger.info(f"MLX device set to: {mx.default_device()}")
98+
set_wired_limit()
9899

99100
# Use MetalPlatform.get_torch_device() to properly support MPS when available.
100101
# This ensures consistency with the platform's device selection logic and

0 commit comments

Comments
 (0)