From da84a1eec218a168017a40c6b55c68992d4fcbe1 Mon Sep 17 00:00:00 2001 From: otarkhan Date: Thu, 4 Dec 2025 01:29:33 +0200 Subject: [PATCH 1/2] Fix slow batch generation in server by setting wired_limit The server's batch generation path was missing the wired_limit setting that stream_generate uses, causing ~10x slower performance and ~30% GPU utilization on large models. When wired_limit is not set, model weights cannot stay pinned in GPU-accessible memory, causing frequent memory paging and GPU stalls. This fix sets max_recommended_working_set_size when creating the BatchGenerator and restores the previous limit when done, matching the behavior of stream_generate. --- mlx_lm/server.py | 11 ++++++++++- 1 file changed, 10 insertions(+), 1 deletion(-) diff --git a/mlx_lm/server.py b/mlx_lm/server.py index 520cc815a..1930bc5e5 100644 --- a/mlx_lm/server.py +++ b/mlx_lm/server.py @@ -33,7 +33,7 @@ from huggingface_hub import scan_cache_dir from ._version import __version__ -from .generate import BatchGenerator, stream_generate +from .generate import BatchGenerator, generation_stream, stream_generate, wired_limit from .models.cache import can_trim_prompt_cache, make_prompt_cache, trim_prompt_cache from .sample_utils import make_logits_processors, make_sampler from .utils import load @@ -509,6 +509,7 @@ def _generate(self): batch_generator = None drain_batch = False batch_results = {} + old_wired_limit = None unprocessed_requests = [] @@ -596,6 +597,10 @@ def progress_callback(info): current_tokenizer = tokenizer current_model_key = self.model_provider.model_key batch_results = {} + if mx.metal.is_available(): + old_wired_limit = mx.set_wired_limit( + mx.metal.device_info()["max_recommended_working_set_size"] + ) batch_generator = BatchGenerator( model, stop_tokens=tokenizer.eos_token_ids, @@ -633,6 +638,10 @@ def progress_callback(info): current_model_key = None batch_generator = None drain_batch = False + if old_wired_limit is not None: + mx.synchronize(generation_stream) + mx.set_wired_limit(old_wired_limit) + old_wired_limit = None continue uids_to_remove = [] From ec8ba235269744f74ac502cf19253c3b605c661d Mon Sep 17 00:00:00 2001 From: otarkhan Date: Thu, 4 Dec 2025 19:13:32 +0200 Subject: [PATCH 2/2] Move wired_limit management to BatchGenerator class - Set wired_limit in __init__() and restore in close() - Add __del__() as fallback safety net - Call close() explicitly in server and batch_generate - Remove redundant wired_limit context manager from batch_generate --- mlx_lm/generate.py | 50 ++++++++++++++++++++++++++++++---------------- mlx_lm/server.py | 12 ++--------- 2 files changed, 35 insertions(+), 27 deletions(-) diff --git a/mlx_lm/generate.py b/mlx_lm/generate.py index 702695798..d0e49cb0f 100644 --- a/mlx_lm/generate.py +++ b/mlx_lm/generate.py @@ -948,6 +948,22 @@ def __init__( self.active_batch = None + if mx.metal.is_available(): + self._old_wired_limit = mx.set_wired_limit( + mx.metal.device_info()["max_recommended_working_set_size"] + ) + else: + self._old_wired_limit = None + + def close(self): + if self._old_wired_limit is not None: + mx.synchronize(generation_stream) + mx.set_wired_limit(self._old_wired_limit) + self._old_wired_limit = None + + def __del__(self): + self.close() + def insert( self, prompts, max_tokens: Union[List[int], int, None] = None, caches=None ): @@ -1196,23 +1212,23 @@ def batch_generate( if verbose: print(f"[batch_generate] Finished processing 0/{num_samples} ...", end="\r") - with wired_limit(model, [generation_stream]): - uids = gen.insert(prompts, max_tokens, caches=prompt_caches) - results = {uid: [] for uid in uids} - prompt_caches = {} - while responses := gen.next(): - for r in responses: - if r.finish_reason is not None: - if return_prompt_caches: - prompt_caches[r.uid] = r.prompt_cache - if verbose: - fin += 1 - print( - f"[batch_generate] Finished processing {fin}/{num_samples} ...", - end="\r", - ) - if r.finish_reason != "stop": - results[r.uid].append(r.token) + uids = gen.insert(prompts, max_tokens, caches=prompt_caches) + results = {uid: [] for uid in uids} + prompt_caches = {} + while responses := gen.next(): + for r in responses: + if r.finish_reason is not None: + if return_prompt_caches: + prompt_caches[r.uid] = r.prompt_cache + if verbose: + fin += 1 + print( + f"[batch_generate] Finished processing {fin}/{num_samples} ...", + end="\r", + ) + if r.finish_reason != "stop": + results[r.uid].append(r.token) + gen.close() if verbose: print(f"[batch_generate] Finished processing {fin}/{num_samples}") diff --git a/mlx_lm/server.py b/mlx_lm/server.py index 1930bc5e5..fa4c307db 100644 --- a/mlx_lm/server.py +++ b/mlx_lm/server.py @@ -33,7 +33,7 @@ from huggingface_hub import scan_cache_dir from ._version import __version__ -from .generate import BatchGenerator, generation_stream, stream_generate, wired_limit +from .generate import BatchGenerator, stream_generate from .models.cache import can_trim_prompt_cache, make_prompt_cache, trim_prompt_cache from .sample_utils import make_logits_processors, make_sampler from .utils import load @@ -509,7 +509,6 @@ def _generate(self): batch_generator = None drain_batch = False batch_results = {} - old_wired_limit = None unprocessed_requests = [] @@ -597,10 +596,6 @@ def progress_callback(info): current_tokenizer = tokenizer current_model_key = self.model_provider.model_key batch_results = {} - if mx.metal.is_available(): - old_wired_limit = mx.set_wired_limit( - mx.metal.device_info()["max_recommended_working_set_size"] - ) batch_generator = BatchGenerator( model, stop_tokens=tokenizer.eos_token_ids, @@ -636,12 +631,9 @@ def progress_callback(info): current_sampling = None current_tokenizer = None current_model_key = None + batch_generator.close() batch_generator = None drain_batch = False - if old_wired_limit is not None: - mx.synchronize(generation_stream) - mx.set_wired_limit(old_wired_limit) - old_wired_limit = None continue uids_to_remove = []