Skip to content

Commit 5c965a5

Browse files
authored
Merge branch 'ml-explore:main' into main
2 parents e06ca01 + ed1fca4 commit 5c965a5

22 files changed

Lines changed: 2245 additions & 239 deletions

‎mlx_lm/_version.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,3 @@
11
# Copyright © 2023-2025 Apple Inc.
22

3-
__version__ = "0.31.2"
3+
__version__ = "0.31.3"

‎mlx_lm/generate.py‎

Lines changed: 48 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -35,6 +35,7 @@
3535
KVCache,
3636
QuantizedKVCache,
3737
RotatingKVCache,
38+
TokenBuffer,
3839
load_prompt_cache,
3940
)
4041
from .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
"""

‎mlx_lm/models/apertus.py‎

Lines changed: 9 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -167,20 +167,27 @@ def __init__(self, args: ModelArgs):
167167
self.args = args
168168
self.model_type = args.model_type
169169
self.model = ApertusModel(args)
170-
self.lm_head = nn.Linear(args.hidden_size, args.vocab_size, bias=False)
170+
if not args.tie_word_embeddings:
171+
self.lm_head = nn.Linear(args.hidden_size, args.vocab_size, bias=False)
171172

172173
def __call__(
173174
self,
174175
inputs: mx.array,
175176
cache: Optional[Any] = None,
176177
) -> mx.array:
177178
out = self.model(inputs, cache)
178-
return self.lm_head(out)
179+
if self.args.tie_word_embeddings:
180+
out = self.model.embed_tokens.as_linear(out)
181+
else:
182+
out = self.lm_head(out)
183+
return out
179184

180185
def sanitize(self, weights):
181186
for k, v in weights.items():
182187
if k.endswith("alpha_p") or k.endswith("alpha_n"):
183188
weights[k] = v.squeeze()
189+
if self.args.tie_word_embeddings:
190+
weights.pop("lm_head.weight", None)
184191
return weights
185192

186193
@property

0 commit comments

Comments
 (0)