3535 KVCache ,
3636 QuantizedKVCache ,
3737 RotatingKVCache ,
38+ TokenBuffer ,
3839 load_prompt_cache ,
3940)
4041from .sample_utils import make_sampler
@@ -222,7 +223,7 @@ def setup_arg_parser():
222223
223224
224225# A stream on the default device just for generation
225- generation_stream = mx .new_stream (mx .default_device ())
226+ generation_stream = mx .new_thread_local_stream (mx .default_device ())
226227
227228
228229@contextlib .contextmanager
@@ -525,6 +526,12 @@ def speculative_generate_step(
525526 model_cache = prompt_cache [: len (model .layers )]
526527 draft_cache = prompt_cache [len (model .layers ) :]
527528
529+ if not cache .can_trim_prompt_cache (model_cache ):
530+ types = {type (c ).__name__ for c in model_cache if not c .is_trimmable ()}
531+ raise ValueError (
532+ f"Speculative decoding requires a trimmable prompt cache " f"(got { types } )."
533+ )
534+
528535 sampler = sampler or (lambda x : mx .argmax (x , axis = - 1 ))
529536
530537 quantize_cache_fn = functools .partial (
@@ -570,11 +577,12 @@ def _step(model, cache, y, n_predict=1):
570577 return _process_and_sample (None , logits .squeeze (0 ))
571578
572579 def _prefill (model , cache , y ):
573- while y .size > prefill_step_size :
574- model (y [:prefill_step_size ][None ], cache = cache )
580+ while y .size > 1 :
581+ n_to_process = min (prefill_step_size , y .size - 1 )
582+ model (y [:n_to_process ][None ], cache = cache )
575583 quantize_cache_fn (cache )
576584 mx .eval ([c .state for c in cache ])
577- y = y [prefill_step_size :]
585+ y = y [n_to_process :]
578586 mx .clear_cache ()
579587 return y
580588
@@ -1272,7 +1280,7 @@ def __init__(
12721280 self ._current_logprobs = []
12731281 self ._next_tokens = inputs
12741282 self ._next_logprobs = []
1275- self ._token_context = [mx . array ( t [ - 256 :] ) for t in tokens ]
1283+ self ._token_context = [TokenBuffer ( t ) for t in tokens ]
12761284 self ._num_tokens = [0 ] * len (self .uids )
12771285 self ._matcher_states = [m .make_state () for m in state_machines ]
12781286
@@ -1320,23 +1328,23 @@ def _step(self) -> Tuple[List[int], List[mx.array]]:
13201328 self ._current_logprobs = self ._next_logprobs
13211329 inputs = self ._current_tokens
13221330
1323- # Update the token context that will be used by the logits processors
1324- for i , ti in enumerate (self ._token_context ):
1325- self ._token_context [i ] = mx .concatenate (
1326- [ti [1 :] if len (ti ) == 256 else ti , inputs [i : i + 1 ]]
1327- )
1328-
13291331 # Forward pass
13301332 logits = self .model (inputs [:, None ], cache = self .prompt_cache )
13311333 logits = logits [:, - 1 , :]
13321334
13331335 # Logits processors
1336+ token_context = []
13341337 if any (self .logits_processors ):
1338+ # Update the token context that will be used by the logits processors
1339+ token_context = [
1340+ tc .update_and_fetch (inputs [i : i + 1 ])
1341+ for i , tc in enumerate (self ._token_context )
1342+ ]
13351343 processed_logits = []
13361344 for e in range (len (self .uids )):
13371345 sample_logits = logits [e : e + 1 ]
13381346 for processor in self .logits_processors [e ]:
1339- sample_logits = processor (self . tokens [e ], sample_logits )
1347+ sample_logits = processor (token_context [e ], sample_logits )
13401348 processed_logits .append (sample_logits )
13411349 logits = mx .concatenate (processed_logits , axis = 0 )
13421350
@@ -1358,7 +1366,7 @@ def _step(self) -> Tuple[List[int], List[mx.array]]:
13581366 # asynchronously
13591367 self ._next_tokens = sampled
13601368 self ._next_logprobs = list (logprobs )
1361- mx .async_eval (self ._next_tokens , self ._next_logprobs , self . _token_context )
1369+ mx .async_eval (self ._next_tokens , self ._next_logprobs , token_context )
13621370
13631371 # Eval the current tokens and current logprobs. After that also add
13641372 # them to self.tokens so that it always represents the tokens contained
@@ -1489,6 +1497,7 @@ class BatchGenerator:
14891497 def __init__ (
14901498 self ,
14911499 model : nn .Module ,
1500+ * ,
14921501 max_tokens : int = 128 ,
14931502 stop_tokens : Optional [Sequence [Sequence [int ]]] = None ,
14941503 sampler : Optional [Callable [[mx .array ], mx .array ]] = None ,
@@ -1498,6 +1507,8 @@ def __init__(
14981507 completion_batch_size : int = 32 ,
14991508 prefill_batch_size : int = 8 ,
15001509 prefill_step_size : int = 2048 ,
1510+ max_kv_size : Optional [int ] = None ,
1511+ stream = None ,
15011512 ):
15021513 self .model = model
15031514 self .max_tokens = max_tokens
@@ -1507,6 +1518,9 @@ def __init__(
15071518 self .prefill_step_size = prefill_step_size
15081519 self .prefill_batch_size = prefill_batch_size
15091520 self .completion_batch_size = max (completion_batch_size , prefill_batch_size )
1521+ self .max_kv_size = max_kv_size
1522+
1523+ self ._stream = stream or generation_stream
15101524
15111525 self ._default_state_machine = SequenceStateMachine (
15121526 {"normal" : [(seq , None ) for seq in stop_tokens ]} if stop_tokens else {},
@@ -1534,9 +1548,13 @@ def __init__(
15341548 else :
15351549 self ._old_wired_limit = None
15361550
1551+ @property
1552+ def stream (self ):
1553+ return self ._stream
1554+
15371555 def close (self ):
15381556 if self ._old_wired_limit is not None :
1539- mx .synchronize (generation_stream )
1557+ mx .synchronize (self . _stream )
15401558 mx .set_wired_limit (self ._old_wired_limit )
15411559 self ._old_wired_limit = None
15421560
@@ -1613,7 +1631,7 @@ def insert_segments(
16131631 caches = caches or [None ] * len (segments )
16141632 for i in range (len (segments )):
16151633 if caches [i ] is None :
1616- caches [i ] = cache . make_prompt_cache ( self .model )
1634+ caches [i ] = self ._make_new_cache ( )
16171635
16181636 for seq , m , c , at , s , lp , sm in zip (
16191637 segments ,
@@ -1636,6 +1654,19 @@ def insert_segments(
16361654
16371655 return uids
16381656
1657+ def _make_new_cache (self ):
1658+ if self .max_kv_size is None :
1659+ return cache .make_prompt_cache (self .model )
1660+
1661+ return [
1662+ (
1663+ RotatingKVCache (max_size = self .max_kv_size )
1664+ if isinstance (ci , KVCache )
1665+ else ci
1666+ )
1667+ for ci in cache .make_prompt_cache (self .model )
1668+ ]
1669+
16391670 def _find_uids (self , uids ):
16401671 uids = set (uids )
16411672 results = {}
@@ -1820,7 +1851,7 @@ def next(self):
18201851 Returns:
18211852 Tuple of prompt processing responses and generation responses.
18221853 """
1823- with mx .stream (generation_stream ):
1854+ with mx .stream (self . _stream ):
18241855 return self ._next ()
18251856
18261857 def next_generated (self ):
@@ -1830,7 +1861,7 @@ def next_generated(self):
18301861 Returns:
18311862 List of GenerationBatch.Response objects
18321863 """
1833- with mx .stream (generation_stream ):
1864+ with mx .stream (self . _stream ):
18341865 while True :
18351866 prompt_responses , generation_responses = self ._next ()
18361867 if not generation_responses and prompt_responses :
@@ -1861,7 +1892,6 @@ def batch_generate(
18611892 max_tokens : Union [int , List [int ]] = 128 ,
18621893 verbose : bool = False ,
18631894 return_prompt_caches : bool = False ,
1864- logits_processors : Optional [List [Callable [[mx .array , mx .array ], mx .array ]]] = None ,
18651895 ** kwargs ,
18661896) -> BatchResponse :
18671897 """
@@ -1880,8 +1910,6 @@ def batch_generate(
18801910 can be per prompt if a list is provided.
18811911 return_prompt_caches (bool): Return the prompt caches in the batch
18821912 responses. Default: ``False``.
1883- logits_processors (List[Callable[[mx.array, mx.array], mx.array]], optional):
1884- A list of functions that take tokens and logits and return the processed logits. Default: ``None``.
18851913 kwargs: The remaining options get passed to :obj:`BatchGenerator`.
18861914 See :obj:`BatchGenerator` for more details.
18871915 """
0 commit comments