← back to Exo
Fix mlx seed (#1094)
c65320acd3df04d14030a2b746802481ae1a76cd · 2026-01-08 19:40:15 -0600 · Chris A
## Motivation
<!-- Why is this change needed? What problem does it solve? -->
<!-- If it fixes an open issue, please link to the issue here -->
## Changes
<!-- Describe what you changed in detail -->
## Why It Works
<!-- Explain why your approach solves the problem -->
## Test Plan
### Manual Testing
<!-- Hardware: (e.g., MacBook Pro M1 Max 32GB, Mac Mini M2 16GB,
connected via Thunderbolt 4) -->
<!-- What you did: -->
<!-- - -->
### Automated Testing
<!-- Describe changes to automated tests, or how existing tests cover
this change -->
<!-- - -->
---------
Co-authored-by: google-labs-jules[bot] <161369871+google-labs-jules[bot]@users.noreply.github.com>
Co-authored-by: Ryuichi Leo Takashige <leo@exolabs.net>
Files touched
M src/exo/worker/engines/mlx/generator/generate.pyM src/exo/worker/engines/mlx/utils_mlx.pyM src/exo/worker/runner/runner.pyM src/exo/worker/tests/unittests/test_runner/test_event_ordering.py
Diff
commit c65320acd3df04d14030a2b746802481ae1a76cd
Author: Chris A <172564514+chriscod3@users.noreply.github.com>
Date: Thu Jan 8 19:40:15 2026 -0600
Fix mlx seed (#1094)
## Motivation
<!-- Why is this change needed? What problem does it solve? -->
<!-- If it fixes an open issue, please link to the issue here -->
## Changes
<!-- Describe what you changed in detail -->
## Why It Works
<!-- Explain why your approach solves the problem -->
## Test Plan
### Manual Testing
<!-- Hardware: (e.g., MacBook Pro M1 Max 32GB, Mac Mini M2 16GB,
connected via Thunderbolt 4) -->
<!-- What you did: -->
<!-- - -->
### Automated Testing
<!-- Describe changes to automated tests, or how existing tests cover
this change -->
<!-- - -->
---------
Co-authored-by: google-labs-jules[bot] <161369871+google-labs-jules[bot]@users.noreply.github.com>
Co-authored-by: Ryuichi Leo Takashige <leo@exolabs.net>
---
src/exo/worker/engines/mlx/generator/generate.py | 14 ++++++++++++--
src/exo/worker/engines/mlx/utils_mlx.py | 11 +++--------
src/exo/worker/runner/runner.py | 9 ++-------
.../tests/unittests/test_runner/test_event_ordering.py | 2 +-
4 files changed, 18 insertions(+), 18 deletions(-)
diff --git a/src/exo/worker/engines/mlx/generator/generate.py b/src/exo/worker/engines/mlx/generator/generate.py
index cbcc288d..c42bf5d4 100644
--- a/src/exo/worker/engines/mlx/generator/generate.py
+++ b/src/exo/worker/engines/mlx/generator/generate.py
@@ -3,6 +3,7 @@ from typing import Any, Callable, Generator, cast, get_args
import mlx.core as mx
from mlx_lm import stream_generate
from mlx_lm.models.cache import KVCache
+from mlx_lm.sample_utils import make_sampler
from mlx_lm.tokenizer_utils import TokenizerWrapper
# from exo.engines.mlx.cache import KVPrefixCache
@@ -47,7 +48,6 @@ def maybe_quantize_kv_cache(
def warmup_inference(
model: Model,
tokenizer: TokenizerWrapper,
- sampler: Callable[[mx.array], mx.array],
) -> int:
content = "Prompt to warm up the inference engine. Repeat this."
@@ -70,6 +70,9 @@ def warmup_inference(
model=model,
)
+ # Use a default sampler for warmup
+ sampler = make_sampler(temp=0.7)
+
logger.info("Generating warmup tokens")
for _r in stream_generate(
model=model,
@@ -115,7 +118,6 @@ def eos_ids_from_tokenizer(tokenizer: TokenizerWrapper) -> list[int]:
def mlx_generate(
model: Model,
tokenizer: TokenizerWrapper,
- sampler: Callable[[mx.array], mx.array],
task: ChatCompletionTaskParams,
) -> Generator[GenerationResponse]:
# Ensure that generation stats only contains peak memory for this generation
@@ -125,6 +127,9 @@ def mlx_generate(
# Currently we support chat-completion tasks only.
logger.info(f"task_params: {task}")
+ if task.seed is not None:
+ mx.random.seed(task.seed)
+
prompt = apply_chat_template(
tokenizer=tokenizer,
chat_task_data=task,
@@ -138,6 +143,11 @@ def mlx_generate(
eos_ids = eos_ids_from_tokenizer(tokenizer)
logits_processors = [ban_token_ids(eos_ids)]
+ sampler = make_sampler(
+ temp=task.temperature if task.temperature is not None else 0.7,
+ top_p=task.top_p if task.top_p is not None else 1.0,
+ )
+
max_tokens = task.max_tokens or MAX_TOKENS
for out in stream_generate(
model=model,
diff --git a/src/exo/worker/engines/mlx/utils_mlx.py b/src/exo/worker/engines/mlx/utils_mlx.py
index 96c8893f..5b2dfbaa 100644
--- a/src/exo/worker/engines/mlx/utils_mlx.py
+++ b/src/exo/worker/engines/mlx/utils_mlx.py
@@ -3,11 +3,10 @@ import os
import resource
import time
from pathlib import Path
-from typing import Any, Callable, cast
+from typing import Any, cast
from mlx_lm.models.cache import KVCache, QuantizedKVCache, RotatingKVCache
from mlx_lm.models.deepseek_v3 import DeepseekV3Model
-from mlx_lm.sample_utils import make_sampler
from mlx_lm.tokenizer_utils import TokenizerWrapper
from exo.worker.engines.mlx.constants import (
@@ -176,11 +175,7 @@ def initialize_mlx(
def load_mlx_items(
bound_instance: BoundInstance, group: Group | None
-) -> tuple[Model, TokenizerWrapper, Callable[[mx.array], mx.array]]:
- # TODO: pass temperature
- sampler: Callable[[mx.array], mx.array] = make_sampler(temp=0.7)
- logger.info("Created a sampler")
-
+) -> tuple[Model, TokenizerWrapper]:
if group is None:
logger.info(f"Single device used for {bound_instance.instance}")
model_path = build_model_path(bound_instance.bound_shard.model_meta.model_id)
@@ -201,7 +196,7 @@ def load_mlx_items(
set_wired_limit_for_model(get_weights_size(bound_instance.bound_shard))
- return cast(Model, model), tokenizer, sampler
+ return cast(Model, model), tokenizer
def shard_and_load(
diff --git a/src/exo/worker/runner/runner.py b/src/exo/worker/runner/runner.py
index f39e7c10..029bbf4c 100644
--- a/src/exo/worker/runner/runner.py
+++ b/src/exo/worker/runner/runner.py
@@ -68,7 +68,6 @@ def main(
model = None
tokenizer = None
- sampler = None
group = None
current_status: RunnerStatus = RunnerIdle()
@@ -110,14 +109,13 @@ def main(
)
)
- model, tokenizer, sampler = load_mlx_items(bound_instance, group)
+ model, tokenizer = load_mlx_items(bound_instance, group)
current_status = RunnerLoaded()
logger.info("runner loaded")
case StartWarmup() if isinstance(current_status, RunnerLoaded):
assert model
assert tokenizer
- assert sampler
current_status = RunnerWarmingUp()
logger.info("runner warming up")
event_sender.send(
@@ -130,7 +128,6 @@ def main(
toks = warmup_inference(
model=model,
tokenizer=tokenizer,
- sampler=sampler,
# kv_prefix_cache=kv_prefix_cache, # supply for warmup-time prefix caching
)
logger.info(f"warmed up by generating {toks} tokens")
@@ -144,7 +141,6 @@ def main(
):
assert model
assert tokenizer
- assert sampler
logger.info(f"received chat request: {str(task)[:500]}")
current_status = RunnerRunning()
logger.info("runner running")
@@ -160,7 +156,6 @@ def main(
for response in mlx_generate(
model=model,
tokenizer=tokenizer,
- sampler=sampler,
task=task_params,
):
match response:
@@ -204,7 +199,7 @@ def main(
RunnerStatusUpdated(runner_id=runner_id, runner_status=current_status)
)
if isinstance(current_status, RunnerShutdown):
- del model, tokenizer, group, sampler
+ del model, tokenizer, group
mx.clear_cache()
import gc
diff --git a/src/exo/worker/tests/unittests/test_runner/test_event_ordering.py b/src/exo/worker/tests/unittests/test_runner/test_event_ordering.py
index 954052c3..e1744d1b 100644
--- a/src/exo/worker/tests/unittests/test_runner/test_event_ordering.py
+++ b/src/exo/worker/tests/unittests/test_runner/test_event_ordering.py
@@ -111,7 +111,7 @@ def assert_events_equal(test_events: Iterable[Event], true_events: Iterable[Even
def patch_out_mlx(monkeypatch: pytest.MonkeyPatch):
# initialize_mlx returns a "group" equal to 1
monkeypatch.setattr(mlx_runner, "initialize_mlx", make_nothin(1))
- monkeypatch.setattr(mlx_runner, "load_mlx_items", make_nothin((1, 1, 1)))
+ monkeypatch.setattr(mlx_runner, "load_mlx_items", make_nothin((1, 1)))
monkeypatch.setattr(mlx_runner, "warmup_inference", make_nothin(1))
monkeypatch.setattr(mlx_runner, "_check_for_debug_prompts", nothin)
← b9a78f6f ci: compute CURRENT_PROJECT_VERSION from semver
·
back to Exo
·
UNKNOWN to PREPARING (#1112) 59e7594e →