From f0433505a8f5f94bc47e7c4541c6084f6d100b7e Mon Sep 17 00:00:00 2001 From: dmcc73 Date: Mon, 30 Mar 2026 18:21:57 +0100 Subject: [PATCH] Fix speculative temp: use request temperature, not global env var The speculative cycle was using EXO_SPECULATIVE_TEMP (global) instead of the request's actual temperature. This caused greedy decoding in speculative while the model sampled at T=0.7, producing different (shorter) output and missing responses after . Now passes task_params.temperature from submit() to MTPBatchGenerator per-request via _request_temp[uid]. Falls back to self.temp (env var) if not set. Co-Authored-By: Claude Opus 4.6 (1M context) --- src/exo/worker/engines/mlx/generator/batch_generate.py | 5 +++++ .../worker/engines/mlx/speculative/mtp_batch_generator.py | 4 +++- 2 files changed, 8 insertions(+), 1 deletion(-) diff --git a/src/exo/worker/engines/mlx/generator/batch_generate.py b/src/exo/worker/engines/mlx/generator/batch_generate.py index 7ab39fb7..28c69c71 100644 --- a/src/exo/worker/engines/mlx/generator/batch_generate.py +++ b/src/exo/worker/engines/mlx/generator/batch_generate.py @@ -332,6 +332,11 @@ class ExoBatchGenerator: uid = uids[0] + # Pass request temperature to speculative cycle + if hasattr(self._exo_gen, '_request_temp'): + request_temp = task_params.temperature if task_params.temperature is not None else 0.7 + self._exo_gen._request_temp[uid] = request_temp + self._active_tasks[uid] = _EngineTask( uid=uid, task_params=task_params, diff --git a/src/exo/worker/engines/mlx/speculative/mtp_batch_generator.py b/src/exo/worker/engines/mlx/speculative/mtp_batch_generator.py index 0a82581e..2e5dc8d4 100644 --- a/src/exo/worker/engines/mlx/speculative/mtp_batch_generator.py +++ b/src/exo/worker/engines/mlx/speculative/mtp_batch_generator.py @@ -45,6 +45,7 @@ class MTPBatchGenerator(BatchGenerator): self._captured = {} # pre_norm / prompt_pre_norm from norm wrapper self._mtp_pre_norm = {} # uid → (B, 1, D) pre-norm hidden state self._mtp_prefilled = set() # uids with MTP cache prefilled + self._request_temp = {} # uid → temperature from request self._setup_hidden_capture() @@ -134,7 +135,7 @@ class MTPBatchGenerator(BatchGenerator): return super()._next() gamma = self.gamma - temp = self.temp + temp = self._request_temp.get(uid, self.temp) alpha = self.alpha # 1. Draft γ tokens (lazy chain, no eval) @@ -303,3 +304,4 @@ class MTPBatchGenerator(BatchGenerator): self._mtp_pre_norm.pop(uid, None) self._mtp_prefilled.discard(uid) self._token_buffer.pop(uid, None) + self._request_temp.pop(uid, None)