File size: 56,532 Bytes
e4f7326
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
911
912
913
914
915
916
917
918
919
920
921
922
923
924
925
926
927
928
929
930
931
932
933
934
935
936
937
938
939
940
941
942
943
944
945
946
947
948
949
950
951
952
953
954
955
956
957
958
959
960
961
962
963
964
965
966
967
968
969
970
971
972
973
974
975
976
977
978
979
980
981
982
983
984
985
986
987
988
989
990
991
992
993
994
995
996
997
998
999
1000
1001
1002
1003
1004
1005
1006
1007
1008
1009
1010
1011
1012
1013
1014
1015
1016
1017
1018
1019
1020
1021
1022
1023
1024
1025
1026
1027
1028
1029
1030
1031
1032
1033
1034
1035
1036
1037
1038
1039
1040
1041
1042
1043
1044
1045
1046
1047
1048
1049
1050
1051
1052
1053
1054
1055
1056
1057
1058
1059
1060
1061
1062
1063
1064
1065
1066
1067
1068
1069
1070
1071
1072
1073
1074
1075
1076
1077
1078
1079
1080
1081
1082
1083
1084
1085
1086
1087
1088
1089
1090
1091
1092
1093
1094
1095
1096
1097
1098
1099
1100
1101
1102
1103
1104
1105
1106
1107
1108
1109
1110
1111
1112
1113
1114
1115
1116
1117
1118
1119
1120
1121
1122
1123
1124
1125
1126
1127
1128
1129
1130
1131
1132
1133
1134
1135
1136
1137
1138
1139
1140
1141
1142
1143
1144
1145
1146
1147
1148
1149
1150
1151
1152
1153
1154
1155
1156
1157
1158
1159
1160
1161
1162
1163
1164
1165
1166
1167
1168
1169
1170
1171
1172
1173
1174
1175
1176
1177
1178
1179
1180
1181
1182
1183
1184
1185
1186
1187
"""Native MLX schema25 inference with bounded KV reuse and continuous admission."""

from __future__ import annotations

import argparse
import copy
import hashlib
import json
import math
import queue
import sys
import threading
import time
from collections import OrderedDict, deque
from dataclasses import dataclass, fields, replace
from pathlib import Path
from typing import Any, Iterator

import mlx.core as mx

from modilify_mk2.configuration_modilify_mk2 import DENOISE_TEMPERATURE
from modilify_mk2.mlx_commit_policy import (
    fused_commit_failure_rate, infer_commit_reason, select_commit_lengths,
)
from modilify_mk2.mlx_model import MLXCanvasOutput, MLXModilifyMk2
from modilify_mk2.runtime import MLXRuntime, load_runtime
from modilify_mk2.mlx_state import MLXLatentState, MLXRollingState
from modilify_mk2.chat import apply_chat_template
def parse_bool(value: str | bool) -> bool:
    """Parse explicit CLI booleans such as ``--think true``."""

    if isinstance(value, bool):
        return value
    normalized = value.strip().lower()
    if normalized in {"1", "true", "yes", "on"}:
        return True
    if normalized in {"0", "false", "no", "off"}:
        return False
    raise argparse.ArgumentTypeError("expected true or false")


REQUEST_FIELDS = frozenset(
    {
        "request_id",
        "prompt",
        "messages",
        "max_new_tokens",
        "max_denoising_steps",
        "seed",
        "think",
    }
)


@dataclass(frozen=True)
class ContinuousRequest:
    """One validated machine-mode request before chat-template tokenization."""

    request_id: str
    messages: list[dict[str, Any]]
    max_new_tokens: int
    max_denoising_steps: int | None
    seed: int
    think: bool
    prompt: str | None = None


def stable_request_seed(base_seed: int, request_id: str) -> int:
    """Derive a scheduler-independent non-negative seed from a stable request ID."""

    payload = f"{base_seed}\0{request_id}".encode("utf-8")
    return int.from_bytes(hashlib.sha256(payload).digest()[:8], "big") & ((1 << 63) - 1)


def _positive_int(record: dict[str, Any], name: str, default: int | None) -> int | None:
    value = record.get(name, default)
    if value is None:
        return default
    if isinstance(value, bool) or not isinstance(value, int) or value <= 0:
        raise ValueError(f"`{name}` must be a positive integer or null.")
    return value


def parse_continuous_request(
    record: Any,
    *,
    default_max_new_tokens: int,
    default_max_denoising_steps: int | None,
    default_seed: int,
    default_think: bool,
) -> ContinuousRequest:
    """Validate one strict request object without silently inventing identity."""

    if not isinstance(record, dict):
        raise ValueError("Each request line must be a JSON object.")
    unknown = sorted(set(record).difference(REQUEST_FIELDS))
    if unknown:
        raise ValueError(f"Unknown request fields: {unknown}")
    request_id = record.get("request_id")
    if not isinstance(request_id, str) or not request_id.strip():
        raise ValueError("`request_id` must be a non-empty string.")
    has_prompt = "prompt" in record
    has_messages = "messages" in record
    if has_prompt == has_messages:
        raise ValueError("Exactly one of `prompt` or `messages` is required.")

    prompt = None
    if has_prompt:
        prompt = record["prompt"]
        if not isinstance(prompt, str):
            raise ValueError("`prompt` must be a string.")
        messages = [{"role": "user", "content": prompt}]
    else:
        messages = record["messages"]
        if not isinstance(messages, list) or not messages:
            raise ValueError("`messages` must be a non-empty list.")
        if any(
            not isinstance(message, dict)
            or "role" not in message
            or "content" not in message
            for message in messages
        ):
            raise ValueError("Every message must be an object with `role` and `content`.")

    think = record.get("think", default_think)
    if not isinstance(think, bool):
        raise ValueError("`think` must be a boolean.")
    supplied_seed = record.get("seed")
    if supplied_seed is not None and (
        isinstance(supplied_seed, bool) or not isinstance(supplied_seed, int)
    ):
        raise ValueError("`seed` must be an integer or null.")
    seed = (
        stable_request_seed(default_seed, request_id)
        if supplied_seed is None
        else supplied_seed
    )
    return ContinuousRequest(
        request_id=request_id,
        prompt=prompt,
        messages=messages,
        max_new_tokens=int(
            _positive_int(record, "max_new_tokens", default_max_new_tokens)
        ),
        max_denoising_steps=_positive_int(
            record, "max_denoising_steps", default_max_denoising_steps
        ),
        seed=seed,
        think=think,
    )


def _extract_input_ids(encoded: Any) -> list[int]:
    value = encoded.get("input_ids") if isinstance(encoded, dict) else encoded
    if hasattr(encoded, "input_ids"):
        value = encoded.input_ids
    if isinstance(value, tuple):
        value = list(value)
    if isinstance(value, list) and len(value) == 1 and isinstance(value[0], list):
        value = value[0]
    if not isinstance(value, list) or any(
        isinstance(token, bool) or not isinstance(token, int) for token in value
    ):
        raise ValueError("Chat template did not return one integer token sequence.")
    if not value:
        raise ValueError("Chat template returned an empty prompt.")
    return [int(token) for token in value]


def _seed(value: int, salt: int) -> mx.array:
    return mx.random.key((int(value) + int(salt)) % (2**32))


def _logical(value: mx.array, head: mx.array) -> mx.array:
    canvas = value.shape[1]
    index = (head[:, None] + mx.arange(canvas)[None, :]) % canvas
    return mx.take_along_axis(value, index, axis=1)


def _physical(value: mx.array, head: mx.array) -> mx.array:
    canvas = value.shape[1]
    index = (mx.arange(canvas)[None, :] - head[:, None]) % canvas
    return mx.take_along_axis(value, index, axis=1)


def _concat_rows(values: list[Any]) -> Any:
    first = values[0]
    if isinstance(first, mx.array):
        return mx.concatenate(values, axis=0)
    return type(first)(**{
        item.name: _concat_rows([getattr(value, item.name) for value in values])
        for item in fields(first)
    })


def _slice_row(value: Any, row: int) -> Any:
    if isinstance(value, mx.array):
        return value[row:row + 1]
    return type(value)(**{
        item.name: _slice_row(getattr(value, item.name), row)
        for item in fields(value)
    })


def _arrays(value: Any) -> list[mx.array]:
    if isinstance(value, mx.array):
        return [value]
    return [array for item in fields(value)
            for array in _arrays(getattr(value, item.name))]


def _clone_cache(cache: list[Any], *, compact: bool = False) -> list[Any]:
    copied = []
    arrays = []
    for layer in cache:
        clone = copy.copy(layer)
        source = layer.state if compact else (layer.keys, layer.values)
        if source[0] is not None:
            clone.keys = mx.array(source[0])
            clone.values = mx.array(source[1])
            arrays.extend((clone.keys, clone.values))
        copied.append(clone)
    if arrays:
        mx.eval(*arrays)
    return copied


@dataclass
class _PrefixEntry:
    cache: list[Any]
    bytes: int


class PrefixKVCache:
    """LRU cache of immutable, block-boundary prompt KV snapshots."""

    def __init__(self, limit_bytes: int):
        self.limit_bytes = limit_bytes
        self.entries: OrderedDict[tuple[int, ...], _PrefixEntry] = OrderedDict()
        self.bytes = 0
        self.hits = 0
        self.reused_tokens = 0

    def longest(self, ids: tuple[int, ...]) -> tuple[int, list[Any] | None]:
        lengths = sorted(
            {len(key) for key in self.entries if len(key) <= len(ids)},
            reverse=True,
        )
        for length in lengths:
            key = ids[:length]
            entry = self.entries.get(key)
            if entry is not None:
                self.entries.move_to_end(key)
                self.hits += 1
                self.reused_tokens += length
                return length, _clone_cache(entry.cache)
        return 0, None

    def put(self, ids: tuple[int, ...], cache: list[Any]) -> None:
        if self.limit_bytes <= 0 or ids in self.entries:
            return
        size = sum(int(value.nbytes) for layer in cache
                   for value in (layer.state if layer.keys is not None else ())
                   if value is not None)
        if size > self.limit_bytes:
            return
        snapshot = _clone_cache(cache, compact=True)
        while self.entries and self.bytes + size > self.limit_bytes:
            _, victim = self.entries.popitem(last=False)
            self.bytes -= victim.bytes
        self.entries[ids] = _PrefixEntry(snapshot, size)
        self.bytes += size


@dataclass
class InferenceRow:
    request: ContinuousRequest
    prompt_ids: tuple[int, ...]
    cache: list[Any]
    rolling: MLXRollingState
    generated: list[int]
    repetition_seen: set[int]
    created: float
    admitted: float
    denoise_steps: int = 0
    jumps: int = 0
    shifts: int = 0
    first_token_at: float | None = None
    stop_reason: str | None = None
    first_denoise_at: float | None = None
    last_denoise_at: float | None = None
    prefill_seconds: float = 0.0


@dataclass
class PrefillRow:
    request: ContinuousRequest
    prompt_ids: tuple[int, ...]
    cache: list[Any]
    offset: int
    reused: int
    created: float
    seconds: float = 0.0


@dataclass
class _DenoiseWork:
    rows: list[InferenceRow]
    rolling: MLXRollingState
    output: Any
    proposal: mx.array
    remaining: mx.array
    physical_positions: mx.array
    next_latent: MLXLatentState
    policy: Any


def _forward_independent_rows(model: MLXModilifyMk2,
                              rows: list[InferenceRow]) -> list[MLXCanvasOutput]:
    """Interleave decoder layers while preserving independent singleton rows."""
    decoder = model.model.decoder
    latent = model.latent_deliberation
    working_bus, persistent_bus = latent.working_memory_bus, latent.persistent_memory_bus
    contexts = []
    for row in rows:
        rolling = row.rolling
        canvas, state, head = rolling.canvas, rolling.latent, rolling.head
        batch, length = canvas.shape
        tokens = decoder.embed_tokens(canvas) * decoder.embed_scale
        working, next_state = latent(token_embeddings=tokens,
                                     confidence=state.confidence,
                                     entropy=state.entropy, state=state,
                                     canvas_head=head)
        physical = (head[:, None] + mx.arange(length)[None, :]) % length
        gather = mx.broadcast_to(physical[:, :, None], tokens.shape)
        logical_tokens = mx.take_along_axis(tokens, mx.stop_gradient(gather), axis=1)
        logical_working = mx.take_along_axis(working, mx.stop_gradient(gather), axis=1)
        hidden = model._merge_context(logical_tokens, logical_working)
        prefix_length = int(row.cache[0].offset)
        full_mask = mx.ones((batch, prefix_length + length), mx.bool_)
        masks = decoder._make_decoder_masks(hidden, row.cache, full_mask)
        seen = latent.logical_seen(state, head)
        contexts.append({"hidden": hidden, "working": working, "state": next_state,
                         "tokens": tokens, "head": head, "masks": masks,
                         "offset": prefix_length,
                         "working_kv": working_bus.prepare_kv((logical_working, seen)),
                         "persistent_kv": persistent_bus.prepare_kv(state.memory_slots, seen)})
    reader = 0
    for index, layer in enumerate(decoder.layers):
        for row, context in zip(rows, contexts, strict=True):
            hidden = layer(context["hidden"], context["masks"][layer.layer_type],
                           row.cache[index], decoder=True, offset=context["offset"])
            if layer.layer_type == "full_attention":
                if context["working_kv"] is not None and reader < working_bus.num_readers:
                    hidden = working_bus.read(hidden, reader, *context["working_kv"])
                if context["persistent_kv"] is not None and reader < persistent_bus.num_readers:
                    hidden = persistent_bus.read(hidden, reader, *context["persistent_kv"])
            context["hidden"] = hidden
        # Explicit layer boundaries keep the execution order interleaved and
        # bound intermediate activations to the configured pipeline depth.
        mx.async_eval(*(context["hidden"] for context in contexts))
        if layer.layer_type == "full_attention":
            reader += 1
    outputs = []
    for context in contexts:
        hidden = decoder.norm(context["hidden"])
        length = hidden.shape[1]
        inverse = (mx.arange(length)[None, :] - context["head"][:, None]) % length
        heavy = mx.take_along_axis(hidden, mx.stop_gradient(mx.broadcast_to(
            inverse[:, :, None], hidden.shape)), axis=1)
        outputs.append(MLXCanvasOutput(heavy, context["working"],
                                       context["state"], context["tokens"]))
    return outputs


def _noise(seed: int, step: int, canvas: int, vocab: int) -> mx.array:
    return mx.random.randint(0, vocab, (1, canvas),
                             key=_seed(seed, 0x51A7 + step * 1000003))


def _empty_rolling(config: Any, seed: int, max_new_tokens: int,
                   pad_token_id: int) -> MLXRollingState:
    canvas = int(config.canvas_length)
    vocab = int(config.text_config.vocab_size)
    latent = MLXLatentState.empty(1, canvas, int(config.latent_memory_slots),
                                  int(config.latent_dim), enable_gdn2=True)
    latent = replace(latent, entropy=mx.full((1, canvas), math.log(vocab), mx.float32))
    initial = mx.where(mx.arange(canvas)[None, :] < max_new_tokens,
                       _noise(seed, 0, canvas, vocab), pad_token_id)
    return MLXRollingState(initial, latent,
                            mx.zeros((1,), mx.int32))


def _inference_statistics(hidden: mx.array, weight: mx.array,
                          rows: list[InferenceRow], *, softcap: float,
                          chunk_size: int, penalty: float,
                          excluded: set[int],
                          top_k: int | None = 40,
                          min_p: float | None = 0.05) -> tuple[mx.array, ...]:
    """Exact chunked Top-K and Min-P Gumbel sampling without a retained full-vocabulary matrix."""
    batch, canvas, dim = hidden.shape
    flat = hidden.reshape(batch * canvas, dim)
    neg_inf = mx.full((batch * canvas,), -mx.inf, mx.float32)
    log_z = neg_inf
    best_gumbel = neg_inf
    chosen_score = mx.zeros_like(log_z)
    chosen = mx.zeros((batch * canvas,), mx.int32)
    greedy_score = neg_inf
    greedy = mx.zeros_like(chosen)
    moment_max = neg_inf
    moment_sum = mx.zeros_like(log_z)
    moment_weighted = mx.zeros_like(log_z)
    repetition_mask = None
    if penalty != 1.0:
        masks = []
        for row in rows:
            mask = mx.zeros((weight.shape[0],), mx.bool_)
            eligible = sorted(row.repetition_seen - excluded)
            if eligible:
                mask[mx.array(eligible, mx.int32)] = True
            masks.append(mask)
        repetition_mask = mx.stack(masks, axis=0)

    use_constrained = top_k is not None and top_k > 0
    chunk_cand_scores = []
    chunk_cand_tokens = []

    for chunk_number, start in enumerate(range(0, weight.shape[0], chunk_size)):
        stop = min(start + chunk_size, weight.shape[0])
        score = mx.tanh((flat @ weight[start:stop].T).astype(mx.float32) / softcap) * softcap
        if penalty != 1.0:
            assert repetition_mask is not None
            row_scores = score.reshape(batch, canvas, stop - start)
            seen = repetition_mask[:, None, start:stop]
            changed = mx.where(row_scores < 0, row_scores * penalty, row_scores / penalty)
            score = mx.where(seen, changed, row_scores).reshape(batch * canvas, stop - start)
        score = score / DENOISE_TEMPERATURE
        log_z = mx.logaddexp(log_z, mx.logsumexp(score, axis=-1))

        local_max = mx.max(score, axis=-1)
        better_greedy = local_max > greedy_score
        greedy = mx.where(better_greedy, mx.argmax(score, axis=-1).astype(mx.int32) + start, greedy)
        greedy_score = mx.maximum(greedy_score, local_max)

        shifted = mx.exp(score - local_max[:, None])
        next_max = mx.maximum(moment_max, local_max)
        old_scale = mx.exp(moment_max - next_max)
        new_scale = mx.exp(local_max - next_max)
        moment_sum = moment_sum * old_scale + mx.sum(shifted, axis=-1) * new_scale
        moment_weighted = moment_weighted * old_scale + mx.sum(shifted * score, axis=-1) * new_scale
        moment_max = next_max

        if use_constrained:
            k = min(top_k, stop - start)
            chunk_top_idx = mx.stop_gradient(mx.argpartition(-score, kth=k - 1, axis=-1)[:, :k]).astype(mx.int32)
            chunk_top_score = mx.take_along_axis(score, chunk_top_idx, axis=-1)
            chunk_cand_scores.append(chunk_top_score)
            chunk_cand_tokens.append(chunk_top_idx + start)
        else:
            uniform = mx.concatenate([
                mx.random.uniform(shape=(canvas, stop - start),
                                  key=_seed(row.request.seed,
                                            0xC09A + row.denoise_steps * 1000003 + chunk_number))
                for row in rows
            ], axis=0)
            uniform = mx.clip(uniform, 1.17549435e-38, 1.0 - 1.19209290e-7)
            gumbel = score - mx.log(-mx.log(uniform))
            local_index = mx.argmax(gumbel, axis=-1)
            local_best = mx.max(gumbel, axis=-1)
            better = local_best > best_gumbel
            chosen_score = mx.where(better, mx.take_along_axis(score, local_index[:, None], axis=-1)[:, 0], chosen_score)
            chosen = mx.where(better, local_index + start, chosen)
            best_gumbel = mx.maximum(best_gumbel, local_best)

        # Materialize each chunk without a CPU/GPU barrier. Include candidates
        # so their argpartition graph does not retain every vocabulary slab.
        live = [log_z, greedy, greedy_score, moment_max, moment_sum, moment_weighted]
        if use_constrained:
            live.extend((chunk_cand_scores[-1], chunk_cand_tokens[-1]))
        else:
            live.extend((best_gumbel, chosen_score, chosen))
        mx.async_eval(*live)

    if use_constrained:
        all_cand_scores = mx.concatenate(chunk_cand_scores, axis=-1)
        all_cand_tokens = mx.concatenate(chunk_cand_tokens, axis=-1)
        global_k = min(top_k, all_cand_scores.shape[-1])
        global_idx = mx.stop_gradient(mx.argpartition(-all_cand_scores, kth=global_k - 1, axis=-1)[:, :global_k]).astype(mx.int32)
        cand_scores = mx.take_along_axis(all_cand_scores, global_idx, axis=-1)
        cand_tokens = mx.take_along_axis(all_cand_tokens, global_idx, axis=-1)
        if min_p is not None and min_p > 0.0:
            thresh = greedy_score + math.log(float(min_p))
            valid_cand = cand_scores >= thresh[:, None]
            eligible_scores = mx.where(valid_cand, cand_scores, -mx.inf)
        else:
            eligible_scores = cand_scores
        uniform = mx.concatenate([
            mx.random.uniform(shape=(canvas, global_k),
                              key=_seed(row.request.seed,
                                        0xC09A + row.denoise_steps * 1000003))
            for row in rows
        ], axis=0)
        uniform = mx.clip(uniform, 1.17549435e-38, 1.0 - 1.19209290e-7)
        gumbel = eligible_scores - mx.log(-mx.log(uniform))
        chosen_local = mx.stop_gradient(mx.argmax(gumbel, axis=-1))
        chosen_score = mx.take_along_axis(cand_scores, chosen_local[:, None], axis=-1)[:, 0]
        chosen = mx.take_along_axis(cand_tokens, chosen_local[:, None], axis=-1)[:, 0]

    confidence = mx.clip(mx.exp(chosen_score - log_z), 0.0, 1.0)
    greedy_confidence = mx.clip(mx.exp(greedy_score - log_z), 0.0, 1.0)
    entropy = log_z - moment_weighted / mx.maximum(moment_sum, 1.17549435e-38)
    shape = (batch, canvas)
    return (chosen.reshape(shape), confidence.reshape(shape), entropy.reshape(shape),
            greedy.reshape(shape), greedy_confidence.reshape(shape))


class MLXContinuousEngine:
    def __init__(self, runtime: MLXRuntime, *, prefix_mib: int,
                 prefill_chunk: int, vocab_chunk: int, max_batch_rows: int,
                 max_batch_tokens: int, emit: Any, pipeline_depth: int = 2,
                 token_events: bool = True):
        self.runtime = runtime
        if min(prefill_chunk, vocab_chunk, max_batch_rows, pipeline_depth) <= 0:
            raise ValueError("Chunk sizes, batch rows and pipeline depth must be positive.")
        if max_batch_tokens < int(runtime.config.canvas_length):
            raise ValueError("--max-batch-tokens must fit at least one canvas.")
        self.prefix = PrefixKVCache(prefix_mib * 1024 * 1024)
        self.prefill_chunk = prefill_chunk
        self.vocab_chunk = vocab_chunk
        self.max_batch_rows = max_batch_rows
        self.max_batch_tokens = max_batch_tokens
        self.emit = emit
        self.token_events = token_events
        self.pipeline_depth = pipeline_depth
        self.forward_count = 0
        self.active_row_steps = 0
        self.denoise_seconds = 0.0
        self.max_inflight = 0
        self.peak_active_requests = 0
        generation = runtime.generation
        pad_token_id = generation.pad_token_id
        if pad_token_id is None:
            pad_token_id = getattr(runtime.config, "pad_token_id", None)
        if isinstance(pad_token_id, (list, tuple)):
            pad_token_id = pad_token_id[0]
        self.pad_token_id = int(0 if pad_token_id is None else pad_token_id)
        configured_eos = generation.eos_token_id or runtime.config.eos_token_id
        if isinstance(configured_eos, int):
            configured_eos = [configured_eos]
        self.turn_end = int(runtime.config.turn_end_token_id if
                            generation.turn_end_token_id is None else
                            generation.turn_end_token_id)
        self.stops = tuple(dict.fromkeys((self.turn_end, *(int(x) for x in configured_eos or ()))))
        token_values = [generation.pad_token_id, generation.bos_token_id,
                        generation.eos_token_id, generation.turn_end_token_id,
                        getattr(runtime.config, "image_token_id", None),
                        generation.repetition_penalty_exclude_token_ids]
        self.excluded = set()
        for value in token_values:
            if isinstance(value, int):
                self.excluded.add(int(value))
            elif isinstance(value, (tuple, list, set)):
                self.excluded.update(int(x) for x in value if x is not None)

    def begin_prefill(self, request: ContinuousRequest, created: float) -> PrefillRow:
        encoded = apply_chat_template(self.runtime.tokenizer, request.messages,
                                      think=request.think)
        prompt = tuple(_extract_input_ids(encoded))
        if len(prompt) + request.max_new_tokens > int(self.runtime.config.text_config.max_position_embeddings):
            raise ValueError("Prompt plus response exceeds the model position limit.")
        prefix_length, cache = self.prefix.longest(prompt)
        if cache is None:
            cache = self.runtime.model.model.encoder.make_cache()
        return PrefillRow(request, prompt, cache, prefix_length, prefix_length, created)

    def prefill_step(self, work: PrefillRow) -> InferenceRow | None:
        """Execute at most one prefill chunk before yielding to active rows."""
        started = time.perf_counter()
        if work.offset < len(work.prompt_ids):
            end = min(work.offset + self.prefill_chunk, len(work.prompt_ids))
            block = mx.array(work.prompt_ids[work.offset:end], mx.int32)[None, :]
            _, work.cache = self.runtime.model.model.encoder(block, cache=work.cache)
            mx.eval(*(value for layer in work.cache for value in (layer.keys, layer.values)
                      if isinstance(value, mx.array)))
            work.offset = end
            self.prefix.put(work.prompt_ids[:end], work.cache)
        if work.offset < len(work.prompt_ids):
            work.seconds += time.perf_counter() - started
            return None
        request, prompt = work.request, work.prompt_ids
        cache = _clone_cache(work.cache, compact=True)
        work.seconds += time.perf_counter() - started
        seen = {token for token in prompt if token not in self.excluded}
        row = InferenceRow(request, prompt, cache,
                           _empty_rolling(self.runtime.config, request.seed,
                                          request.max_new_tokens, self.pad_token_id), [],
                           seen, work.created, time.perf_counter(),
                           prefill_seconds=work.seconds)
        self.emit({"event": "request_started", "request_id": request.request_id,
                   "prompt_tokens": len(prompt), "prefix_cache_tokens": work.reused,
                   "prefill_seconds": row.prefill_seconds})
        return row

    def admit(self, request: ContinuousRequest, created: float) -> InferenceRow:
        work = self.begin_prefill(request, created)
        while True:
            row = self.prefill_step(work)
            if row is not None:
                return row

    def _append_encoder(self, rows: list[InferenceRow], blocks: list[list[int]]) -> None:
        # Keep the encoder's GEMM and expert routing shapes independent of the
        # cohort too. Only the scheduler and GPU submission are concurrent.
        for row, tokens in zip(rows, blocks, strict=True):
            if tokens:
                ids = mx.array(tokens, mx.int32)[None, :]
                _, row.cache = self.runtime.model.model.encoder(ids, cache=row.cache)
                mx.async_eval(*(value for layer in row.cache
                                for value in (layer.keys, layer.values)
                                if isinstance(value, mx.array)))

    def step(self, rows: list[InferenceRow]) -> list[InferenceRow]:
        if not rows:
            return []
        if len(rows) > min(self.max_batch_rows, self.max_batch_tokens //
                           int(self.runtime.config.canvas_length)):
            raise ValueError("Denoise cohort exceeds the configured row/token budget.")
        if len({id(row) for row in rows}) != len(rows):
            raise ValueError("A request may appear only once in a denoise cohort.")
        started = time.perf_counter()
        finished = []
        for start in range(0, len(rows), self.pipeline_depth):
            cohort = rows[start:start + self.pipeline_depth]
            outputs = _forward_independent_rows(self.runtime.model, cohort)
            pending = [self._prepare_step([row], output=output)
                       for row, output in zip(cohort, outputs, strict=True)]
            self.max_inflight = max(self.max_inflight, len(pending))
            for work in pending:
                finished.extend(self._finish_step(work))
        self.denoise_seconds += time.perf_counter() - started
        return finished

    def _prepare_step(self, rows: list[InferenceRow], *,
                      output: MLXCanvasOutput) -> _DenoiseWork:
        model, config, generation = (self.runtime.model, self.runtime.config,
                                     self.runtime.generation)
        rolling = _concat_rows([row.rolling for row in rows])
        batch, canvas = rolling.canvas.shape
        stats = _inference_statistics(
            output.heavy_hidden, model.model.decoder.embed_tokens.weight,
            rows, softcap=float(config.text_config.final_logit_softcapping),
            chunk_size=self.vocab_chunk, penalty=float(generation.repetition_penalty),
            excluded=self.excluded,
            top_k=getattr(config, "commit_top_k", 40),
            min_p=getattr(config, "commit_min_p", 0.05),
        )
        proposal, confidence, entropy, greedy, greedy_confidence = stats
        confidence = confidence.astype(mx.float32)
        entropy = entropy.astype(mx.float32)
        changed = (proposal != rolling.canvas).astype(mx.float32)
        remaining = mx.array([row.request.max_new_tokens - len(row.generated)
                              for row in rows], mx.int32)
        physical_positions = (mx.arange(canvas)[None, :] - rolling.head[:, None]) % canvas
        valid = physical_positions < remaining[:, None]
        next_latent = replace(
            output.next_latent_state, confidence=confidence, entropy=entropy,
            age=rolling.latent.age + 1, token_changed=changed,
            confidence_delta=confidence - rolling.latent.confidence,
            entropy_delta=entropy - rolling.latent.entropy,
        )
        next_latent = _concat_rows([
            model.latent_deliberation.observe_state(
                _slice_row(next_latent, index), output.heavy_hidden[index:index + 1],
                output.working_state[index:index + 1], valid[index:index + 1],
                rolling.head[index:index + 1]) for index in range(batch)
        ])
        policy = select_commit_lengths(
            _logical(proposal, rolling.head),
            _logical(fused_commit_failure_rate(
                confidence, entropy, entropy_weight=config.commit_entropy_weight,
                confidence_power=config.commit_confidence_power,
                top_k=getattr(config, "commit_top_k", None),
                min_p=getattr(config, "commit_min_p", None),
                target_confidence=getattr(config, "commit_target_confidence", None),
                failure_budget=float(config.commit_failure_budget)), rolling.head),
            _logical(fused_commit_failure_rate(
                rolling.latent.confidence, rolling.latent.entropy,
                entropy_weight=config.commit_entropy_weight,
                confidence_power=config.commit_confidence_power,
                top_k=getattr(config, "commit_top_k", None),
                min_p=getattr(config, "commit_min_p", None),
                target_confidence=getattr(config, "commit_target_confidence", None),
                failure_budget=float(config.commit_failure_budget)), rolling.head),
            _logical(greedy, rolling.head),
            _logical(fused_commit_failure_rate(
                greedy_confidence, entropy, entropy_weight=config.commit_entropy_weight,
                confidence_power=config.commit_confidence_power,
                top_k=getattr(config, "commit_top_k", None),
                min_p=getattr(config, "commit_min_p", None),
                target_confidence=getattr(config, "commit_target_confidence", None),
                failure_budget=float(config.commit_failure_budget)), rolling.head),
            ponder_steps=rolling.latent.ponder_steps,
            stagnation_steps=rolling.latent.stagnation_steps,
            active_rows=mx.ones((batch,), mx.bool_), remaining_lengths=remaining,
            failure_budget=float(config.commit_failure_budget),
            stop_token_id=self.stops,
            stagnation_threshold=int(generation.jump_on_no_progress_after),
            min_progress=float(generation.min_trajectory_progress),
            max_ponder_steps=int(generation.max_ponder_steps),
            valid_mask=mx.arange(canvas)[None, :] < remaining[:, None],
        )
        mx.async_eval(policy.commit_lengths, policy.commit_token_ids,
                      policy.jump_rows, policy.ponder_steps, policy.stagnation_steps,
                      *_arrays(next_latent), output.heavy_hidden, output.working_state)
        self.forward_count += 1
        return _DenoiseWork(rows, rolling, output, proposal, remaining,
                            physical_positions, next_latent, policy)

    def _finish_step(self, work: _DenoiseWork) -> list[InferenceRow]:
        rows, rolling, output = work.rows, work.rolling, work.output
        proposal, remaining = work.proposal, work.remaining
        physical_positions, next_latent, policy = (
            work.physical_positions, work.next_latent, work.policy)
        model, config, generation = (self.runtime.model, self.runtime.config,
                                     self.runtime.generation)
        batch, canvas = rolling.canvas.shape
        lengths = policy.commit_lengths.astype(mx.int32)
        selected = policy.commit_token_ids
        mx.eval(lengths, selected, policy.jump_rows)
        completed = time.perf_counter()
        for row in rows:
            row.last_denoise_at = completed
            if row.first_denoise_at is None:
                row.first_denoise_at = completed
                self.emit({"event": "first_denoise", "request_id": row.request.request_id,
                           "ttfd_seconds": completed - row.created})
        lengths_host = [int(x) for x in lengths.tolist()]
        maximum = max(lengths_host)
        selected_host = selected.tolist()
        blocks = [[int(x) for x in selected_host[index][:lengths_host[index]]]
                  for index in range(batch)]
        # The commit policy has fixed these tokens. Stream them before the
        # writer and encoder append, which are only needed by the next step.
        for index, row in enumerate(rows):
            block = blocks[index]
            if not block:
                continue
            if row.first_token_at is None:
                row.first_token_at = time.perf_counter()
            row.generated.extend(block)
            row.repetition_seen.update(token for token in block
                                       if token not in self.excluded)
            if self.token_events:
                self.emit({"event": "token", "request_id": row.request.request_id,
                           "token_ids": block, "generated_tokens": len(row.generated),
                           "text": self.runtime.tokenizer.decode(
                               row.generated, skip_special_tokens=False),
                           "denoise_steps": row.denoise_steps + 1})
        next_canvas = mx.where(
            (physical_positions < lengths[:, None]) & policy.jump_rows[:, None],
            _physical(selected, rolling.head), proposal,
        )
        next_latent = replace(next_latent, ponder_steps=policy.ponder_steps,
                              stagnation_steps=policy.stagnation_steps)
        if maximum:
            embeddings = model.model.decoder.embed_tokens(selected[:, :maximum])
            embeddings = embeddings * model.model.decoder.embed_scale
            memory, _ = model.latent_deliberation.commit_write(
                memory=next_latent.memory_slots,
                working_state=output.working_state,


                heavy_hidden=output.heavy_hidden,
                committed_token_embeddings=embeddings,
                commit_lengths=lengths,
                prefix_lengths=mx.array([int(row.cache[0].offset) for row in rows], mx.int32),
                commit_reason=infer_commit_reason(
                    lengths, jump_rows=policy.jump_rows,
                    commit_token_ids=selected, terminal_token_ids=self.stops,
                ),
                canvas_head=rolling.head, max_commit=maximum,
            )
            next_latent = replace(
                next_latent, memory_slots=memory,
                gdn2=replace(next_latent.gdn2, persistent=memory),
            )
        self._append_encoder(rows, blocks)
        noise = mx.concatenate([
            _noise(row.request.seed, row.denoise_steps + 1, canvas,
                   int(config.text_config.vocab_size)) for row in rows
        ], axis=0)
        next_rolling = MLXRollingState(next_canvas, next_latent,
                                        rolling.head).advance_ring(
            lengths, noise, entropy_fill_value=math.log(config.text_config.vocab_size),
        )
        logical_positions = (mx.arange(canvas)[None, :] - next_rolling.head[:, None]) % canvas
        newly_exposed = logical_positions >= (canvas - lengths)[:, None]
        next_rolling = replace(
            next_rolling,
            canvas=mx.where(newly_exposed &
                            (logical_positions >= (remaining - lengths)[:, None]),
                            self.pad_token_id, next_rolling.canvas),
        )
        # The following step depends on these arrays on the same MLX stream.
        # Submit now and overlap the writer/refill with host stop/output work.
        mx.async_eval(*_arrays(next_rolling))
        finished = []
        jumped = policy.jump_rows.tolist()
        for index, row in enumerate(rows):
            row.rolling = _slice_row(next_rolling, index)
            row.denoise_steps += 1
            self.active_row_steps += 1
            row.jumps += int(jumped[index])
            row.shifts += int(bool(blocks[index]))
            if self.turn_end in blocks[index]:
                row.stop_reason = "turn_end"
            elif any(token in self.stops for token in blocks[index]):
                row.stop_reason = "eos"
            elif len(row.generated) >= row.request.max_new_tokens:
                row.stop_reason = "max_new_tokens"
            elif (row.request.max_denoising_steps is not None and
                  row.denoise_steps >= row.request.max_denoising_steps):
                row.stop_reason = "max_denoising_steps"
            else:
                bound = row.request.max_new_tokens * int(generation.max_ponder_steps)
                if row.denoise_steps >= bound:
                    row.stop_reason = "episode_watchdog"
            if row.stop_reason is not None:
                finished.append(row)
        return finished

    def result(self, row: InferenceRow) -> dict[str, Any]:
        now = time.perf_counter()
        return {"event": "generation_result", "request_id": row.request.request_id,
                "status": "finished", "checkpoint": str(self.runtime.checkpoint),
                "step": self.runtime.step, "thinking": "on" if row.request.think else "off",
                "stop_reason": row.stop_reason, "prompt_tokens": len(row.prompt_ids),
                "generated_tokens": len(row.generated), "token_ids": row.generated,
                "text": self.runtime.tokenizer.decode(row.generated, skip_special_tokens=True),
                "denoise_steps": row.denoise_steps, "jump_count": row.jumps,
                "state_shift_count": row.shifts,
                "tokens_per_forward": len(row.generated) / max(row.denoise_steps, 1),
                "queue_seconds": row.admitted - row.created,
                "prefill_seconds": row.prefill_seconds,
                "ttfd_seconds": None if row.first_denoise_at is None else row.first_denoise_at - row.created,
                "denoise_steps_per_second": row.denoise_steps / max(
                    (row.last_denoise_at or now) - row.admitted, 1e-9),
                "ttft_seconds": None if row.first_token_at is None else row.first_token_at - row.created,
                "elapsed_seconds": now - row.created,
                "prefix_cache_hits": self.prefix.hits,
                "prefix_cache_reused_tokens": self.prefix.reused_tokens,
                "mlx_active_memory_bytes": int(mx.get_active_memory()),
                "mlx_peak_memory_bytes": int(mx.get_peak_memory())}


def iter_scheduler(engine: MLXContinuousEngine, incoming: queue.Queue[Any],
                   max_queue_size: int, *,
                   continue_on_error: bool = False) -> Iterator[dict[str, Any]]:
    """Bounded continuous admission with a prefill token budget per decode turn."""
    active: deque[InferenceRow] = deque()
    prefilling: deque[PrefillRow] = deque()
    pending: deque[tuple[ContinuousRequest, float]] = deque()
    seen_ids: set[str] = set()
    input_done = False
    started = time.perf_counter()
    capacity = min(engine.max_batch_rows, engine.max_batch_tokens //
                   int(engine.runtime.config.canvas_length))

    def accept(item: Any, created: float) -> dict[str, Any] | None:
        nonlocal input_done
        if item is None:
            input_done = True
        elif isinstance(item, dict):
            return item
        elif item.request_id in seen_ids:
            return {"event": "request_error", "request_id": item.request_id,
                    "error": "Duplicate request_id."}
        else:
            seen_ids.add(item.request_id)
            pending.append((item, created))

    while not input_done or pending or prefilling or active:
        while not input_done and len(pending) < max_queue_size:
            try:
                error = accept(*incoming.get_nowait())
                if error is not None:
                    yield error
            except queue.Empty:
                break
        prefill_budget = engine.prefill_chunk
        had_active = bool(active)
        while prefill_budget > 0:
            if pending and len(active) + len(prefilling) < engine.max_batch_rows:
                request, created = pending.popleft()
                try:
                    # A short newcomer gets one chunk promptly; incomplete
                    # prompts then rotate in the bounded prefill cohort.
                    prefilling.appendleft(engine.begin_prefill(request, created))
                except Exception as error:
                    yield {"event": "request_error", "request_id": request.request_id,
                           "error": str(error)}
                    continue
            if not prefilling:
                break
            work = prefilling.popleft()
            tokens = min(engine.prefill_chunk, len(work.prompt_ids) - work.offset)
            if tokens > prefill_budget:
                # Keep original chunk boundaries and singleton GEMM shapes.
                prefilling.appendleft(work)
                break
            prefill_budget -= tokens
            try:
                row = engine.prefill_step(work)
                if row is None:
                    prefilling.append(work)
                else:
                    active.append(row)
            except Exception as error:
                yield {"event": "request_error", "request_id": work.request.request_id,
                       "error": str(error)}
            if not had_active and active:
                # Deliver the first request's initial denoise promptly. Once
                # decoding, use the remaining token budget to fill short rows.
                break
        if active:
            engine.peak_active_requests = max(getattr(engine, "peak_active_requests", 0),
                                               len(active) + len(prefilling))
            selected = [active.popleft() for _ in range(min(capacity, len(active)))]
            try:
                finished = engine.step(selected)
                finished_ids = {id(row) for row in finished}
                for row in selected:
                    if id(row) in finished_ids:
                        yield engine.result(row)
                    else:
                        active.append(row)
            except Exception as error:
                for row in selected:
                    yield {"event": "generation_result", "request_id": row.request.request_id,
                           "status": "failed", "error": str(error)}
                if not continue_on_error:
                    raise
        elif not pending and not prefilling and not input_done:
            # Wake immediately for new input instead of polling every 10 ms.
            error = accept(*incoming.get())
            if error is not None:
                yield error
    tail_started = time.perf_counter()
    mx.synchronize()
    engine.denoise_seconds += time.perf_counter() - tail_started
    elapsed = time.perf_counter() - started
    yield {"event": "inference_summary", "elapsed_seconds": elapsed,
                 "active_row_denoises": engine.active_row_steps,
                 "heavy_forward_count": engine.forward_count,
                 "denoise_steps_per_second": engine.active_row_steps / max(elapsed, 1e-9),
                 "denoise_scheduler_seconds": engine.denoise_seconds,
                 "scheduler_denoises_per_second": engine.active_row_steps /
                     max(engine.denoise_seconds, 1e-9),
                 "pipeline_depth": engine.pipeline_depth,
                 "max_inflight_denoises": engine.max_inflight,
                 "peak_active_requests": getattr(engine, "peak_active_requests", 0),
                 "mlx_peak_memory_bytes": int(mx.get_peak_memory())}


def _run_scheduler(engine: MLXContinuousEngine, incoming: queue.Queue[Any],
                   max_queue_size: int) -> int:
    had_error = False
    for record in iter_scheduler(engine, incoming, max_queue_size):
        had_error |= record["event"] == "request_error" or record.get("status") == "failed"
        engine.emit(record)
    return int(had_error)


def _parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description="Modilify Mk2 native MLX text generation")
    parser.add_argument("--model", "--checkpoint", dest="checkpoint",
                        default=(str(Path(__file__).resolve().parent)
                                 if (Path(__file__).resolve().parent / "export_manifest.json").is_file()
                                 else "Modilify/Modilify-Mk2-preview-mlx"),
                        help="Local model directory or Hugging Face Hub model ID.")
    parser.add_argument("--prompt", default="Why is the sky blue?")
    parser.add_argument("--requests-jsonl", metavar="PATH|-",
                        help="Read JSONL requests continuously; '-' reads stdin.")
    parser.add_argument("--stream", action="store_true",
                        help="Stream generated text to stdout for a single --prompt request.")
    parser.add_argument("--think", type=parse_bool, default=True)
    parser.add_argument("--seed", type=int, default=42)
    parser.add_argument("--max-new-tokens", type=int, default=8192)
    parser.add_argument("--max-denoising-steps", type=int)
    parser.add_argument("--canvas-length", type=int)
    parser.add_argument("--repetition-penalty", type=float, default=1.0)
    parser.add_argument("--batch-size", type=int, default=4)
    parser.add_argument("--pipeline-depth", type=int, default=2,
                        help="Bound in-flight singleton denoises; 1 minimizes activation memory.")
    parser.add_argument("--max-batch-tokens", type=int, default=1024)
    parser.add_argument("--prefill-chunk-size", type=int, default=256)
    parser.add_argument("--vocab-chunk-size", type=int, default=4096)
    parser.add_argument("--prefix-cache-mib", type=int, default=512)
    parser.add_argument("--max-queue-size", type=int, default=128)
    parser.add_argument("--commit-failure-budget", type=float, default=None,
                        help="Override commit failure budget (defaults to checkpoint config).")
    parser.add_argument("--commit-top-k", type=int, default=None,
                        help="Override commit top_k bound (defaults to checkpoint config).")
    parser.add_argument("--commit-min-p", type=float, default=None,
                        help="Override commit min_p bound (defaults to checkpoint config).")
    parser.add_argument("--commit-target-confidence", type=float, default=None,
                        help="Override commit target confidence (defaults to checkpoint config).")
    return parser


def _validate_args(parser: argparse.ArgumentParser, args: Any) -> None:
    if args.stream and args.requests_jsonl is not None:
        parser.error("--stream can only be used with a single --prompt request.")
    positive = ("max_new_tokens", "batch_size", "max_batch_tokens",
                "prefill_chunk_size", "vocab_chunk_size", "max_queue_size", "pipeline_depth")
    for name in positive:
        if getattr(args, name) <= 0:
            parser.error(f"--{name.replace('_', '-')} must be positive.")
    if args.max_denoising_steps is not None and args.max_denoising_steps <= 0:
        parser.error("--max-denoising-steps must be positive.")
    if args.prefix_cache_mib < 0:
        parser.error("--prefix-cache-mib must be nonnegative.")
    if not math.isfinite(args.repetition_penalty) or args.repetition_penalty <= 0:
        parser.error("--repetition-penalty must be finite and positive.")
    if args.commit_failure_budget is not None and args.commit_failure_budget <= 0:
        parser.error("--commit-failure-budget must be positive.")
    if args.commit_top_k is not None and args.commit_top_k <= 0:
        parser.error("--commit-top-k must be positive.")
    if args.commit_min_p is not None and not (0.0 < args.commit_min_p < 1.0):
        parser.error("--commit-min-p must be in (0, 1).")
    if args.commit_target_confidence is not None and not (0.0 < args.commit_target_confidence < 1.0):
        parser.error("--commit-target-confidence must be in (0, 1).")


def _producer(path: str, output: queue.Queue[Any], args: Any) -> None:
    stream = None
    try:
        stream = sys.stdin if path == "-" else open(path, encoding="utf-8")
        for line_number, line in enumerate(stream, 1):
            if not line.strip():
                continue
            try:
                record = json.loads(line)
                request = parse_continuous_request(
                    record, default_max_new_tokens=args.max_new_tokens,
                    default_max_denoising_steps=args.max_denoising_steps,
                    default_seed=args.seed, default_think=args.think,
                )
                output.put((request, time.perf_counter()))
            except Exception as error:
                output.put(({"event": "request_error", "line": line_number,
                             "error": str(error)}, time.perf_counter()))
    except Exception as error:
        output.put(({"event": "request_error", "error": str(error)},
                    time.perf_counter()))
    finally:
        if stream is not None and stream is not sys.stdin:
            stream.close()
        output.put((None, time.perf_counter()))


def main(argv: list[str] | None = None) -> int:
    parser = _parser()
    args = parser.parse_args(argv)
    _validate_args(parser, args)
    runtime = load_runtime(args.checkpoint, canvas_length=args.canvas_length,
                           max_new_tokens=args.max_new_tokens,
                           max_denoising_steps=args.max_denoising_steps,
                           repetition_penalty=args.repetition_penalty,
                           commit_failure_budget=args.commit_failure_budget,
                           commit_top_k=args.commit_top_k,
                           commit_min_p=args.commit_min_p,
                           commit_target_confidence=args.commit_target_confidence)
    if args.max_denoising_steps is None:
        args.max_denoising_steps = runtime.generation.max_denoising_steps

    streamed_token_ids: dict[str, list[int]] = {}
    streamed_text: dict[str, str] = {}

    def emit(record: dict[str, Any]) -> None:
        if args.stream:
            event = record.get("event")
            request_id = str(record.get("request_id", "single"))
            if event == "token":
                ids = streamed_token_ids.setdefault(request_id, [])
                ids.extend(int(token_id) for token_id in record["token_ids"])
                current = runtime.tokenizer.decode(ids, skip_special_tokens=True)
                previous = streamed_text.get(request_id, "")
                if current.startswith(previous):
                    delta = current[len(previous):]
                else:
                    # Keep output append-only if a tokenizer revises a prior decode.
                    common = 0
                    for old_char, new_char in zip(previous, current):
                        if old_char != new_char:
                            break
                        common += 1
                    delta = current[common:]
                if delta:
                    sys.stdout.write(delta)
                    sys.stdout.flush()
                streamed_text[request_id] = current
            elif event == "generation_result":
                if record.get("status") == "failed":
                    sys.stderr.write(f"\nInference failed: {record.get('error', 'unknown error')}\n")
                    sys.stderr.flush()
                else:
                    sys.stdout.write("\n")
                    sys.stdout.flush()
            elif event == "request_error":
                sys.stderr.write(f"\nInference error: {record.get('error', 'unknown error')}\n")
                sys.stderr.flush()
            return
        sys.stdout.write(json.dumps(record, ensure_ascii=False) + "\n")
        sys.stdout.flush()

    emit({"event": "runtime_loaded", "checkpoint": str(runtime.checkpoint),
          "step": runtime.step, "restored_trainable_tensors": runtime.tensor_count,
          "backend": "mlx", "canvas_length": runtime.config.canvas_length,
          "pipeline_depth": args.pipeline_depth,
          "execution": "layer_interleaved_singleton",
          "commit_failure_budget": runtime.config.commit_failure_budget,
          "commit_target_confidence": getattr(runtime.config, "commit_target_confidence", None),
          "commit_top_k": getattr(runtime.config, "commit_top_k", None),
          "commit_min_p": getattr(runtime.config, "commit_min_p", None)})
    engine = MLXContinuousEngine(
        runtime, prefix_mib=args.prefix_cache_mib,
        prefill_chunk=args.prefill_chunk_size,
        vocab_chunk=args.vocab_chunk_size,
        max_batch_rows=args.batch_size,
        max_batch_tokens=args.max_batch_tokens, emit=emit, pipeline_depth=args.pipeline_depth,
    )
    incoming: queue.Queue[Any] = queue.Queue(maxsize=max(2, args.max_queue_size))
    if args.requests_jsonl is None:
        request = ContinuousRequest("single", [{"role": "user", "content": args.prompt}],
                                    args.max_new_tokens, args.max_denoising_steps,
                                    args.seed, args.think, args.prompt)
        incoming.put((request, time.perf_counter()))
        incoming.put((None, time.perf_counter()))
    else:
        threading.Thread(target=_producer, args=(args.requests_jsonl, incoming, args),
                         daemon=True).start()
    return _run_scheduler(engine, incoming, args.max_queue_size)


def load_model(
    model: str = "Modilify/Modilify-Mk2-preview-mlx", *,
    canvas_length: int | None = None,
) -> MLXRuntime:
    """Load the complete local model or its Hugging Face snapshot once."""
    return load_runtime(model, canvas_length=canvas_length,
                        max_new_tokens=256, max_denoising_steps=None,
                        repetition_penalty=1.0)


def generate(
    runtime: MLXRuntime, prompt: str | None = None, *,
    messages: list[dict[str, Any]] | None = None,
    max_new_tokens: int = 256, max_denoising_steps: int | None = None,
    seed: int = 42, think: bool = True,
) -> dict[str, Any]:
    """Generate one independent response, returning text, tokens, and metrics."""
    record = {"request_id": "single", "max_new_tokens": max_new_tokens,
              "max_denoising_steps": max_denoising_steps, "seed": seed, "think": think}
    if prompt is not None:
        record["prompt"] = prompt
    if messages is not None:
        record["messages"] = messages
    request = parse_continuous_request(
        record, default_max_new_tokens=256,
        default_max_denoising_steps=runtime.generation.max_denoising_steps,
        default_seed=42, default_think=True,
    )
    engine = MLXContinuousEngine(
        runtime, prefix_mib=0, prefill_chunk=256, vocab_chunk=4096,
        max_batch_rows=1, max_batch_tokens=int(runtime.config.canvas_length),
        emit=lambda event: None, pipeline_depth=1, token_events=False,
    )
    row = engine.admit(request, time.perf_counter())
    while not engine.step([row]):
        pass
    return engine.result(row)


if __name__ == "__main__":
    raise SystemExit(main())