Address ring slowdown by turning on FAST SYNCH (#1594)

## Motivation

Large models + large prompts + pipeline RING = 0.2tps generation speed
Large models + large prompts + pipeline JACCL = 15 tps generation speed

Why? Well, MLX_METAL_FAST_SYNCH is on in pipeline JACCL.

## Changes

Just turn on fast synch everywhere, especially as GPU locks are old news

Also, changed to use mx.device_info as mx.metal.device_info is going to
be deprecated.

## Why It Works

Some magic thing that happens in the mlx backend. I really tried to find
a regression but couldn't. I will probably try again at some point.

## Test Plan

### Manual Testing
Did a bunch, no longer 0.2tps

### Automated Testing
We'll do that today.
This commit is contained in:
rltakashige
2026-02-23 15:14:58 +00:00
committed by GitHub
parent 8d94eab6c6
commit fbb80e1cc9
3 changed files with 19 additions and 9 deletions
+16
View File
@@ -215,6 +215,22 @@ class StreamContext:
traceback: object | None = ...,
) -> None: ...
def device_info() -> dict[str, str | int]:
"""
Get information about the GPU device and system settings.
Currently returns:
* ``architecture``
* ``max_buffer_size``
* ``max_recommended_working_set_size``
* ``memory_size``
* ``resource_limit``
Returns:
dict: A dictionary with string keys and string or integer values.
"""
def abs(a: array, /, *, stream: Stream | Device | None = ...) -> array:
"""
Element-wise absolute value.
+1 -1
View File
@@ -669,7 +669,7 @@ def set_wired_limit_for_model(model_size: Memory):
return
max_rec_size = Memory.from_bytes(
int(mx.metal.device_info()["max_recommended_working_set_size"])
int(mx.device_info()["max_recommended_working_set_size"])
)
if model_size > 0.9 * max_rec_size:
logger.warning(
+2 -8
View File
@@ -4,7 +4,7 @@ import loguru
from exo.shared.types.events import Event, RunnerStatusUpdated
from exo.shared.types.tasks import Task, TaskId
from exo.shared.types.worker.instances import BoundInstance, MlxJacclInstance
from exo.shared.types.worker.instances import BoundInstance
from exo.shared.types.worker.runners import RunnerFailed
from exo.utils.channels import ClosedResourceError, MpReceiver, MpSender
@@ -19,13 +19,7 @@ def entrypoint(
_logger: "loguru.Logger",
) -> None:
fast_synch_override = os.environ.get("EXO_FAST_SYNCH")
if fast_synch_override == "on" or (
fast_synch_override != "off"
and (
isinstance(bound_instance.instance, MlxJacclInstance)
and len(bound_instance.instance.jaccl_devices) >= 2
)
):
if fast_synch_override != "off":
os.environ["MLX_METAL_FAST_SYNCH"] = "1"
else:
os.environ["MLX_METAL_FAST_SYNCH"] = "0"