From c65320acd3df04d14030a2b746802481ae1a76cd Mon Sep 17 00:00:00 2001 From: Chris A <172564514+chriscod3@users.noreply.github.com> Date: Thu, 8 Jan 2026 19:40:15 -0600 Subject: [PATCH] Fix mlx seed (#1094) ## Motivation ## Changes ## Why It Works ## Test Plan ### Manual Testing ### Automated Testing --------- Co-authored-by: google-labs-jules[bot] <161369871+google-labs-jules[bot]@users.noreply.github.com> Co-authored-by: Ryuichi Leo Takashige --- 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 ++------- .../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)