[object Object]

← 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

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 →