ydy9038074 commited on
Commit
e4f7326
·
verified ·
1 Parent(s): 51d7858

Publish Modilify Mk2 Preview MLX

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +2 -34
  2. .gitignore +11 -0
  3. LICENSE +125 -0
  4. NOTICE.md +255 -0
  5. README.md +126 -0
  6. assets/01-LOGO.jpg +3 -0
  7. chat_template.jinja +387 -0
  8. config.json +134 -0
  9. export_manifest.json +0 -0
  10. generation_config.json +10 -0
  11. inference.py +1186 -0
  12. model-00001.safetensors +3 -0
  13. model-00002.safetensors +3 -0
  14. model-00003.safetensors +3 -0
  15. model-00004.safetensors +3 -0
  16. model-00005.safetensors +3 -0
  17. model-00006.safetensors +3 -0
  18. model-00007.safetensors +3 -0
  19. model-00008.safetensors +3 -0
  20. model-00009.safetensors +3 -0
  21. model-00010.safetensors +3 -0
  22. model-00011.safetensors +3 -0
  23. model-00012.safetensors +3 -0
  24. model-00013.safetensors +3 -0
  25. model-00014.safetensors +3 -0
  26. model-00015.safetensors +3 -0
  27. model-00016.safetensors +3 -0
  28. model-00017.safetensors +3 -0
  29. model-00018.safetensors +3 -0
  30. model-00019.safetensors +3 -0
  31. model-00020.safetensors +3 -0
  32. model-00021.safetensors +3 -0
  33. model-00022.safetensors +3 -0
  34. model-00023.safetensors +3 -0
  35. model-00024.safetensors +3 -0
  36. model-00025.safetensors +3 -0
  37. model-00026.safetensors +3 -0
  38. model-00027.safetensors +3 -0
  39. model-00028.safetensors +3 -0
  40. model-00029.safetensors +3 -0
  41. model-00030.safetensors +3 -0
  42. model-00031.safetensors +3 -0
  43. model-00032.safetensors +3 -0
  44. model.safetensors.index.json +0 -0
  45. modilify_mk2/__init__.py +9 -0
  46. modilify_mk2/chat.py +287 -0
  47. modilify_mk2/configuration_modilify_mk2.py +442 -0
  48. modilify_mk2/mlx_commit_policy.py +255 -0
  49. modilify_mk2/mlx_gdn2_memory.py +117 -0
  50. modilify_mk2/mlx_gdn2_trajectory.py +87 -0
.gitattributes CHANGED
@@ -1,35 +1,3 @@
1
- *.7z filter=lfs diff=lfs merge=lfs -text
2
- *.arrow filter=lfs diff=lfs merge=lfs -text
3
- *.bin filter=lfs diff=lfs merge=lfs -text
4
- *.bz2 filter=lfs diff=lfs merge=lfs -text
5
- *.ckpt filter=lfs diff=lfs merge=lfs -text
6
- *.ftz filter=lfs diff=lfs merge=lfs -text
7
- *.gz filter=lfs diff=lfs merge=lfs -text
8
- *.h5 filter=lfs diff=lfs merge=lfs -text
9
- *.joblib filter=lfs diff=lfs merge=lfs -text
10
- *.lfs.* filter=lfs diff=lfs merge=lfs -text
11
- *.mlmodel filter=lfs diff=lfs merge=lfs -text
12
- *.model filter=lfs diff=lfs merge=lfs -text
13
- *.msgpack filter=lfs diff=lfs merge=lfs -text
14
- *.npy filter=lfs diff=lfs merge=lfs -text
15
- *.npz filter=lfs diff=lfs merge=lfs -text
16
- *.onnx filter=lfs diff=lfs merge=lfs -text
17
- *.ot filter=lfs diff=lfs merge=lfs -text
18
- *.parquet filter=lfs diff=lfs merge=lfs -text
19
- *.pb filter=lfs diff=lfs merge=lfs -text
20
- *.pickle filter=lfs diff=lfs merge=lfs -text
21
- *.pkl filter=lfs diff=lfs merge=lfs -text
22
- *.pt filter=lfs diff=lfs merge=lfs -text
23
- *.pth filter=lfs diff=lfs merge=lfs -text
24
- *.rar filter=lfs diff=lfs merge=lfs -text
25
  *.safetensors filter=lfs diff=lfs merge=lfs -text
26
- saved_model/**/* filter=lfs diff=lfs merge=lfs -text
27
- *.tar.* filter=lfs diff=lfs merge=lfs -text
28
- *.tar filter=lfs diff=lfs merge=lfs -text
29
- *.tflite filter=lfs diff=lfs merge=lfs -text
30
- *.tgz filter=lfs diff=lfs merge=lfs -text
31
- *.wasm filter=lfs diff=lfs merge=lfs -text
32
- *.xz filter=lfs diff=lfs merge=lfs -text
33
- *.zip filter=lfs diff=lfs merge=lfs -text
34
- *.zst filter=lfs diff=lfs merge=lfs -text
35
- *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  *.safetensors filter=lfs diff=lfs merge=lfs -text
2
+ tokenizer.json filter=lfs diff=lfs merge=lfs -text
3
+ assets/01-LOGO.jpg filter=lfs diff=lfs merge=lfs -text
 
 
 
 
 
 
 
 
.gitignore ADDED
@@ -0,0 +1,11 @@
 
 
 
 
 
 
 
 
 
 
 
 
1
+ .DS_Store
2
+ __pycache__/
3
+ .vscode/
4
+ *.py[cod]
5
+ .venv/
6
+ build/
7
+ dist/
8
+ *.egg-info/
9
+ .cache/
10
+ .env
11
+ /.upload-fast.sh
LICENSE ADDED
@@ -0,0 +1,125 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ Modilify Open Model License 1.0
2
+
3
+ Copyright 2026 Modilify
4
+
5
+ 1. Definitions
6
+
7
+ "Model" means the weights, configuration, inference code, tokenizer, processor,
8
+ and documentation distributed with this license. "Derivative Model" means a
9
+ modified, fine-tuned, distilled, merged, quantized, or otherwise adapted version
10
+ of the Model. "You" means the individual or legal entity exercising permissions
11
+ under this license. "High-Risk Use" means a use that can materially affect a
12
+ person's safety, liberty, access to essential services, employment, housing,
13
+ credit, education, legal rights, or medical care, or that controls critical
14
+ infrastructure, weapons, or large-scale biometric surveillance.
15
+
16
+ 2. Copyright Grant
17
+
18
+ Subject to this license, Modilify grants You a worldwide, perpetual,
19
+ non-exclusive, royalty-free, irrevocable copyright license to use, reproduce,
20
+ prepare derivative works of, publicly display, publicly perform, sublicense,
21
+ host as a service, and distribute the Model and Derivative Models, including for
22
+ commercial purposes.
23
+
24
+ 3. Patent Grant
25
+
26
+ Each contributor grants You a worldwide, perpetual, non-exclusive, royalty-free,
27
+ irrevocable patent license, except as stated in this section, to make, have made,
28
+ use, offer to sell, sell, import, and otherwise transfer the Model where the
29
+ license applies only to those patent claims licensable by that contributor that
30
+ are necessarily infringed by that contributor's contribution alone or in
31
+ combination with the Model. If You institute patent litigation alleging that the
32
+ Model or a contribution constitutes patent infringement, patent licenses granted
33
+ to You under this license terminate as of the filing date.
34
+
35
+ 4. Conditions on Redistribution
36
+
37
+ If You distribute the Model or a Derivative Model, You must:
38
+
39
+ a. provide recipients a copy of this license;
40
+ b. retain copyright, patent, attribution, and NOTICE statements;
41
+ c. state clearly that You modified the Model and identify material modifications;
42
+ d. preserve applicable third-party license and attribution notices; and
43
+ e. publish with the distributed model a reasonably accessible impact statement
44
+ describing intended uses, material limitations, evaluation scope, known
45
+ safety risks, and risk mitigations for the Derivative Model.
46
+
47
+ The impact statement may be maintained in a public model card or equivalent
48
+ document. You are not required to submit it separately to Modilify.
49
+
50
+ 5. Responsible Use and High-Risk Uses
51
+
52
+ You must not use the Model or a Derivative Model:
53
+
54
+ a. to develop, operate, or materially facilitate weapons, autonomous targeting,
55
+ or systems intended to cause physical harm;
56
+ b. for unlawful mass surveillance, biometric identification without lawful
57
+ authority and appropriate safeguards, or social scoring that determines
58
+ access to rights or essential services;
59
+ c. to exploit children or vulnerable persons, facilitate human trafficking, or
60
+ generate non-consensual intimate content;
61
+ d. to impersonate a person or deceptively represent machine output as an
62
+ authentic human communication where the deception is reasonably likely to
63
+ cause material harm; or
64
+ e. to make a final decision in a High-Risk Use without meaningful qualified
65
+ human review, proportionate testing, monitoring, appeal or correction paths,
66
+ and compliance with applicable law.
67
+
68
+ Before deploying the Model in a High-Risk Use, You must perform safety and impact
69
+ due diligence proportionate to foreseeable harm. At minimum, document the use
70
+ context, evaluate relevant failure modes and affected groups, apply reasonable
71
+ technical and organizational safeguards, monitor material incidents, and update
72
+ or suspend the deployment when its residual risk is not reasonable. Research,
73
+ testing, auditing, and defensive safety work are permitted when conducted with
74
+ appropriate safeguards.
75
+
76
+ 6. Trademarks
77
+
78
+ This license does not grant permission to use the trade names, trademarks,
79
+ service marks, or product names of Modilify or any contributor, except as needed
80
+ for reasonable and customary attribution or to describe the origin of the Model.
81
+
82
+ 7. Third-Party Components
83
+
84
+ The Model includes or derives from third-party components identified in
85
+ NOTICE.md. Those components remain subject to their applicable
86
+ licenses and terms. In particular, rights and obligations associated with the
87
+ Google DiffusionGemma base are not removed, narrowed, or replaced by this
88
+ license. You are responsible for complying with all applicable upstream terms.
89
+
90
+ 8. Termination and Reinstatement
91
+
92
+ Your rights terminate automatically if You materially violate this license and
93
+ do not cure the violation within 30 days after becoming aware of it. Rights are
94
+ reinstated upon timely cure unless a rights holder provides written notice of a
95
+ substantially similar repeated violation. Sections intended by their nature to
96
+ survive termination remain effective.
97
+
98
+ 9. Disclaimer of Warranty
99
+
100
+ THE MODEL IS PROVIDED "AS IS," WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND,
101
+ EXPRESS OR IMPLIED, INCLUDING WARRANTIES OF TITLE, NON-INFRINGEMENT,
102
+ MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE, ACCURACY, OR SAFETY. YOU ARE
103
+ SOLELY RESPONSIBLE FOR DETERMINING THE APPROPRIATENESS OF USING OR REDISTRIBUTING
104
+ THE MODEL AND ASSUME ALL RISKS ASSOCIATED WITH YOUR EXERCISE OF PERMISSIONS.
105
+
106
+ 10. Limitation of Liability
107
+
108
+ TO THE MAXIMUM EXTENT PERMITTED BY LAW, NO COPYRIGHT HOLDER OR CONTRIBUTOR SHALL
109
+ BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR
110
+ CONSEQUENTIAL DAMAGES ARISING FROM THIS LICENSE OR THE USE OR INABILITY TO USE
111
+ THE MODEL, HOWEVER CAUSED AND UNDER ANY THEORY OF LIABILITY, EVEN IF ADVISED OF
112
+ THE POSSIBILITY OF SUCH DAMAGES.
113
+
114
+ 11. Governing Law and Venue
115
+
116
+ This license is governed by the laws of the State of California, excluding its
117
+ conflict-of-law rules. Any dispute arising from this license must be brought in
118
+ the state or federal courts located in Santa Clara County, California, and each
119
+ party consents to their personal jurisdiction and venue.
120
+
121
+ 12. Entire License; Severability
122
+
123
+ This document states the complete Modilify license for the Model, subject to
124
+ applicable third-party terms. If a provision is unenforceable, it will be limited
125
+ to the minimum extent necessary and the remaining provisions remain effective.
NOTICE.md ADDED
@@ -0,0 +1,255 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Modilify Mk2 — Notices
2
+
3
+ Copyright 2026 Modilify
4
+
5
+ This distribution is derived from `google/diffusiongemma-26B-A4B-it`, published by Google DeepMind under Apache License 2.0. This MLX distribution retains the text encoder/decoder backbone, tokenizer, and chat template. It does not include vision tower, vision projection, or image/video inference path. Modilify added trained LoRA adapters, dual-timescale GDN2 trajectory memory, working and persistent memory readers, and the confidence-and-entropy prefix commit policy. The native MLX implementation uses the DiffusionGemma text architecture from mlx-vlm and unfused adapters from mlx-lm.
6
+
7
+ The Modilify Open Model License 1.0 applies to Modilify's distribution and original contributions. It does not erase, narrow, or replace rights and notices applicable to upstream components. Users remain responsible for complying with all applicable upstream terms.
8
+
9
+ - Upstream model: https://huggingface.co/google/diffusiongemma-26B-A4B-it
10
+ - Transformers project: https://github.com/huggingface/transformers
11
+ - MLX project: https://github.com/ml-explore/mlx
12
+ - mlx-vlm project: https://github.com/Blaizzy/mlx-vlm
13
+ - mlx-lm project: https://github.com/ml-explore/mlx-lm
14
+
15
+ The configuration subclasses public DiffusionGemma interfaces in Hugging Face Transformers, distributed under Apache License 2.0. MLX and mlx-lm are distributed under the MIT License, and mlx-vlm under the MIT License. These libraries are installed as dependencies; their original notices remain in their distributions.
16
+
17
+ ## Derivative Model Impact Statement Template
18
+
19
+ When distributing a derivative of Modilify Mk2, include a public impact statement covering the following items. No separate submission to Modilify is required.
20
+
21
+ ### Identity and modifications
22
+
23
+ - Model name, version, publisher, and contact.
24
+ - Base version.
25
+ - Material modifications, data sources, merges, quantization, or adaptation.
26
+
27
+ ### Intended and excluded uses
28
+
29
+ - Intended users and use cases.
30
+ - Explicitly excluded uses.
31
+ - Deployment context and degree of human oversight.
32
+
33
+ ### Evaluation scope
34
+
35
+ - Evaluated capabilities and datasets.
36
+ - Languages, modalities, populations, or contexts not evaluated.
37
+ - Hardware and software used.
38
+
39
+ ### Known limitations and foreseeable risks
40
+
41
+ - Reliability limitations.
42
+ - Safety, bias, privacy, security, and misuse risks.
43
+ - High-risk decisions the model must not make autonomously.
44
+
45
+ ### Mitigations and monitoring
46
+
47
+ - Technical and organizational safeguards.
48
+ - Human review, appeal, and correction mechanisms.
49
+ - Monitoring, incident response, and update policy.
50
+
51
+ ## Apache License 2.0
52
+
53
+ The complete license text applicable to the upstream components follows.
54
+
55
+ Apache License
56
+ Version 2.0, January 2004
57
+ http://www.apache.org/licenses/
58
+
59
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
60
+
61
+ 1. Definitions.
62
+
63
+ "License" shall mean the terms and conditions for use, reproduction,
64
+ and distribution as defined by Sections 1 through 9 of this document.
65
+
66
+ "Licensor" shall mean the copyright owner or entity authorized by
67
+ the copyright owner that is granting the License.
68
+
69
+ "Legal Entity" shall mean the union of the acting entity and all
70
+ other entities that control, are controlled by, or are under common
71
+ control with that entity. For the purposes of this definition,
72
+ "control" means (i) the power, direct or indirect, to cause the
73
+ direction or management of such entity, whether by contract or
74
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
75
+ outstanding shares, or (iii) beneficial ownership of such entity.
76
+
77
+ "You" (or "Your") shall mean an individual or Legal Entity
78
+ exercising permissions granted by this License.
79
+
80
+ "Source" form shall mean the preferred form for making modifications,
81
+ including but not limited to software source code, documentation
82
+ source, and configuration files.
83
+
84
+ "Object" form shall mean any form resulting from mechanical
85
+ transformation or translation of a Source form, including but
86
+ not limited to compiled object code, generated documentation,
87
+ and conversions to other media types.
88
+
89
+ "Work" shall mean the work of authorship, whether in Source or
90
+ Object form, made available under the License, as indicated by a
91
+ copyright notice that is included in or attached to the work
92
+ (an example is provided in the Appendix below).
93
+
94
+ "Derivative Works" shall mean any work, whether in Source or Object
95
+ form, that is based on (or derived from) the Work and for which the
96
+ editorial revisions, annotations, elaborations, or other modifications
97
+ represent, as a whole, an original work of authorship. For the purposes
98
+ of this License, Derivative Works shall not include works that remain
99
+ separable from, or merely link (or bind by name) to the interfaces of,
100
+ the Work and Derivative Works thereof.
101
+
102
+ "Contribution" shall mean any work of authorship, including
103
+ the original version of the Work and any modifications or additions
104
+ to that Work or Derivative Works thereof, that is intentionally
105
+ submitted to Licensor for inclusion in the Work by the copyright owner
106
+ or by an individual or Legal Entity authorized to submit on behalf of
107
+ the copyright owner. For the purposes of this definition, "submitted"
108
+ means any form of electronic, verbal, or written communication sent
109
+ to the Licensor or its representatives, including but not limited to
110
+ communication on electronic mailing lists, source code control systems,
111
+ and issue tracking systems that are managed by, or on behalf of, the
112
+ Licensor for the purpose of discussing and improving the Work, but
113
+ excluding communication that is conspicuously marked or otherwise
114
+ designated in writing by the copyright owner as "Not a Contribution."
115
+
116
+ "Contributor" shall mean Licensor and any individual or Legal Entity
117
+ on behalf of whom a Contribution has been received by Licensor and
118
+ subsequently incorporated within the Work.
119
+
120
+ 2. Grant of Copyright License. Subject to the terms and conditions of
121
+ this License, each Contributor hereby grants to You a perpetual,
122
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
123
+ copyright license to reproduce, prepare Derivative Works of,
124
+ publicly display, publicly perform, sublicense, and distribute the
125
+ Work and such Derivative Works in Source or Object form.
126
+
127
+ 3. Grant of Patent License. Subject to the terms and conditions of
128
+ this License, each Contributor hereby grants to You a perpetual,
129
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
130
+ (except as stated in this section) patent license to make, have made,
131
+ use, offer to sell, sell, import, and otherwise transfer the Work,
132
+ where such license applies only to those patent claims licensable
133
+ by such Contributor that are necessarily infringed by their
134
+ Contribution(s) alone or by combination of their Contribution(s)
135
+ with the Work to which such Contribution(s) was submitted. If You
136
+ institute patent litigation against any entity (including a
137
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
138
+ or a Contribution incorporated within the Work constitutes direct
139
+ or contributory patent infringement, then any patent licenses
140
+ granted to You under this License for that Work shall terminate
141
+ as of the date such litigation is filed.
142
+
143
+ 4. Redistribution. You may reproduce and distribute copies of the
144
+ Work or Derivative Works thereof in any medium, with or without
145
+ modifications, and in Source or Object form, provided that You
146
+ meet the following conditions:
147
+
148
+ (a) You must give any other recipients of the Work or
149
+ Derivative Works a copy of this License; and
150
+
151
+ (b) You must cause any modified files to carry prominent notices
152
+ stating that You changed the files; and
153
+
154
+ (c) You must retain, in the Source form of any Derivative Works
155
+ that You distribute, all copyright, patent, trademark, and
156
+ attribution notices from the Source form of the Work,
157
+ excluding those notices that do not pertain to any part of
158
+ the Derivative Works; and
159
+
160
+ (d) If the Work includes a "NOTICE" text file as part of its
161
+ distribution, then any Derivative Works that You distribute must
162
+ include a readable copy of the attribution notices contained
163
+ within such NOTICE file, excluding those notices that do not
164
+ pertain to any part of the Derivative Works, in at least one
165
+ of the following places: within a NOTICE text file distributed
166
+ as part of the Derivative Works; within the Source form or
167
+ documentation, if provided along with the Derivative Works; or,
168
+ within a display generated by the Derivative Works, if and
169
+ wherever such third-party notices normally appear. The contents
170
+ of the NOTICE file are for informational purposes only and
171
+ do not modify the License. You may add Your own attribution
172
+ notices within Derivative Works that You distribute, alongside
173
+ or as an addendum to the NOTICE text from the Work, provided
174
+ that such additional attribution notices cannot be construed
175
+ as modifying the License.
176
+
177
+ You may add Your own copyright statement to Your modifications and
178
+ may provide additional or different license terms and conditions
179
+ for use, reproduction, or distribution of Your modifications, or
180
+ for any such Derivative Works as a whole, provided Your use,
181
+ reproduction, and distribution of the Work otherwise complies with
182
+ the conditions stated in this License.
183
+
184
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
185
+ any Contribution intentionally submitted for inclusion in the Work
186
+ by You to the Licensor shall be under the terms and conditions of
187
+ this License, without any additional terms or conditions.
188
+ Notwithstanding the above, nothing herein shall supersede or modify
189
+ the terms of any separate license agreement you may have executed
190
+ with Licensor regarding such Contributions.
191
+
192
+ 6. Trademarks. This License does not grant permission to use the trade
193
+ names, trademarks, service marks, or product names of the Licensor,
194
+ except as required for reasonable and customary use in describing the
195
+ origin of the Work and reproducing the content of the NOTICE file.
196
+
197
+ 7. Disclaimer of Warranty. Unless required by applicable law or
198
+ agreed to in writing, Licensor provides the Work (and each
199
+ Contributor provides its Contributions) on an "AS IS" BASIS,
200
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
201
+ implied, including, without limitation, any warranties or conditions
202
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
203
+ PARTICULAR PURPOSE. You are solely responsible for determining the
204
+ appropriateness of using or redistributing the Work and assume any
205
+ risks associated with Your exercise of permissions under this License.
206
+
207
+ 8. Limitation of Liability. In no event and under no legal theory,
208
+ whether in tort (including negligence), contract, or otherwise,
209
+ unless required by applicable law (such as deliberate and grossly
210
+ negligent acts) or agreed to in writing, shall any Contributor be
211
+ liable to You for damages, including any direct, indirect, special,
212
+ incidental, or consequential damages of any character arising as a
213
+ result of this License or out of the use or inability to use the
214
+ Work (including but not limited to damages for loss of goodwill,
215
+ work stoppage, computer failure or malfunction, or any and all
216
+ other commercial damages or losses), even if such Contributor
217
+ has been advised of the possibility of such damages.
218
+
219
+ 9. Accepting Warranty or Additional Liability. While redistributing
220
+ the Work or Derivative Works thereof, You may choose to offer,
221
+ and charge a fee for, acceptance of support, warranty, indemnity,
222
+ or other liability obligations and/or rights consistent with this
223
+ License. However, in accepting such obligations, You may act only
224
+ on Your own behalf and on Your sole responsibility, not on behalf
225
+ of any other Contributor, and only if You agree to indemnify,
226
+ defend, and hold each Contributor harmless for any liability
227
+ incurred by, or claims asserted against, such Contributor by reason
228
+ of your accepting any such warranty or additional liability.
229
+
230
+ END OF TERMS AND CONDITIONS
231
+
232
+ APPENDIX: How to apply the Apache License to your work.
233
+
234
+ To apply the Apache License to your work, attach the following
235
+ boilerplate notice, with the fields enclosed by brackets "[]"
236
+ replaced with your own identifying information. (Don't include
237
+ the brackets!) The text should be enclosed in the appropriate
238
+ comment syntax for the file format. We also recommend that a
239
+ file or class name and description of purpose be included on the
240
+ same "printed page" as the copyright notice for easier
241
+ identification within third-party archives.
242
+
243
+ Copyright [yyyy] [name of copyright owner]
244
+
245
+ Licensed under the Apache License, Version 2.0 (the "License");
246
+ you may not use this file except in compliance with the License.
247
+ You may obtain a copy of the License at
248
+
249
+ http://www.apache.org/licenses/LICENSE-2.0
250
+
251
+ Unless required by applicable law or agreed to in writing, software
252
+ distributed under the License is distributed on an "AS IS" BASIS,
253
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
254
+ See the License for the specific language governing permissions and
255
+ limitations under the License.
README.md ADDED
@@ -0,0 +1,126 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: other
3
+ license_name: modilify-open-model-license-1.0
4
+ license_link: LICENSE
5
+ library_name: mlx
6
+ pipeline_tag: text-generation
7
+ tags:
8
+ - mlx
9
+ - diffusion
10
+ - mixture-of-experts
11
+ - custom-code
12
+ - modilify-mk2
13
+ ---
14
+
15
+ ![Modilify](assets/01-LOGO.jpg)
16
+
17
+ # Modilify Mk2 Preview · MLX
18
+
19
+ Modilify Mk2 is a rolling-canvas text diffusion model with learned trajectory memory. This release provides complete, unquantized native MLX weights and an inference package for Apple Silicon Macs. It builds on [Google DiffusionGemma](https://huggingface.co/google/diffusiongemma-26B-A4B-it), adding LoRA adapters, GDN2 memory at two timescales, and a confidence-and-entropy prefix commit policy.
20
+
21
+ The model revises a 256-token canvas over successive denoising passes, commits a variable-length prefix, and carries memory forward as the canvas advances. Per-position and row-level working states update on every denoising pass; persistent memory updates only when tokens commit.
22
+
23
+ This is an experimental **text-only preview**. The complete model is included; no separate base-model download is required. PyTorch is not required.
24
+
25
+ ## Download and run
26
+
27
+ Use Python 3.12 on an Apple Silicon Mac with macOS 26 or newer for the pinned MLX environment. The weights occupy **48.23 GiB**. Inference needs additional memory for KV caches, recurrent states, and computation. See [local comparison](#local-comparison) for measured short-request memory usage.
28
+
29
+ ```bash
30
+ python3.12 -m venv .venv
31
+ source .venv/bin/activate
32
+ python -m pip install "huggingface-hub==1.20.1"
33
+
34
+ hf download Modilify/Modilify-Mk2-preview-mlx --local-dir ./Modilify-Mk2-preview-mlx
35
+ cd Modilify-Mk2-preview-mlx
36
+ python -m pip install .
37
+
38
+ python inference.py --prompt "Why is the sky blue?" \
39
+ --max-new-tokens 256 --max-denoising-steps 256 --stream
40
+ ```
41
+
42
+ The source entrypoint defaults to the model in the same directory. The installed command accepts a local model directory or a Hub model ID:
43
+
44
+ ```bash
45
+ modilify-mlx --model ./ --prompt "你好,请介绍一下自己。" --stream
46
+ modilify-mlx --model Modilify/Modilify-Mk2-preview-mlx \
47
+ --prompt "What is 2 + 3?" --max-new-tokens 128 --max-denoising-steps 128
48
+ ```
49
+
50
+ The Hub ID path downloads the model snapshot into the Hugging Face cache. Later runs reuse the cached files. All inference computation runs locally.
51
+
52
+ Thinking is enabled by default; use `--think False` to disable it. `--max-new-tokens` limits output length; `--max-denoising-steps` limits denoising work. Reaching either budget may truncate a response. The CLI defaults to 8192 output tokens and the generation configuration's 48 denoising steps. Set both budgets explicitly for your workload. `--canvas-length` can reduce the canvas from its maximum of 256. Without `--stream`, the CLI emits JSONL events and a final result with text, token IDs, stop reason, and timing/memory metrics.
53
+
54
+ ## Python API
55
+
56
+ Load once and reuse the runtime for independent requests:
57
+
58
+ ```python
59
+ from inference import load_model, generate
60
+
61
+ runtime = load_model("./")
62
+ result = generate(
63
+ runtime,
64
+ "Explain diffusion models in one paragraph.",
65
+ max_new_tokens=256,
66
+ max_denoising_steps=256,
67
+ seed=42,
68
+ think=True,
69
+ )
70
+ print(result["text"])
71
+ print(result["stop_reason"])
72
+ ```
73
+
74
+ Pass `messages=[{"role": "user", "content": "..."}]` instead of `prompt` for a conversation. The tokenizer's Modilify template handles system, user, and assistant turns. With thinking enabled, generated text may include reasoning before the answer.
75
+
76
+ ## Continuous requests
77
+
78
+ ```bash
79
+ printf '%s\n' \
80
+ '{"request_id":"a","prompt":"你好","max_new_tokens":128,"max_denoising_steps":128}' \
81
+ '{"request_id":"b","prompt":"What is 2 + 3?","seed":42,"think":false,"max_new_tokens":128,"max_denoising_steps":128}' | \
82
+ modilify-mlx --model ./ --requests-jsonl - \
83
+ --batch-size 4 --pipeline-depth 2 --max-batch-tokens 1024
84
+ ```
85
+
86
+ Each line requires a unique `request_id` and either `prompt` or `messages`. Optional fields are `max_new_tokens`, `max_denoising_steps`, `seed`, and `think`. The default scheduler interleaves independent rows and keeps their random streams and states separate. `--batch-size` controls resident requests; `--pipeline-depth` bounds in-flight denoises.
87
+
88
+ ## Release details
89
+
90
+ | Item | Value |
91
+ | --- | --- |
92
+ | Base model | DiffusionGemma 26B-A4B-it text backbone |
93
+ | Saved training step | 1250 |
94
+ | Memory topology | schema25 / `compact_gdn2_v2` |
95
+ | Memory scheme | `dual_timescale_gdn2_trajectory_memory` |
96
+ | Canvas | Up to 256 tokens |
97
+ | Expert routing | 8 of 128 experts per token |
98
+ | Dense / expert LoRA rank | 16 / 8 |
99
+ | Weight precision | BF16, with small trajectory norms and dynamics parameters in FP32 |
100
+ | Proposal constraints | top_k 40, min_p 0.05 |
101
+ | Commit defaults | failure budget 0.2, target confidence 0.5 |
102
+ | Weights | 32 Safetensors shards, 1323 tensors |
103
+
104
+ The complete backbone, adapters, and memory parameters are stored in the native MLX parameter layout. Adapters remain unfused to preserve the reference BF16 computation. Loading validates model topology, tensor shape/dtype, and shard SHA256 hashes. `export_manifest.json` contains the weight index metadata and release provenance; no optimizer state or training data is shipped.
105
+
106
+ This model uses its own rolling-canvas generation loop. Use the included `inference.py`, `modilify-mlx`, or Python API. Generic `mlx_lm.generate`, Transformers `AutoModel.from_pretrained`, and standard autoregressive serving backends do not implement this protocol. Image, audio, and video inputs are not supported by this release.
107
+
108
+ ## Training, evaluation, and limitations
109
+
110
+ The text backbone is adapted with dense and expert LoRA and trained trajectory memory. Adaptation uses response supervision over the valid rolling canvas, with terminal, confidence, and commit-readiness objectives. This release uses the saved step 1250 weights. The adaptation dataset is not distributed and its source composition is not documented in this release; refer to the upstream model card for base-model training information.
111
+
112
+ Local comparisons establish export correctness for the tested cases, not task quality or safety. No capability benchmark score is claimed for this preview, and upstream benchmark scores do not describe this modified model. Long-context quality, broad language coverage, and deployment safety have not been established by those checks.
113
+
114
+ ## Local comparison
115
+
116
+ On 2026-10-06, the package installed in a fresh Python 3.12 environment without PyTorch and ran offline. Two requests matched the strict step 1250 source checkpoint at 1/2/5/10/20 denoising steps, including token IDs, stopping and canvas shifts. At five steps, interleaved B2 also matched independent B1 outputs and all three GDN2 state hashes. Short-request MLX peak allocations were about 50.79 GiB; longer workloads require additional memory.
117
+
118
+ ## Intended use and impact statement
119
+
120
+ This preview is intended for research on text diffusion, trajectory memory, local generation, and reproducible inference experiments. Generated statements can be incorrect, biased, or inappropriate; verify consequential outputs and keep qualified human review for sensitive decisions. The model is not intended to make autonomous high-risk decisions. Do not assume local execution alone provides safety, privacy, or resistance to prompt injection. Deployment should include use-specific testing, appropriate safeguards, and incident monitoring.
121
+
122
+ ## License and attribution
123
+
124
+ Modilify's original contributions and distribution are covered by the [Modilify Open Model License 1.0](LICENSE), a custom license. Upstream components retain their own terms; the DiffusionGemma base is published under Apache 2.0. See [NOTICE.md](NOTICE.md) for attribution and upstream notices.
125
+
126
+ Maintained by **Modilify**. Report model and runtime issues through the [Hub discussions](https://huggingface.co/Modilify/Modilify-Mk2-preview-mlx/discussions).
assets/01-LOGO.jpg ADDED

Git LFS Details

  • SHA256: dbbe5819f47d5a367c9bb6bb448d82afe5bdb881ef8676f4d4c22b2a6b20e261
  • Pointer size: 131 Bytes
  • Size of remote file: 335 kB
chat_template.jinja ADDED
@@ -0,0 +1,387 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {#
2
+ Template: Google Gemma 4 Canonical Chat Template
3
+ Author: Google Gemma Engineering Team
4
+ Published: 2026-07-09
5
+ Context: Fixed tool-calling loops, turn closures, and thinking content-ordering.
6
+ #}
7
+ {%- macro format_parameters(properties, required, filter_keys=false) -%}
8
+ {%- set standard_keys = ['description', 'type', 'properties', 'required', 'nullable'] -%}
9
+ {%- set ns = namespace(found_first=false) -%}
10
+ {%- for key, value in properties | dictsort -%}
11
+ {%- set add_comma = false -%}
12
+ {%- if not filter_keys or key not in standard_keys -%}
13
+ {%- if ns.found_first %},{% endif -%}
14
+ {%- set ns.found_first = true -%}
15
+ {{ key }}:{
16
+ {%- if value['description'] -%}
17
+ description:<|"|>{{ value['description'] }}<|"|>
18
+ {%- set add_comma = true -%}
19
+ {%- endif -%}
20
+ {%- if value['type'] | upper == 'STRING' -%}
21
+ {%- if value['enum'] -%}
22
+ {%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
23
+ enum:{{ format_argument(value['enum']) }}
24
+ {%- endif -%}
25
+ {%- elif value['type'] | upper == 'ARRAY' -%}
26
+ {%- if value['items'] is mapping and value['items'] -%}
27
+ {%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
28
+ items:{
29
+ {%- set ns_items = namespace(found_first=false) -%}
30
+ {%- for item_key, item_value in value['items'] | dictsort -%}
31
+ {%- if item_value is not none -%}
32
+ {%- if ns_items.found_first %},{% endif -%}
33
+ {%- set ns_items.found_first = true -%}
34
+ {%- if item_key == 'properties' -%}
35
+ properties:{
36
+ {%- if item_value is mapping -%}
37
+ {{- format_parameters(item_value, value['items']['required'] | default([])) -}}
38
+ {%- endif -%}
39
+ }
40
+ {%- elif item_key == 'required' -%}
41
+ required:[
42
+ {%- for req_item in item_value -%}
43
+ <|"|>{{- req_item -}}<|"|>
44
+ {%- if not loop.last %},{% endif -%}
45
+ {%- endfor -%}
46
+ ]
47
+ {%- elif item_key == 'type' -%}
48
+ {%- if item_value is string -%}
49
+ type:{{ format_argument(item_value | upper) }}
50
+ {%- else -%}
51
+ type:{{ format_argument(item_value | map('upper') | list) }}
52
+ {%- endif -%}
53
+ {%- else -%}
54
+ {{ item_key }}:{{ format_argument(item_value) }}
55
+ {%- endif -%}
56
+ {%- endif -%}
57
+ {%- endfor -%}
58
+ }
59
+ {%- endif -%}
60
+ {%- endif -%}
61
+ {%- if value['nullable'] %}
62
+ {%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
63
+ nullable:true
64
+ {%- endif -%}
65
+ {%- if value['type'] | upper == 'OBJECT' -%}
66
+ {%- if value['properties'] is defined and value['properties'] is mapping -%}
67
+ {%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
68
+ properties:{
69
+ {{- format_parameters(value['properties'], value['required'] | default([])) -}}
70
+ }
71
+ {%- elif value is mapping -%}
72
+ {%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
73
+ properties:{
74
+ {{- format_parameters(value, value['required'] | default([]), filter_keys=true) -}}
75
+ }
76
+ {%- endif -%}
77
+ {%- if value['required'] -%}
78
+ {%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
79
+ required:[
80
+ {%- for item in value['required'] | default([]) -%}
81
+ <|"|>{{- item -}}<|"|>
82
+ {%- if not loop.last %},{% endif -%}
83
+ {%- endfor -%}
84
+ ]
85
+ {%- endif -%}
86
+ {%- endif -%}
87
+ {%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
88
+ type:<|"|>{{ value['type'] | upper }}<|"|>}
89
+ {%- endif -%}
90
+ {%- endfor -%}
91
+ {%- endmacro -%}
92
+ {%- macro format_function_declaration(tool_data) -%}
93
+ declaration:{{- tool_data['function']['name'] -}}{description:<|"|>{{- tool_data['function']['description'] -}}<|"|>
94
+ {%- set params = tool_data['function']['parameters'] -%}
95
+ {%- if params -%}
96
+ ,parameters:{
97
+ {%- if params['properties'] -%}
98
+ properties:{ {{- format_parameters(params['properties'], params['required']) -}} },
99
+ {%- endif -%}
100
+ {%- if params['required'] -%}
101
+ required:[
102
+ {%- for item in params['required'] -%}
103
+ <|"|>{{- item -}}<|"|>
104
+ {{- ',' if not loop.last -}}
105
+ {%- endfor -%}
106
+ ],
107
+ {%- endif -%}
108
+ {%- if params['type'] -%}
109
+ type:<|"|>{{- params['type'] | upper -}}<|"|>}
110
+ {%- endif -%}
111
+ {%- endif -%}
112
+ {%- if 'response' in tool_data['function'] -%}
113
+ {%- set response_declaration = tool_data['function']['response'] -%}
114
+ ,response:{
115
+ {%- if response_declaration['description'] -%}
116
+ description:<|"|>{{- response_declaration['description'] -}}<|"|>,
117
+ {%- endif -%}
118
+ {%- if response_declaration['type'] | upper == 'OBJECT' -%}
119
+ type:<|"|>{{- response_declaration['type'] | upper -}}<|"|>}
120
+ {%- endif -%}
121
+ {%- endif -%}
122
+ }
123
+ {%- endmacro -%}
124
+ {%- macro format_argument(argument, escape_keys=True) -%}
125
+ {%- if argument is none -%}
126
+ {{- 'null' -}}
127
+ {%- elif argument is string -%}
128
+ {{- '<|"|>' + argument + '<|"|>' -}}
129
+ {%- elif argument is boolean -%}
130
+ {{- 'true' if argument else 'false' -}}
131
+ {%- elif argument is mapping -%}
132
+ {{- '{' -}}
133
+ {%- set ns = namespace(found_first=false) -%}
134
+ {%- for key, value in argument | dictsort -%}
135
+ {%- if ns.found_first %},{% endif -%}
136
+ {%- set ns.found_first = true -%}
137
+ {%- if escape_keys -%}
138
+ {{- '<|"|>' + key + '<|"|>' -}}
139
+ {%- else -%}
140
+ {{- key -}}
141
+ {%- endif -%}
142
+ :{{- format_argument(value, escape_keys=escape_keys) -}}
143
+ {%- endfor -%}
144
+ {{- '}' -}}
145
+ {%- elif argument is sequence -%}
146
+ {{- '[' -}}
147
+ {%- for item in argument -%}
148
+ {{- format_argument(item, escape_keys=escape_keys) -}}
149
+ {%- if not loop.last %},{% endif -%}
150
+ {%- endfor -%}
151
+ {{- ']' -}}
152
+ {%- else -%}
153
+ {{- argument -}}
154
+ {%- endif -%}
155
+ {%- endmacro -%}
156
+ {%- macro strip_thinking(text) -%}
157
+ {%- set ns = namespace(result='') -%}
158
+ {%- for part in text.split('<channel|>') -%}
159
+ {%- if '<|channel>' in part -%}
160
+ {%- set ns.result = ns.result + part.split('<|channel>')[0] -%}
161
+ {%- else -%}
162
+ {%- set ns.result = ns.result + part -%}
163
+ {%- endif -%}
164
+ {%- endfor -%}
165
+ {{- ns.result | trim -}}
166
+ {%- endmacro -%}
167
+
168
+ {%- macro format_tool_response_block(tool_name, response) -%}
169
+ {{- '<|tool_response>' -}}
170
+ {%- if response is mapping -%}
171
+ {{- 'response:' + tool_name + '{' -}}
172
+ {%- for key, value in response | dictsort -%}
173
+ {{- key -}}:{{- format_argument(value, escape_keys=False) -}}
174
+ {%- if not loop.last %},{% endif -%}
175
+ {%- endfor -%}
176
+ {{- '}' -}}
177
+ {%- else -%}
178
+ {{- 'response:' + tool_name + '{value:' + format_argument(response, escape_keys=False) + '}' -}}
179
+ {%- endif -%}
180
+ {{- '<tool_response|>' -}}
181
+ {%- endmacro -%}
182
+
183
+ {#- ===== SETUP ===== -#}
184
+ {%- set ns = namespace(prev_message_type=None, prev_non_tool_role=None) -%}
185
+ {%- set loop_messages = messages -%}
186
+ {%- set enable_thinking = enable_thinking | default(false) -%}
187
+ {%- set preserve_thinking = preserve_thinking | default(false) -%}
188
+ {{- bos_token -}}
189
+ {#- Handle System/Tool Definitions Block -#}
190
+ {%- if enable_thinking or tools or (messages and messages[0]['role'] in ['system', 'developer']) -%}
191
+ {{- '<|turn>system\n' -}}
192
+ {#- Inject Thinking token at the very top of the FIRST system turn -#}
193
+ {%- if enable_thinking -%}
194
+ {{- '<|think|>\n' -}}
195
+ {%- set ns.prev_message_type = 'think' -%}
196
+ {%- endif -%}
197
+ {%- if messages and messages[0]['role'] in ['system', 'developer'] -%}
198
+ {%- if messages[0]['content'] is string -%}
199
+ {{- messages[0]['content'] | trim -}}
200
+ {%- elif messages[0]['content'] is sequence -%}
201
+ {%- for item in messages[0]['content'] -%}
202
+ {{- item['text'] | trim + ' '-}}
203
+ {%- endfor -%}
204
+ {%- endif -%}
205
+ {%- set loop_messages = messages[1:] -%}
206
+ {%- endif -%}
207
+ {%- if tools -%}
208
+ {%- for tool in tools %}
209
+ {{- '<|tool>' -}}
210
+ {{- format_function_declaration(tool) | trim -}}
211
+ {{- '<tool|>' -}}
212
+ {%- endfor %}
213
+ {%- set ns.prev_message_type = 'tool' -%}
214
+ {%- endif -%}
215
+ {{- '<turn|>\n' -}}
216
+ {%- endif %}
217
+
218
+ {#- Pre-scan: find last user message index for reasoning guard -#}
219
+ {%- set ns_turn = namespace(last_user_idx=-1) -%}
220
+ {%- for i in range(loop_messages | length) -%}
221
+ {%- if loop_messages[i]['role'] == 'user' -%}
222
+ {%- set ns_turn.last_user_idx = i -%}
223
+ {%- endif -%}
224
+ {%- endfor -%}
225
+
226
+ {#- Loop through messages -#}
227
+ {%- for message in loop_messages -%}
228
+ {%- if message['role'] != 'tool' -%}
229
+ {%- set ns.prev_message_type = None -%}
230
+ {%- set role = 'model' if message['role'] == 'assistant' else message['role'] -%}
231
+ {#- Detect continuation using tracked state — O(1) instead of O(n) backward scan -#}
232
+ {%- set continue_same_model_turn = (role == 'model' and ns.prev_non_tool_role == 'assistant') -%}
233
+ {%- if not continue_same_model_turn -%}
234
+ {{- '<|turn>' + role + '\n' }}
235
+
236
+ {%- endif -%}
237
+
238
+ {#- Render reasoning/reasoning_content as thinking channel -#}
239
+ {%- set thinking_text = message.get('reasoning') or message.get('reasoning_content') -%}
240
+ {%- set thinking_gate = (loop.index0 > ns_turn.last_user_idx) or (preserve_thinking and message.get('tool_calls')) -%}
241
+ {%- if thinking_text and thinking_gate -%}
242
+ {{- '<|channel>thought\n' + thinking_text + '\n<channel|>' -}}
243
+ {%- endif -%}
244
+
245
+ {%- if message.get('tool_calls') -%}
246
+ {%- for tool_call in message.get('tool_calls') -%}
247
+ {%- set function = tool_call['function'] -%}
248
+ {{- '<|tool_call>call:' + function['name'] + '{' -}}
249
+ {%- if function['arguments'] is mapping -%}
250
+ {%- set ns_args = namespace(found_first=false) -%}
251
+ {%- for key, value in function['arguments'] | dictsort -%}
252
+ {%- if ns_args.found_first %},{% endif -%}
253
+ {%- set ns_args.found_first = true -%}
254
+ {{- key -}}:{{- format_argument(value, escape_keys=False) -}}
255
+ {%- endfor -%}
256
+ {%- elif function['arguments'] is none -%}
257
+ {%- else -%}
258
+ {{- raise_exception(
259
+ "chat_template: tool_calls[].function.arguments must be a "
260
+ "JSON object (mapping), not a string. Deserialize arguments "
261
+ "before passing to the template."
262
+ ) -}}
263
+ {%- endif -%}
264
+ {{- '}<tool_call|>' -}}
265
+ {%- endfor -%}
266
+ {%- set ns.prev_message_type = 'tool_call' -%}
267
+ {%- endif -%}
268
+
269
+ {%- set ns_tr_out = namespace(flag=false) -%}
270
+ {%- if message.get('tool_responses') -%}
271
+ {#- Legacy: tool_responses embedded on the assistant message (Google/Gemma native) -#}
272
+ {%- for tool_response in message.get('tool_responses') -%}
273
+ {{- format_tool_response_block(tool_response['name'] | default('unknown', true), tool_response['response']) -}}
274
+ {%- set ns_tr_out.flag = true -%}
275
+ {%- set ns.prev_message_type = 'tool_response' -%}
276
+ {%- endfor -%}
277
+ {%- elif message.get('tool_calls') -%}
278
+ {#- OpenAI Chat Completions: forward-scan consecutive role:tool messages -#}
279
+ {%- set ns_tool_scan = namespace(stopped=false) -%}
280
+ {%- for k in range(loop.index0 + 1, loop_messages | length) -%}
281
+ {%- if ns_tool_scan.stopped -%}
282
+ {%- elif loop_messages[k]['role'] != 'tool' -%}
283
+ {%- set ns_tool_scan.stopped = true -%}
284
+ {%- else -%}
285
+ {%- set follow = loop_messages[k] -%}
286
+ {#- Resolve tool_call_id to function name -#}
287
+ {%- set ns_tname = namespace(name=follow.get('name') or 'unknown') -%}
288
+ {%- for tc in message.get('tool_calls') -%}
289
+ {%- if tc.get('id') == follow.get('tool_call_id') -%}
290
+ {%- set ns_tname.name = tc['function']['name'] -%}
291
+ {%- endif -%}
292
+ {%- endfor -%}
293
+ {#- Handle content as string or content-parts array -#}
294
+ {%- set tool_body = follow.get('content') -%}
295
+ {%- if tool_body is string -%}
296
+ {{- format_tool_response_block(ns_tname.name, tool_body) -}}
297
+ {%- elif tool_body is sequence and tool_body is not string -%}
298
+ {%- set ns_txt = namespace(s='') -%}
299
+ {%- for part in tool_body -%}
300
+ {%- if part.get('type') == 'text' -%}
301
+ {%- set ns_txt.s = ns_txt.s + (part.get('text') | default('')) -%}
302
+ {%- endif -%}
303
+ {%- endfor -%}
304
+ {{- format_tool_response_block(ns_tname.name, ns_txt.s) -}}
305
+ {%- for part in tool_body -%}
306
+ {%- if part.get('type') in ['image', 'image_url'] -%}
307
+ {{- '<|image|>' -}}
308
+ {%- elif part.get('type') in ['audio', 'input_audio'] -%}
309
+ {{- '<|audio|>' -}}
310
+ {%- elif part.get('type') == 'video' -%}
311
+ {{- '<|video|>' -}}
312
+ {%- endif -%}
313
+ {%- endfor -%}
314
+ {%- else -%}
315
+ {{- format_tool_response_block(ns_tname.name, tool_body) -}}
316
+ {%- endif -%}
317
+ {%- set ns_tr_out.flag = true -%}
318
+ {%- set ns.prev_message_type = 'tool_response' -%}
319
+ {%- endif -%}
320
+ {%- endfor -%}
321
+ {%- endif -%}
322
+
323
+ {%- set captured_content -%}
324
+ {%- if message.get('content') is string -%}
325
+ {%- if role == 'model' -%}
326
+ {{- strip_thinking(message['content']) -}}
327
+ {%- else -%}
328
+ {{- message['content'] | trim -}}
329
+ {%- endif -%}
330
+ {%- elif message.get('content') is sequence -%}
331
+ {%- for item in message['content'] -%}
332
+ {%- if item.get('type') == 'text' -%}
333
+ {%- if role == 'model' -%}
334
+ {{- strip_thinking(item['text']) -}}
335
+ {%- else -%}
336
+ {{- item['text'] | trim -}}
337
+ {%- endif -%}
338
+ {%- elif item.get('type') in ['image', 'image_url'] -%}
339
+ {{- '<|image|>' -}}
340
+ {%- elif item.get('type') in ['audio', 'input_audio'] -%}
341
+ {{- '<|audio|>' -}}
342
+ {%- elif item.get('type') == 'video' -%}
343
+ {{- '<|video|>' -}}
344
+ {%- endif -%}
345
+ {%- endfor -%}
346
+ {%- endif -%}
347
+ {%- endset -%}
348
+
349
+ {{- captured_content -}}
350
+ {%- set has_content = captured_content | trim | length > 0 -%}
351
+
352
+ {#- Forward-scan: find next non-tool message role for continuation detection -#}
353
+ {%- set next_nt = namespace(role=None, found=false) -%}
354
+ {%- for j in range(loop.index0 + 1, loop_messages | length) -%}
355
+ {%- if not next_nt.found -%}
356
+ {%- if loop_messages[j]['role'] != 'tool' -%}
357
+ {%- set next_nt.role = loop_messages[j]['role'] -%}
358
+ {%- set next_nt.found = true -%}
359
+ {%- endif -%}
360
+ {%- endif -%}
361
+ {%- endfor -%}
362
+
363
+ {%- set continues_into_next = (
364
+ role == 'model'
365
+ and next_nt.role == 'assistant'
366
+ and (not message.get('tool_calls') or ns_tr_out.flag)
367
+ ) -%}
368
+
369
+ {%- if ns.prev_message_type == 'tool_call' and not ns_tr_out.flag -%}
370
+ {{- '<|tool_response>' -}}
371
+ {%- elif continues_into_next -%}
372
+ {%- elif not (ns_tr_out.flag and not has_content and not next_nt.found) -%}
373
+ {{- '<turn|>\n' -}}
374
+ {%- endif -%}
375
+ {%- endif -%}
376
+
377
+ {#- Track previous non-tool role for next iteration (avoids O(n) backward scan) -#}
378
+ {%- set ns.prev_non_tool_role = message['role'] -%}
379
+ {%- endfor -%}
380
+
381
+ {%- if add_generation_prompt -%}
382
+ {%- if ns.prev_message_type != 'tool_response' and ns.prev_message_type != 'tool_call' -%}
383
+ {{- '<|turn>model\n' -}}
384
+ {%- elif ns.prev_message_type == 'tool_response' and enable_thinking -%}
385
+ {{- '<|channel>thought\n' -}}
386
+ {%- endif -%}
387
+ {%- endif -%}
config.json ADDED
@@ -0,0 +1,134 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "canvas_length": 256,
3
+ "channel_end_token_id": 101,
4
+ "commit_confidence_power": 1.0,
5
+ "commit_entropy_weight": 1.0,
6
+ "commit_failure_budget": 0.2,
7
+ "commit_gold_alpha": 0.4,
8
+ "commit_gold_weight": 0.2,
9
+ "commit_min_p": 0.05,
10
+ "commit_readiness_beta": 0.02,
11
+ "commit_readiness_budget_margin": 0.02,
12
+ "commit_readiness_failure_weight": 5.0,
13
+ "commit_readiness_loss_weight": 0.5,
14
+ "commit_readiness_objective": "frontier_prefix_budget_v2",
15
+ "commit_readiness_target_tokens": 16,
16
+ "commit_sequence_dim": 1024,
17
+ "commit_sequence_layers": 2,
18
+ "commit_target_confidence": 0.5,
19
+ "commit_top_k": 40,
20
+ "confidence_calibration_loss_weight": 0.01,
21
+ "eos_token_id": 1,
22
+ "initializer_range": 0.02,
23
+ "kv_cache_bucket_size": 128,
24
+ "latent_dim": 2816,
25
+ "latent_ffn_dim": 7168,
26
+ "latent_history_kv_rank": 1024,
27
+ "latent_history_length": 16,
28
+ "latent_history_views": 4,
29
+ "latent_local_attention_window": 128,
30
+ "latent_memory_slots": 256,
31
+ "latent_num_heads": 16,
32
+ "latent_num_layers": 4,
33
+ "latent_persistent_bus_unfreeze_steps": 0,
34
+ "latent_tape_probes": 4,
35
+ "latent_tape_scheme": "gdn2_spatial_probe_v1",
36
+ "latent_working_bus_unfreeze_steps": 0,
37
+ "latent_working_last_block_global": true,
38
+ "memory_architecture": "compact_gdn2_v2",
39
+ "memory_scheme": "dual_timescale_gdn2_trajectory_memory",
40
+ "model_type": "modilify_mk2",
41
+ "persistent_memory_bus": true,
42
+ "persistent_memory_write": "commit_only_transformer",
43
+ "state_schema_version": 25,
44
+ "terminal_stop_loss_weight": 0.01,
45
+ "terminal_stop_target_probability": 0.9,
46
+ "terminal_token_ids": [
47
+ 106,
48
+ 50
49
+ ],
50
+ "text_config": {
51
+ "attention_bias": false,
52
+ "attention_dropout": 0.0,
53
+ "bos_token_id": 2,
54
+ "dtype": "bfloat16",
55
+ "eos_token_id": 1,
56
+ "final_logit_softcapping": 30.0,
57
+ "global_head_dim": 512,
58
+ "head_dim": 256,
59
+ "hidden_activation": "gelu_pytorch_tanh",
60
+ "hidden_size": 2816,
61
+ "initializer_range": 0.02,
62
+ "intermediate_size": 2112,
63
+ "layer_types": [
64
+ "sliding_attention",
65
+ "sliding_attention",
66
+ "sliding_attention",
67
+ "sliding_attention",
68
+ "sliding_attention",
69
+ "full_attention",
70
+ "sliding_attention",
71
+ "sliding_attention",
72
+ "sliding_attention",
73
+ "sliding_attention",
74
+ "sliding_attention",
75
+ "full_attention",
76
+ "sliding_attention",
77
+ "sliding_attention",
78
+ "sliding_attention",
79
+ "sliding_attention",
80
+ "sliding_attention",
81
+ "full_attention",
82
+ "sliding_attention",
83
+ "sliding_attention",
84
+ "sliding_attention",
85
+ "sliding_attention",
86
+ "sliding_attention",
87
+ "full_attention",
88
+ "sliding_attention",
89
+ "sliding_attention",
90
+ "sliding_attention",
91
+ "sliding_attention",
92
+ "sliding_attention",
93
+ "full_attention"
94
+ ],
95
+ "max_position_embeddings": 262144,
96
+ "model_type": "modilify_mk2_text",
97
+ "moe_intermediate_size": 704,
98
+ "num_attention_heads": 16,
99
+ "num_experts": 128,
100
+ "num_global_key_value_heads": 2,
101
+ "num_hidden_layers": 30,
102
+ "num_key_value_heads": 8,
103
+ "pad_token_id": 0,
104
+ "rms_norm_eps": 1e-06,
105
+ "rope_parameters": {
106
+ "full_attention": {
107
+ "partial_rotary_factor": 0.25,
108
+ "rope_theta": 1000000.0,
109
+ "rope_type": "proportional"
110
+ },
111
+ "sliding_attention": {
112
+ "rope_theta": 10000.0,
113
+ "rope_type": "default"
114
+ }
115
+ },
116
+ "sliding_window": 1024,
117
+ "tie_word_embeddings": true,
118
+ "top_k_experts": 8,
119
+ "use_bidirectional_attention": null,
120
+ "vocab_size": 262144
121
+ },
122
+ "tie_word_embeddings": true,
123
+ "token_ce_normalization": "per_sample_exposure_v1",
124
+ "token_ce_supervision": "valid_canvas_v1",
125
+ "token_loss_weight": 1.0,
126
+ "training_bptt_steps": 16,
127
+ "training_prefix_cache": "incremental",
128
+ "training_scheme": "gold_prefix_shared_commit_0_256_committed_ce_calibration_causal_throughput_terminal_sft",
129
+ "transformers_version": "5.14.1",
130
+ "turn_end_token_id": 106,
131
+ "vocab_chunk_size": 32768,
132
+ "working_memory_bus": true,
133
+ "writer_slot_gate": "per_slot"
134
+ }
export_manifest.json ADDED
The diff for this file is too large to render. See raw diff
 
generation_config.json ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "eos_token_id": 1,
3
+ "jump_on_no_progress_after": 12,
4
+ "max_denoising_steps": null,
5
+ "max_ponder_steps": 64,
6
+ "min_trajectory_progress": 0.005,
7
+ "repetition_penalty": 1.0,
8
+ "repetition_penalty_exclude_token_ids": [],
9
+ "turn_end_token_id": 106
10
+ }
inference.py ADDED
@@ -0,0 +1,1186 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Native MLX schema25 inference with bounded KV reuse and continuous admission."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ import copy
7
+ import hashlib
8
+ import json
9
+ import math
10
+ import queue
11
+ import sys
12
+ import threading
13
+ import time
14
+ from collections import OrderedDict, deque
15
+ from dataclasses import dataclass, fields, replace
16
+ from pathlib import Path
17
+ from typing import Any, Iterator
18
+
19
+ import mlx.core as mx
20
+
21
+ from modilify_mk2.configuration_modilify_mk2 import DENOISE_TEMPERATURE
22
+ from modilify_mk2.mlx_commit_policy import (
23
+ fused_commit_failure_rate, infer_commit_reason, select_commit_lengths,
24
+ )
25
+ from modilify_mk2.mlx_model import MLXCanvasOutput, MLXModilifyMk2
26
+ from modilify_mk2.runtime import MLXRuntime, load_runtime
27
+ from modilify_mk2.mlx_state import MLXLatentState, MLXRollingState
28
+ from modilify_mk2.chat import apply_chat_template
29
+ def parse_bool(value: str | bool) -> bool:
30
+ """Parse explicit CLI booleans such as ``--think true``."""
31
+
32
+ if isinstance(value, bool):
33
+ return value
34
+ normalized = value.strip().lower()
35
+ if normalized in {"1", "true", "yes", "on"}:
36
+ return True
37
+ if normalized in {"0", "false", "no", "off"}:
38
+ return False
39
+ raise argparse.ArgumentTypeError("expected true or false")
40
+
41
+
42
+ REQUEST_FIELDS = frozenset(
43
+ {
44
+ "request_id",
45
+ "prompt",
46
+ "messages",
47
+ "max_new_tokens",
48
+ "max_denoising_steps",
49
+ "seed",
50
+ "think",
51
+ }
52
+ )
53
+
54
+
55
+ @dataclass(frozen=True)
56
+ class ContinuousRequest:
57
+ """One validated machine-mode request before chat-template tokenization."""
58
+
59
+ request_id: str
60
+ messages: list[dict[str, Any]]
61
+ max_new_tokens: int
62
+ max_denoising_steps: int | None
63
+ seed: int
64
+ think: bool
65
+ prompt: str | None = None
66
+
67
+
68
+ def stable_request_seed(base_seed: int, request_id: str) -> int:
69
+ """Derive a scheduler-independent non-negative seed from a stable request ID."""
70
+
71
+ payload = f"{base_seed}\0{request_id}".encode("utf-8")
72
+ return int.from_bytes(hashlib.sha256(payload).digest()[:8], "big") & ((1 << 63) - 1)
73
+
74
+
75
+ def _positive_int(record: dict[str, Any], name: str, default: int | None) -> int | None:
76
+ value = record.get(name, default)
77
+ if value is None:
78
+ return default
79
+ if isinstance(value, bool) or not isinstance(value, int) or value <= 0:
80
+ raise ValueError(f"`{name}` must be a positive integer or null.")
81
+ return value
82
+
83
+
84
+ def parse_continuous_request(
85
+ record: Any,
86
+ *,
87
+ default_max_new_tokens: int,
88
+ default_max_denoising_steps: int | None,
89
+ default_seed: int,
90
+ default_think: bool,
91
+ ) -> ContinuousRequest:
92
+ """Validate one strict request object without silently inventing identity."""
93
+
94
+ if not isinstance(record, dict):
95
+ raise ValueError("Each request line must be a JSON object.")
96
+ unknown = sorted(set(record).difference(REQUEST_FIELDS))
97
+ if unknown:
98
+ raise ValueError(f"Unknown request fields: {unknown}")
99
+ request_id = record.get("request_id")
100
+ if not isinstance(request_id, str) or not request_id.strip():
101
+ raise ValueError("`request_id` must be a non-empty string.")
102
+ has_prompt = "prompt" in record
103
+ has_messages = "messages" in record
104
+ if has_prompt == has_messages:
105
+ raise ValueError("Exactly one of `prompt` or `messages` is required.")
106
+
107
+ prompt = None
108
+ if has_prompt:
109
+ prompt = record["prompt"]
110
+ if not isinstance(prompt, str):
111
+ raise ValueError("`prompt` must be a string.")
112
+ messages = [{"role": "user", "content": prompt}]
113
+ else:
114
+ messages = record["messages"]
115
+ if not isinstance(messages, list) or not messages:
116
+ raise ValueError("`messages` must be a non-empty list.")
117
+ if any(
118
+ not isinstance(message, dict)
119
+ or "role" not in message
120
+ or "content" not in message
121
+ for message in messages
122
+ ):
123
+ raise ValueError("Every message must be an object with `role` and `content`.")
124
+
125
+ think = record.get("think", default_think)
126
+ if not isinstance(think, bool):
127
+ raise ValueError("`think` must be a boolean.")
128
+ supplied_seed = record.get("seed")
129
+ if supplied_seed is not None and (
130
+ isinstance(supplied_seed, bool) or not isinstance(supplied_seed, int)
131
+ ):
132
+ raise ValueError("`seed` must be an integer or null.")
133
+ seed = (
134
+ stable_request_seed(default_seed, request_id)
135
+ if supplied_seed is None
136
+ else supplied_seed
137
+ )
138
+ return ContinuousRequest(
139
+ request_id=request_id,
140
+ prompt=prompt,
141
+ messages=messages,
142
+ max_new_tokens=int(
143
+ _positive_int(record, "max_new_tokens", default_max_new_tokens)
144
+ ),
145
+ max_denoising_steps=_positive_int(
146
+ record, "max_denoising_steps", default_max_denoising_steps
147
+ ),
148
+ seed=seed,
149
+ think=think,
150
+ )
151
+
152
+
153
+ def _extract_input_ids(encoded: Any) -> list[int]:
154
+ value = encoded.get("input_ids") if isinstance(encoded, dict) else encoded
155
+ if hasattr(encoded, "input_ids"):
156
+ value = encoded.input_ids
157
+ if isinstance(value, tuple):
158
+ value = list(value)
159
+ if isinstance(value, list) and len(value) == 1 and isinstance(value[0], list):
160
+ value = value[0]
161
+ if not isinstance(value, list) or any(
162
+ isinstance(token, bool) or not isinstance(token, int) for token in value
163
+ ):
164
+ raise ValueError("Chat template did not return one integer token sequence.")
165
+ if not value:
166
+ raise ValueError("Chat template returned an empty prompt.")
167
+ return [int(token) for token in value]
168
+
169
+
170
+ def _seed(value: int, salt: int) -> mx.array:
171
+ return mx.random.key((int(value) + int(salt)) % (2**32))
172
+
173
+
174
+ def _logical(value: mx.array, head: mx.array) -> mx.array:
175
+ canvas = value.shape[1]
176
+ index = (head[:, None] + mx.arange(canvas)[None, :]) % canvas
177
+ return mx.take_along_axis(value, index, axis=1)
178
+
179
+
180
+ def _physical(value: mx.array, head: mx.array) -> mx.array:
181
+ canvas = value.shape[1]
182
+ index = (mx.arange(canvas)[None, :] - head[:, None]) % canvas
183
+ return mx.take_along_axis(value, index, axis=1)
184
+
185
+
186
+ def _concat_rows(values: list[Any]) -> Any:
187
+ first = values[0]
188
+ if isinstance(first, mx.array):
189
+ return mx.concatenate(values, axis=0)
190
+ return type(first)(**{
191
+ item.name: _concat_rows([getattr(value, item.name) for value in values])
192
+ for item in fields(first)
193
+ })
194
+
195
+
196
+ def _slice_row(value: Any, row: int) -> Any:
197
+ if isinstance(value, mx.array):
198
+ return value[row:row + 1]
199
+ return type(value)(**{
200
+ item.name: _slice_row(getattr(value, item.name), row)
201
+ for item in fields(value)
202
+ })
203
+
204
+
205
+ def _arrays(value: Any) -> list[mx.array]:
206
+ if isinstance(value, mx.array):
207
+ return [value]
208
+ return [array for item in fields(value)
209
+ for array in _arrays(getattr(value, item.name))]
210
+
211
+
212
+ def _clone_cache(cache: list[Any], *, compact: bool = False) -> list[Any]:
213
+ copied = []
214
+ arrays = []
215
+ for layer in cache:
216
+ clone = copy.copy(layer)
217
+ source = layer.state if compact else (layer.keys, layer.values)
218
+ if source[0] is not None:
219
+ clone.keys = mx.array(source[0])
220
+ clone.values = mx.array(source[1])
221
+ arrays.extend((clone.keys, clone.values))
222
+ copied.append(clone)
223
+ if arrays:
224
+ mx.eval(*arrays)
225
+ return copied
226
+
227
+
228
+ @dataclass
229
+ class _PrefixEntry:
230
+ cache: list[Any]
231
+ bytes: int
232
+
233
+
234
+ class PrefixKVCache:
235
+ """LRU cache of immutable, block-boundary prompt KV snapshots."""
236
+
237
+ def __init__(self, limit_bytes: int):
238
+ self.limit_bytes = limit_bytes
239
+ self.entries: OrderedDict[tuple[int, ...], _PrefixEntry] = OrderedDict()
240
+ self.bytes = 0
241
+ self.hits = 0
242
+ self.reused_tokens = 0
243
+
244
+ def longest(self, ids: tuple[int, ...]) -> tuple[int, list[Any] | None]:
245
+ lengths = sorted(
246
+ {len(key) for key in self.entries if len(key) <= len(ids)},
247
+ reverse=True,
248
+ )
249
+ for length in lengths:
250
+ key = ids[:length]
251
+ entry = self.entries.get(key)
252
+ if entry is not None:
253
+ self.entries.move_to_end(key)
254
+ self.hits += 1
255
+ self.reused_tokens += length
256
+ return length, _clone_cache(entry.cache)
257
+ return 0, None
258
+
259
+ def put(self, ids: tuple[int, ...], cache: list[Any]) -> None:
260
+ if self.limit_bytes <= 0 or ids in self.entries:
261
+ return
262
+ size = sum(int(value.nbytes) for layer in cache
263
+ for value in (layer.state if layer.keys is not None else ())
264
+ if value is not None)
265
+ if size > self.limit_bytes:
266
+ return
267
+ snapshot = _clone_cache(cache, compact=True)
268
+ while self.entries and self.bytes + size > self.limit_bytes:
269
+ _, victim = self.entries.popitem(last=False)
270
+ self.bytes -= victim.bytes
271
+ self.entries[ids] = _PrefixEntry(snapshot, size)
272
+ self.bytes += size
273
+
274
+
275
+ @dataclass
276
+ class InferenceRow:
277
+ request: ContinuousRequest
278
+ prompt_ids: tuple[int, ...]
279
+ cache: list[Any]
280
+ rolling: MLXRollingState
281
+ generated: list[int]
282
+ repetition_seen: set[int]
283
+ created: float
284
+ admitted: float
285
+ denoise_steps: int = 0
286
+ jumps: int = 0
287
+ shifts: int = 0
288
+ first_token_at: float | None = None
289
+ stop_reason: str | None = None
290
+ first_denoise_at: float | None = None
291
+ last_denoise_at: float | None = None
292
+ prefill_seconds: float = 0.0
293
+
294
+
295
+ @dataclass
296
+ class PrefillRow:
297
+ request: ContinuousRequest
298
+ prompt_ids: tuple[int, ...]
299
+ cache: list[Any]
300
+ offset: int
301
+ reused: int
302
+ created: float
303
+ seconds: float = 0.0
304
+
305
+
306
+ @dataclass
307
+ class _DenoiseWork:
308
+ rows: list[InferenceRow]
309
+ rolling: MLXRollingState
310
+ output: Any
311
+ proposal: mx.array
312
+ remaining: mx.array
313
+ physical_positions: mx.array
314
+ next_latent: MLXLatentState
315
+ policy: Any
316
+
317
+
318
+ def _forward_independent_rows(model: MLXModilifyMk2,
319
+ rows: list[InferenceRow]) -> list[MLXCanvasOutput]:
320
+ """Interleave decoder layers while preserving independent singleton rows."""
321
+ decoder = model.model.decoder
322
+ latent = model.latent_deliberation
323
+ working_bus, persistent_bus = latent.working_memory_bus, latent.persistent_memory_bus
324
+ contexts = []
325
+ for row in rows:
326
+ rolling = row.rolling
327
+ canvas, state, head = rolling.canvas, rolling.latent, rolling.head
328
+ batch, length = canvas.shape
329
+ tokens = decoder.embed_tokens(canvas) * decoder.embed_scale
330
+ working, next_state = latent(token_embeddings=tokens,
331
+ confidence=state.confidence,
332
+ entropy=state.entropy, state=state,
333
+ canvas_head=head)
334
+ physical = (head[:, None] + mx.arange(length)[None, :]) % length
335
+ gather = mx.broadcast_to(physical[:, :, None], tokens.shape)
336
+ logical_tokens = mx.take_along_axis(tokens, mx.stop_gradient(gather), axis=1)
337
+ logical_working = mx.take_along_axis(working, mx.stop_gradient(gather), axis=1)
338
+ hidden = model._merge_context(logical_tokens, logical_working)
339
+ prefix_length = int(row.cache[0].offset)
340
+ full_mask = mx.ones((batch, prefix_length + length), mx.bool_)
341
+ masks = decoder._make_decoder_masks(hidden, row.cache, full_mask)
342
+ seen = latent.logical_seen(state, head)
343
+ contexts.append({"hidden": hidden, "working": working, "state": next_state,
344
+ "tokens": tokens, "head": head, "masks": masks,
345
+ "offset": prefix_length,
346
+ "working_kv": working_bus.prepare_kv((logical_working, seen)),
347
+ "persistent_kv": persistent_bus.prepare_kv(state.memory_slots, seen)})
348
+ reader = 0
349
+ for index, layer in enumerate(decoder.layers):
350
+ for row, context in zip(rows, contexts, strict=True):
351
+ hidden = layer(context["hidden"], context["masks"][layer.layer_type],
352
+ row.cache[index], decoder=True, offset=context["offset"])
353
+ if layer.layer_type == "full_attention":
354
+ if context["working_kv"] is not None and reader < working_bus.num_readers:
355
+ hidden = working_bus.read(hidden, reader, *context["working_kv"])
356
+ if context["persistent_kv"] is not None and reader < persistent_bus.num_readers:
357
+ hidden = persistent_bus.read(hidden, reader, *context["persistent_kv"])
358
+ context["hidden"] = hidden
359
+ # Explicit layer boundaries keep the execution order interleaved and
360
+ # bound intermediate activations to the configured pipeline depth.
361
+ mx.async_eval(*(context["hidden"] for context in contexts))
362
+ if layer.layer_type == "full_attention":
363
+ reader += 1
364
+ outputs = []
365
+ for context in contexts:
366
+ hidden = decoder.norm(context["hidden"])
367
+ length = hidden.shape[1]
368
+ inverse = (mx.arange(length)[None, :] - context["head"][:, None]) % length
369
+ heavy = mx.take_along_axis(hidden, mx.stop_gradient(mx.broadcast_to(
370
+ inverse[:, :, None], hidden.shape)), axis=1)
371
+ outputs.append(MLXCanvasOutput(heavy, context["working"],
372
+ context["state"], context["tokens"]))
373
+ return outputs
374
+
375
+
376
+ def _noise(seed: int, step: int, canvas: int, vocab: int) -> mx.array:
377
+ return mx.random.randint(0, vocab, (1, canvas),
378
+ key=_seed(seed, 0x51A7 + step * 1000003))
379
+
380
+
381
+ def _empty_rolling(config: Any, seed: int, max_new_tokens: int,
382
+ pad_token_id: int) -> MLXRollingState:
383
+ canvas = int(config.canvas_length)
384
+ vocab = int(config.text_config.vocab_size)
385
+ latent = MLXLatentState.empty(1, canvas, int(config.latent_memory_slots),
386
+ int(config.latent_dim), enable_gdn2=True)
387
+ latent = replace(latent, entropy=mx.full((1, canvas), math.log(vocab), mx.float32))
388
+ initial = mx.where(mx.arange(canvas)[None, :] < max_new_tokens,
389
+ _noise(seed, 0, canvas, vocab), pad_token_id)
390
+ return MLXRollingState(initial, latent,
391
+ mx.zeros((1,), mx.int32))
392
+
393
+
394
+ def _inference_statistics(hidden: mx.array, weight: mx.array,
395
+ rows: list[InferenceRow], *, softcap: float,
396
+ chunk_size: int, penalty: float,
397
+ excluded: set[int],
398
+ top_k: int | None = 40,
399
+ min_p: float | None = 0.05) -> tuple[mx.array, ...]:
400
+ """Exact chunked Top-K and Min-P Gumbel sampling without a retained full-vocabulary matrix."""
401
+ batch, canvas, dim = hidden.shape
402
+ flat = hidden.reshape(batch * canvas, dim)
403
+ neg_inf = mx.full((batch * canvas,), -mx.inf, mx.float32)
404
+ log_z = neg_inf
405
+ best_gumbel = neg_inf
406
+ chosen_score = mx.zeros_like(log_z)
407
+ chosen = mx.zeros((batch * canvas,), mx.int32)
408
+ greedy_score = neg_inf
409
+ greedy = mx.zeros_like(chosen)
410
+ moment_max = neg_inf
411
+ moment_sum = mx.zeros_like(log_z)
412
+ moment_weighted = mx.zeros_like(log_z)
413
+ repetition_mask = None
414
+ if penalty != 1.0:
415
+ masks = []
416
+ for row in rows:
417
+ mask = mx.zeros((weight.shape[0],), mx.bool_)
418
+ eligible = sorted(row.repetition_seen - excluded)
419
+ if eligible:
420
+ mask[mx.array(eligible, mx.int32)] = True
421
+ masks.append(mask)
422
+ repetition_mask = mx.stack(masks, axis=0)
423
+
424
+ use_constrained = top_k is not None and top_k > 0
425
+ chunk_cand_scores = []
426
+ chunk_cand_tokens = []
427
+
428
+ for chunk_number, start in enumerate(range(0, weight.shape[0], chunk_size)):
429
+ stop = min(start + chunk_size, weight.shape[0])
430
+ score = mx.tanh((flat @ weight[start:stop].T).astype(mx.float32) / softcap) * softcap
431
+ if penalty != 1.0:
432
+ assert repetition_mask is not None
433
+ row_scores = score.reshape(batch, canvas, stop - start)
434
+ seen = repetition_mask[:, None, start:stop]
435
+ changed = mx.where(row_scores < 0, row_scores * penalty, row_scores / penalty)
436
+ score = mx.where(seen, changed, row_scores).reshape(batch * canvas, stop - start)
437
+ score = score / DENOISE_TEMPERATURE
438
+ log_z = mx.logaddexp(log_z, mx.logsumexp(score, axis=-1))
439
+
440
+ local_max = mx.max(score, axis=-1)
441
+ better_greedy = local_max > greedy_score
442
+ greedy = mx.where(better_greedy, mx.argmax(score, axis=-1).astype(mx.int32) + start, greedy)
443
+ greedy_score = mx.maximum(greedy_score, local_max)
444
+
445
+ shifted = mx.exp(score - local_max[:, None])
446
+ next_max = mx.maximum(moment_max, local_max)
447
+ old_scale = mx.exp(moment_max - next_max)
448
+ new_scale = mx.exp(local_max - next_max)
449
+ moment_sum = moment_sum * old_scale + mx.sum(shifted, axis=-1) * new_scale
450
+ moment_weighted = moment_weighted * old_scale + mx.sum(shifted * score, axis=-1) * new_scale
451
+ moment_max = next_max
452
+
453
+ if use_constrained:
454
+ k = min(top_k, stop - start)
455
+ chunk_top_idx = mx.stop_gradient(mx.argpartition(-score, kth=k - 1, axis=-1)[:, :k]).astype(mx.int32)
456
+ chunk_top_score = mx.take_along_axis(score, chunk_top_idx, axis=-1)
457
+ chunk_cand_scores.append(chunk_top_score)
458
+ chunk_cand_tokens.append(chunk_top_idx + start)
459
+ else:
460
+ uniform = mx.concatenate([
461
+ mx.random.uniform(shape=(canvas, stop - start),
462
+ key=_seed(row.request.seed,
463
+ 0xC09A + row.denoise_steps * 1000003 + chunk_number))
464
+ for row in rows
465
+ ], axis=0)
466
+ uniform = mx.clip(uniform, 1.17549435e-38, 1.0 - 1.19209290e-7)
467
+ gumbel = score - mx.log(-mx.log(uniform))
468
+ local_index = mx.argmax(gumbel, axis=-1)
469
+ local_best = mx.max(gumbel, axis=-1)
470
+ better = local_best > best_gumbel
471
+ chosen_score = mx.where(better, mx.take_along_axis(score, local_index[:, None], axis=-1)[:, 0], chosen_score)
472
+ chosen = mx.where(better, local_index + start, chosen)
473
+ best_gumbel = mx.maximum(best_gumbel, local_best)
474
+
475
+ # Materialize each chunk without a CPU/GPU barrier. Include candidates
476
+ # so their argpartition graph does not retain every vocabulary slab.
477
+ live = [log_z, greedy, greedy_score, moment_max, moment_sum, moment_weighted]
478
+ if use_constrained:
479
+ live.extend((chunk_cand_scores[-1], chunk_cand_tokens[-1]))
480
+ else:
481
+ live.extend((best_gumbel, chosen_score, chosen))
482
+ mx.async_eval(*live)
483
+
484
+ if use_constrained:
485
+ all_cand_scores = mx.concatenate(chunk_cand_scores, axis=-1)
486
+ all_cand_tokens = mx.concatenate(chunk_cand_tokens, axis=-1)
487
+ global_k = min(top_k, all_cand_scores.shape[-1])
488
+ global_idx = mx.stop_gradient(mx.argpartition(-all_cand_scores, kth=global_k - 1, axis=-1)[:, :global_k]).astype(mx.int32)
489
+ cand_scores = mx.take_along_axis(all_cand_scores, global_idx, axis=-1)
490
+ cand_tokens = mx.take_along_axis(all_cand_tokens, global_idx, axis=-1)
491
+ if min_p is not None and min_p > 0.0:
492
+ thresh = greedy_score + math.log(float(min_p))
493
+ valid_cand = cand_scores >= thresh[:, None]
494
+ eligible_scores = mx.where(valid_cand, cand_scores, -mx.inf)
495
+ else:
496
+ eligible_scores = cand_scores
497
+ uniform = mx.concatenate([
498
+ mx.random.uniform(shape=(canvas, global_k),
499
+ key=_seed(row.request.seed,
500
+ 0xC09A + row.denoise_steps * 1000003))
501
+ for row in rows
502
+ ], axis=0)
503
+ uniform = mx.clip(uniform, 1.17549435e-38, 1.0 - 1.19209290e-7)
504
+ gumbel = eligible_scores - mx.log(-mx.log(uniform))
505
+ chosen_local = mx.stop_gradient(mx.argmax(gumbel, axis=-1))
506
+ chosen_score = mx.take_along_axis(cand_scores, chosen_local[:, None], axis=-1)[:, 0]
507
+ chosen = mx.take_along_axis(cand_tokens, chosen_local[:, None], axis=-1)[:, 0]
508
+
509
+ confidence = mx.clip(mx.exp(chosen_score - log_z), 0.0, 1.0)
510
+ greedy_confidence = mx.clip(mx.exp(greedy_score - log_z), 0.0, 1.0)
511
+ entropy = log_z - moment_weighted / mx.maximum(moment_sum, 1.17549435e-38)
512
+ shape = (batch, canvas)
513
+ return (chosen.reshape(shape), confidence.reshape(shape), entropy.reshape(shape),
514
+ greedy.reshape(shape), greedy_confidence.reshape(shape))
515
+
516
+
517
+ class MLXContinuousEngine:
518
+ def __init__(self, runtime: MLXRuntime, *, prefix_mib: int,
519
+ prefill_chunk: int, vocab_chunk: int, max_batch_rows: int,
520
+ max_batch_tokens: int, emit: Any, pipeline_depth: int = 2,
521
+ token_events: bool = True):
522
+ self.runtime = runtime
523
+ if min(prefill_chunk, vocab_chunk, max_batch_rows, pipeline_depth) <= 0:
524
+ raise ValueError("Chunk sizes, batch rows and pipeline depth must be positive.")
525
+ if max_batch_tokens < int(runtime.config.canvas_length):
526
+ raise ValueError("--max-batch-tokens must fit at least one canvas.")
527
+ self.prefix = PrefixKVCache(prefix_mib * 1024 * 1024)
528
+ self.prefill_chunk = prefill_chunk
529
+ self.vocab_chunk = vocab_chunk
530
+ self.max_batch_rows = max_batch_rows
531
+ self.max_batch_tokens = max_batch_tokens
532
+ self.emit = emit
533
+ self.token_events = token_events
534
+ self.pipeline_depth = pipeline_depth
535
+ self.forward_count = 0
536
+ self.active_row_steps = 0
537
+ self.denoise_seconds = 0.0
538
+ self.max_inflight = 0
539
+ self.peak_active_requests = 0
540
+ generation = runtime.generation
541
+ pad_token_id = generation.pad_token_id
542
+ if pad_token_id is None:
543
+ pad_token_id = getattr(runtime.config, "pad_token_id", None)
544
+ if isinstance(pad_token_id, (list, tuple)):
545
+ pad_token_id = pad_token_id[0]
546
+ self.pad_token_id = int(0 if pad_token_id is None else pad_token_id)
547
+ configured_eos = generation.eos_token_id or runtime.config.eos_token_id
548
+ if isinstance(configured_eos, int):
549
+ configured_eos = [configured_eos]
550
+ self.turn_end = int(runtime.config.turn_end_token_id if
551
+ generation.turn_end_token_id is None else
552
+ generation.turn_end_token_id)
553
+ self.stops = tuple(dict.fromkeys((self.turn_end, *(int(x) for x in configured_eos or ()))))
554
+ token_values = [generation.pad_token_id, generation.bos_token_id,
555
+ generation.eos_token_id, generation.turn_end_token_id,
556
+ getattr(runtime.config, "image_token_id", None),
557
+ generation.repetition_penalty_exclude_token_ids]
558
+ self.excluded = set()
559
+ for value in token_values:
560
+ if isinstance(value, int):
561
+ self.excluded.add(int(value))
562
+ elif isinstance(value, (tuple, list, set)):
563
+ self.excluded.update(int(x) for x in value if x is not None)
564
+
565
+ def begin_prefill(self, request: ContinuousRequest, created: float) -> PrefillRow:
566
+ encoded = apply_chat_template(self.runtime.tokenizer, request.messages,
567
+ think=request.think)
568
+ prompt = tuple(_extract_input_ids(encoded))
569
+ if len(prompt) + request.max_new_tokens > int(self.runtime.config.text_config.max_position_embeddings):
570
+ raise ValueError("Prompt plus response exceeds the model position limit.")
571
+ prefix_length, cache = self.prefix.longest(prompt)
572
+ if cache is None:
573
+ cache = self.runtime.model.model.encoder.make_cache()
574
+ return PrefillRow(request, prompt, cache, prefix_length, prefix_length, created)
575
+
576
+ def prefill_step(self, work: PrefillRow) -> InferenceRow | None:
577
+ """Execute at most one prefill chunk before yielding to active rows."""
578
+ started = time.perf_counter()
579
+ if work.offset < len(work.prompt_ids):
580
+ end = min(work.offset + self.prefill_chunk, len(work.prompt_ids))
581
+ block = mx.array(work.prompt_ids[work.offset:end], mx.int32)[None, :]
582
+ _, work.cache = self.runtime.model.model.encoder(block, cache=work.cache)
583
+ mx.eval(*(value for layer in work.cache for value in (layer.keys, layer.values)
584
+ if isinstance(value, mx.array)))
585
+ work.offset = end
586
+ self.prefix.put(work.prompt_ids[:end], work.cache)
587
+ if work.offset < len(work.prompt_ids):
588
+ work.seconds += time.perf_counter() - started
589
+ return None
590
+ request, prompt = work.request, work.prompt_ids
591
+ cache = _clone_cache(work.cache, compact=True)
592
+ work.seconds += time.perf_counter() - started
593
+ seen = {token for token in prompt if token not in self.excluded}
594
+ row = InferenceRow(request, prompt, cache,
595
+ _empty_rolling(self.runtime.config, request.seed,
596
+ request.max_new_tokens, self.pad_token_id), [],
597
+ seen, work.created, time.perf_counter(),
598
+ prefill_seconds=work.seconds)
599
+ self.emit({"event": "request_started", "request_id": request.request_id,
600
+ "prompt_tokens": len(prompt), "prefix_cache_tokens": work.reused,
601
+ "prefill_seconds": row.prefill_seconds})
602
+ return row
603
+
604
+ def admit(self, request: ContinuousRequest, created: float) -> InferenceRow:
605
+ work = self.begin_prefill(request, created)
606
+ while True:
607
+ row = self.prefill_step(work)
608
+ if row is not None:
609
+ return row
610
+
611
+ def _append_encoder(self, rows: list[InferenceRow], blocks: list[list[int]]) -> None:
612
+ # Keep the encoder's GEMM and expert routing shapes independent of the
613
+ # cohort too. Only the scheduler and GPU submission are concurrent.
614
+ for row, tokens in zip(rows, blocks, strict=True):
615
+ if tokens:
616
+ ids = mx.array(tokens, mx.int32)[None, :]
617
+ _, row.cache = self.runtime.model.model.encoder(ids, cache=row.cache)
618
+ mx.async_eval(*(value for layer in row.cache
619
+ for value in (layer.keys, layer.values)
620
+ if isinstance(value, mx.array)))
621
+
622
+ def step(self, rows: list[InferenceRow]) -> list[InferenceRow]:
623
+ if not rows:
624
+ return []
625
+ if len(rows) > min(self.max_batch_rows, self.max_batch_tokens //
626
+ int(self.runtime.config.canvas_length)):
627
+ raise ValueError("Denoise cohort exceeds the configured row/token budget.")
628
+ if len({id(row) for row in rows}) != len(rows):
629
+ raise ValueError("A request may appear only once in a denoise cohort.")
630
+ started = time.perf_counter()
631
+ finished = []
632
+ for start in range(0, len(rows), self.pipeline_depth):
633
+ cohort = rows[start:start + self.pipeline_depth]
634
+ outputs = _forward_independent_rows(self.runtime.model, cohort)
635
+ pending = [self._prepare_step([row], output=output)
636
+ for row, output in zip(cohort, outputs, strict=True)]
637
+ self.max_inflight = max(self.max_inflight, len(pending))
638
+ for work in pending:
639
+ finished.extend(self._finish_step(work))
640
+ self.denoise_seconds += time.perf_counter() - started
641
+ return finished
642
+
643
+ def _prepare_step(self, rows: list[InferenceRow], *,
644
+ output: MLXCanvasOutput) -> _DenoiseWork:
645
+ model, config, generation = (self.runtime.model, self.runtime.config,
646
+ self.runtime.generation)
647
+ rolling = _concat_rows([row.rolling for row in rows])
648
+ batch, canvas = rolling.canvas.shape
649
+ stats = _inference_statistics(
650
+ output.heavy_hidden, model.model.decoder.embed_tokens.weight,
651
+ rows, softcap=float(config.text_config.final_logit_softcapping),
652
+ chunk_size=self.vocab_chunk, penalty=float(generation.repetition_penalty),
653
+ excluded=self.excluded,
654
+ top_k=getattr(config, "commit_top_k", 40),
655
+ min_p=getattr(config, "commit_min_p", 0.05),
656
+ )
657
+ proposal, confidence, entropy, greedy, greedy_confidence = stats
658
+ confidence = confidence.astype(mx.float32)
659
+ entropy = entropy.astype(mx.float32)
660
+ changed = (proposal != rolling.canvas).astype(mx.float32)
661
+ remaining = mx.array([row.request.max_new_tokens - len(row.generated)
662
+ for row in rows], mx.int32)
663
+ physical_positions = (mx.arange(canvas)[None, :] - rolling.head[:, None]) % canvas
664
+ valid = physical_positions < remaining[:, None]
665
+ next_latent = replace(
666
+ output.next_latent_state, confidence=confidence, entropy=entropy,
667
+ age=rolling.latent.age + 1, token_changed=changed,
668
+ confidence_delta=confidence - rolling.latent.confidence,
669
+ entropy_delta=entropy - rolling.latent.entropy,
670
+ )
671
+ next_latent = _concat_rows([
672
+ model.latent_deliberation.observe_state(
673
+ _slice_row(next_latent, index), output.heavy_hidden[index:index + 1],
674
+ output.working_state[index:index + 1], valid[index:index + 1],
675
+ rolling.head[index:index + 1]) for index in range(batch)
676
+ ])
677
+ policy = select_commit_lengths(
678
+ _logical(proposal, rolling.head),
679
+ _logical(fused_commit_failure_rate(
680
+ confidence, entropy, entropy_weight=config.commit_entropy_weight,
681
+ confidence_power=config.commit_confidence_power,
682
+ top_k=getattr(config, "commit_top_k", None),
683
+ min_p=getattr(config, "commit_min_p", None),
684
+ target_confidence=getattr(config, "commit_target_confidence", None),
685
+ failure_budget=float(config.commit_failure_budget)), rolling.head),
686
+ _logical(fused_commit_failure_rate(
687
+ rolling.latent.confidence, rolling.latent.entropy,
688
+ entropy_weight=config.commit_entropy_weight,
689
+ confidence_power=config.commit_confidence_power,
690
+ top_k=getattr(config, "commit_top_k", None),
691
+ min_p=getattr(config, "commit_min_p", None),
692
+ target_confidence=getattr(config, "commit_target_confidence", None),
693
+ failure_budget=float(config.commit_failure_budget)), rolling.head),
694
+ _logical(greedy, rolling.head),
695
+ _logical(fused_commit_failure_rate(
696
+ greedy_confidence, entropy, entropy_weight=config.commit_entropy_weight,
697
+ confidence_power=config.commit_confidence_power,
698
+ top_k=getattr(config, "commit_top_k", None),
699
+ min_p=getattr(config, "commit_min_p", None),
700
+ target_confidence=getattr(config, "commit_target_confidence", None),
701
+ failure_budget=float(config.commit_failure_budget)), rolling.head),
702
+ ponder_steps=rolling.latent.ponder_steps,
703
+ stagnation_steps=rolling.latent.stagnation_steps,
704
+ active_rows=mx.ones((batch,), mx.bool_), remaining_lengths=remaining,
705
+ failure_budget=float(config.commit_failure_budget),
706
+ stop_token_id=self.stops,
707
+ stagnation_threshold=int(generation.jump_on_no_progress_after),
708
+ min_progress=float(generation.min_trajectory_progress),
709
+ max_ponder_steps=int(generation.max_ponder_steps),
710
+ valid_mask=mx.arange(canvas)[None, :] < remaining[:, None],
711
+ )
712
+ mx.async_eval(policy.commit_lengths, policy.commit_token_ids,
713
+ policy.jump_rows, policy.ponder_steps, policy.stagnation_steps,
714
+ *_arrays(next_latent), output.heavy_hidden, output.working_state)
715
+ self.forward_count += 1
716
+ return _DenoiseWork(rows, rolling, output, proposal, remaining,
717
+ physical_positions, next_latent, policy)
718
+
719
+ def _finish_step(self, work: _DenoiseWork) -> list[InferenceRow]:
720
+ rows, rolling, output = work.rows, work.rolling, work.output
721
+ proposal, remaining = work.proposal, work.remaining
722
+ physical_positions, next_latent, policy = (
723
+ work.physical_positions, work.next_latent, work.policy)
724
+ model, config, generation = (self.runtime.model, self.runtime.config,
725
+ self.runtime.generation)
726
+ batch, canvas = rolling.canvas.shape
727
+ lengths = policy.commit_lengths.astype(mx.int32)
728
+ selected = policy.commit_token_ids
729
+ mx.eval(lengths, selected, policy.jump_rows)
730
+ completed = time.perf_counter()
731
+ for row in rows:
732
+ row.last_denoise_at = completed
733
+ if row.first_denoise_at is None:
734
+ row.first_denoise_at = completed
735
+ self.emit({"event": "first_denoise", "request_id": row.request.request_id,
736
+ "ttfd_seconds": completed - row.created})
737
+ lengths_host = [int(x) for x in lengths.tolist()]
738
+ maximum = max(lengths_host)
739
+ selected_host = selected.tolist()
740
+ blocks = [[int(x) for x in selected_host[index][:lengths_host[index]]]
741
+ for index in range(batch)]
742
+ # The commit policy has fixed these tokens. Stream them before the
743
+ # writer and encoder append, which are only needed by the next step.
744
+ for index, row in enumerate(rows):
745
+ block = blocks[index]
746
+ if not block:
747
+ continue
748
+ if row.first_token_at is None:
749
+ row.first_token_at = time.perf_counter()
750
+ row.generated.extend(block)
751
+ row.repetition_seen.update(token for token in block
752
+ if token not in self.excluded)
753
+ if self.token_events:
754
+ self.emit({"event": "token", "request_id": row.request.request_id,
755
+ "token_ids": block, "generated_tokens": len(row.generated),
756
+ "text": self.runtime.tokenizer.decode(
757
+ row.generated, skip_special_tokens=False),
758
+ "denoise_steps": row.denoise_steps + 1})
759
+ next_canvas = mx.where(
760
+ (physical_positions < lengths[:, None]) & policy.jump_rows[:, None],
761
+ _physical(selected, rolling.head), proposal,
762
+ )
763
+ next_latent = replace(next_latent, ponder_steps=policy.ponder_steps,
764
+ stagnation_steps=policy.stagnation_steps)
765
+ if maximum:
766
+ embeddings = model.model.decoder.embed_tokens(selected[:, :maximum])
767
+ embeddings = embeddings * model.model.decoder.embed_scale
768
+ memory, _ = model.latent_deliberation.commit_write(
769
+ memory=next_latent.memory_slots,
770
+ working_state=output.working_state,
771
+
772
+
773
+ heavy_hidden=output.heavy_hidden,
774
+ committed_token_embeddings=embeddings,
775
+ commit_lengths=lengths,
776
+ prefix_lengths=mx.array([int(row.cache[0].offset) for row in rows], mx.int32),
777
+ commit_reason=infer_commit_reason(
778
+ lengths, jump_rows=policy.jump_rows,
779
+ commit_token_ids=selected, terminal_token_ids=self.stops,
780
+ ),
781
+ canvas_head=rolling.head, max_commit=maximum,
782
+ )
783
+ next_latent = replace(
784
+ next_latent, memory_slots=memory,
785
+ gdn2=replace(next_latent.gdn2, persistent=memory),
786
+ )
787
+ self._append_encoder(rows, blocks)
788
+ noise = mx.concatenate([
789
+ _noise(row.request.seed, row.denoise_steps + 1, canvas,
790
+ int(config.text_config.vocab_size)) for row in rows
791
+ ], axis=0)
792
+ next_rolling = MLXRollingState(next_canvas, next_latent,
793
+ rolling.head).advance_ring(
794
+ lengths, noise, entropy_fill_value=math.log(config.text_config.vocab_size),
795
+ )
796
+ logical_positions = (mx.arange(canvas)[None, :] - next_rolling.head[:, None]) % canvas
797
+ newly_exposed = logical_positions >= (canvas - lengths)[:, None]
798
+ next_rolling = replace(
799
+ next_rolling,
800
+ canvas=mx.where(newly_exposed &
801
+ (logical_positions >= (remaining - lengths)[:, None]),
802
+ self.pad_token_id, next_rolling.canvas),
803
+ )
804
+ # The following step depends on these arrays on the same MLX stream.
805
+ # Submit now and overlap the writer/refill with host stop/output work.
806
+ mx.async_eval(*_arrays(next_rolling))
807
+ finished = []
808
+ jumped = policy.jump_rows.tolist()
809
+ for index, row in enumerate(rows):
810
+ row.rolling = _slice_row(next_rolling, index)
811
+ row.denoise_steps += 1
812
+ self.active_row_steps += 1
813
+ row.jumps += int(jumped[index])
814
+ row.shifts += int(bool(blocks[index]))
815
+ if self.turn_end in blocks[index]:
816
+ row.stop_reason = "turn_end"
817
+ elif any(token in self.stops for token in blocks[index]):
818
+ row.stop_reason = "eos"
819
+ elif len(row.generated) >= row.request.max_new_tokens:
820
+ row.stop_reason = "max_new_tokens"
821
+ elif (row.request.max_denoising_steps is not None and
822
+ row.denoise_steps >= row.request.max_denoising_steps):
823
+ row.stop_reason = "max_denoising_steps"
824
+ else:
825
+ bound = row.request.max_new_tokens * int(generation.max_ponder_steps)
826
+ if row.denoise_steps >= bound:
827
+ row.stop_reason = "episode_watchdog"
828
+ if row.stop_reason is not None:
829
+ finished.append(row)
830
+ return finished
831
+
832
+ def result(self, row: InferenceRow) -> dict[str, Any]:
833
+ now = time.perf_counter()
834
+ return {"event": "generation_result", "request_id": row.request.request_id,
835
+ "status": "finished", "checkpoint": str(self.runtime.checkpoint),
836
+ "step": self.runtime.step, "thinking": "on" if row.request.think else "off",
837
+ "stop_reason": row.stop_reason, "prompt_tokens": len(row.prompt_ids),
838
+ "generated_tokens": len(row.generated), "token_ids": row.generated,
839
+ "text": self.runtime.tokenizer.decode(row.generated, skip_special_tokens=True),
840
+ "denoise_steps": row.denoise_steps, "jump_count": row.jumps,
841
+ "state_shift_count": row.shifts,
842
+ "tokens_per_forward": len(row.generated) / max(row.denoise_steps, 1),
843
+ "queue_seconds": row.admitted - row.created,
844
+ "prefill_seconds": row.prefill_seconds,
845
+ "ttfd_seconds": None if row.first_denoise_at is None else row.first_denoise_at - row.created,
846
+ "denoise_steps_per_second": row.denoise_steps / max(
847
+ (row.last_denoise_at or now) - row.admitted, 1e-9),
848
+ "ttft_seconds": None if row.first_token_at is None else row.first_token_at - row.created,
849
+ "elapsed_seconds": now - row.created,
850
+ "prefix_cache_hits": self.prefix.hits,
851
+ "prefix_cache_reused_tokens": self.prefix.reused_tokens,
852
+ "mlx_active_memory_bytes": int(mx.get_active_memory()),
853
+ "mlx_peak_memory_bytes": int(mx.get_peak_memory())}
854
+
855
+
856
+ def iter_scheduler(engine: MLXContinuousEngine, incoming: queue.Queue[Any],
857
+ max_queue_size: int, *,
858
+ continue_on_error: bool = False) -> Iterator[dict[str, Any]]:
859
+ """Bounded continuous admission with a prefill token budget per decode turn."""
860
+ active: deque[InferenceRow] = deque()
861
+ prefilling: deque[PrefillRow] = deque()
862
+ pending: deque[tuple[ContinuousRequest, float]] = deque()
863
+ seen_ids: set[str] = set()
864
+ input_done = False
865
+ started = time.perf_counter()
866
+ capacity = min(engine.max_batch_rows, engine.max_batch_tokens //
867
+ int(engine.runtime.config.canvas_length))
868
+
869
+ def accept(item: Any, created: float) -> dict[str, Any] | None:
870
+ nonlocal input_done
871
+ if item is None:
872
+ input_done = True
873
+ elif isinstance(item, dict):
874
+ return item
875
+ elif item.request_id in seen_ids:
876
+ return {"event": "request_error", "request_id": item.request_id,
877
+ "error": "Duplicate request_id."}
878
+ else:
879
+ seen_ids.add(item.request_id)
880
+ pending.append((item, created))
881
+
882
+ while not input_done or pending or prefilling or active:
883
+ while not input_done and len(pending) < max_queue_size:
884
+ try:
885
+ error = accept(*incoming.get_nowait())
886
+ if error is not None:
887
+ yield error
888
+ except queue.Empty:
889
+ break
890
+ prefill_budget = engine.prefill_chunk
891
+ had_active = bool(active)
892
+ while prefill_budget > 0:
893
+ if pending and len(active) + len(prefilling) < engine.max_batch_rows:
894
+ request, created = pending.popleft()
895
+ try:
896
+ # A short newcomer gets one chunk promptly; incomplete
897
+ # prompts then rotate in the bounded prefill cohort.
898
+ prefilling.appendleft(engine.begin_prefill(request, created))
899
+ except Exception as error:
900
+ yield {"event": "request_error", "request_id": request.request_id,
901
+ "error": str(error)}
902
+ continue
903
+ if not prefilling:
904
+ break
905
+ work = prefilling.popleft()
906
+ tokens = min(engine.prefill_chunk, len(work.prompt_ids) - work.offset)
907
+ if tokens > prefill_budget:
908
+ # Keep original chunk boundaries and singleton GEMM shapes.
909
+ prefilling.appendleft(work)
910
+ break
911
+ prefill_budget -= tokens
912
+ try:
913
+ row = engine.prefill_step(work)
914
+ if row is None:
915
+ prefilling.append(work)
916
+ else:
917
+ active.append(row)
918
+ except Exception as error:
919
+ yield {"event": "request_error", "request_id": work.request.request_id,
920
+ "error": str(error)}
921
+ if not had_active and active:
922
+ # Deliver the first request's initial denoise promptly. Once
923
+ # decoding, use the remaining token budget to fill short rows.
924
+ break
925
+ if active:
926
+ engine.peak_active_requests = max(getattr(engine, "peak_active_requests", 0),
927
+ len(active) + len(prefilling))
928
+ selected = [active.popleft() for _ in range(min(capacity, len(active)))]
929
+ try:
930
+ finished = engine.step(selected)
931
+ finished_ids = {id(row) for row in finished}
932
+ for row in selected:
933
+ if id(row) in finished_ids:
934
+ yield engine.result(row)
935
+ else:
936
+ active.append(row)
937
+ except Exception as error:
938
+ for row in selected:
939
+ yield {"event": "generation_result", "request_id": row.request.request_id,
940
+ "status": "failed", "error": str(error)}
941
+ if not continue_on_error:
942
+ raise
943
+ elif not pending and not prefilling and not input_done:
944
+ # Wake immediately for new input instead of polling every 10 ms.
945
+ error = accept(*incoming.get())
946
+ if error is not None:
947
+ yield error
948
+ tail_started = time.perf_counter()
949
+ mx.synchronize()
950
+ engine.denoise_seconds += time.perf_counter() - tail_started
951
+ elapsed = time.perf_counter() - started
952
+ yield {"event": "inference_summary", "elapsed_seconds": elapsed,
953
+ "active_row_denoises": engine.active_row_steps,
954
+ "heavy_forward_count": engine.forward_count,
955
+ "denoise_steps_per_second": engine.active_row_steps / max(elapsed, 1e-9),
956
+ "denoise_scheduler_seconds": engine.denoise_seconds,
957
+ "scheduler_denoises_per_second": engine.active_row_steps /
958
+ max(engine.denoise_seconds, 1e-9),
959
+ "pipeline_depth": engine.pipeline_depth,
960
+ "max_inflight_denoises": engine.max_inflight,
961
+ "peak_active_requests": getattr(engine, "peak_active_requests", 0),
962
+ "mlx_peak_memory_bytes": int(mx.get_peak_memory())}
963
+
964
+
965
+ def _run_scheduler(engine: MLXContinuousEngine, incoming: queue.Queue[Any],
966
+ max_queue_size: int) -> int:
967
+ had_error = False
968
+ for record in iter_scheduler(engine, incoming, max_queue_size):
969
+ had_error |= record["event"] == "request_error" or record.get("status") == "failed"
970
+ engine.emit(record)
971
+ return int(had_error)
972
+
973
+
974
+ def _parser() -> argparse.ArgumentParser:
975
+ parser = argparse.ArgumentParser(description="Modilify Mk2 native MLX text generation")
976
+ parser.add_argument("--model", "--checkpoint", dest="checkpoint",
977
+ default=(str(Path(__file__).resolve().parent)
978
+ if (Path(__file__).resolve().parent / "export_manifest.json").is_file()
979
+ else "Modilify/Modilify-Mk2-preview-mlx"),
980
+ help="Local model directory or Hugging Face Hub model ID.")
981
+ parser.add_argument("--prompt", default="Why is the sky blue?")
982
+ parser.add_argument("--requests-jsonl", metavar="PATH|-",
983
+ help="Read JSONL requests continuously; '-' reads stdin.")
984
+ parser.add_argument("--stream", action="store_true",
985
+ help="Stream generated text to stdout for a single --prompt request.")
986
+ parser.add_argument("--think", type=parse_bool, default=True)
987
+ parser.add_argument("--seed", type=int, default=42)
988
+ parser.add_argument("--max-new-tokens", type=int, default=8192)
989
+ parser.add_argument("--max-denoising-steps", type=int)
990
+ parser.add_argument("--canvas-length", type=int)
991
+ parser.add_argument("--repetition-penalty", type=float, default=1.0)
992
+ parser.add_argument("--batch-size", type=int, default=4)
993
+ parser.add_argument("--pipeline-depth", type=int, default=2,
994
+ help="Bound in-flight singleton denoises; 1 minimizes activation memory.")
995
+ parser.add_argument("--max-batch-tokens", type=int, default=1024)
996
+ parser.add_argument("--prefill-chunk-size", type=int, default=256)
997
+ parser.add_argument("--vocab-chunk-size", type=int, default=4096)
998
+ parser.add_argument("--prefix-cache-mib", type=int, default=512)
999
+ parser.add_argument("--max-queue-size", type=int, default=128)
1000
+ parser.add_argument("--commit-failure-budget", type=float, default=None,
1001
+ help="Override commit failure budget (defaults to checkpoint config).")
1002
+ parser.add_argument("--commit-top-k", type=int, default=None,
1003
+ help="Override commit top_k bound (defaults to checkpoint config).")
1004
+ parser.add_argument("--commit-min-p", type=float, default=None,
1005
+ help="Override commit min_p bound (defaults to checkpoint config).")
1006
+ parser.add_argument("--commit-target-confidence", type=float, default=None,
1007
+ help="Override commit target confidence (defaults to checkpoint config).")
1008
+ return parser
1009
+
1010
+
1011
+ def _validate_args(parser: argparse.ArgumentParser, args: Any) -> None:
1012
+ if args.stream and args.requests_jsonl is not None:
1013
+ parser.error("--stream can only be used with a single --prompt request.")
1014
+ positive = ("max_new_tokens", "batch_size", "max_batch_tokens",
1015
+ "prefill_chunk_size", "vocab_chunk_size", "max_queue_size", "pipeline_depth")
1016
+ for name in positive:
1017
+ if getattr(args, name) <= 0:
1018
+ parser.error(f"--{name.replace('_', '-')} must be positive.")
1019
+ if args.max_denoising_steps is not None and args.max_denoising_steps <= 0:
1020
+ parser.error("--max-denoising-steps must be positive.")
1021
+ if args.prefix_cache_mib < 0:
1022
+ parser.error("--prefix-cache-mib must be nonnegative.")
1023
+ if not math.isfinite(args.repetition_penalty) or args.repetition_penalty <= 0:
1024
+ parser.error("--repetition-penalty must be finite and positive.")
1025
+ if args.commit_failure_budget is not None and args.commit_failure_budget <= 0:
1026
+ parser.error("--commit-failure-budget must be positive.")
1027
+ if args.commit_top_k is not None and args.commit_top_k <= 0:
1028
+ parser.error("--commit-top-k must be positive.")
1029
+ if args.commit_min_p is not None and not (0.0 < args.commit_min_p < 1.0):
1030
+ parser.error("--commit-min-p must be in (0, 1).")
1031
+ if args.commit_target_confidence is not None and not (0.0 < args.commit_target_confidence < 1.0):
1032
+ parser.error("--commit-target-confidence must be in (0, 1).")
1033
+
1034
+
1035
+ def _producer(path: str, output: queue.Queue[Any], args: Any) -> None:
1036
+ stream = None
1037
+ try:
1038
+ stream = sys.stdin if path == "-" else open(path, encoding="utf-8")
1039
+ for line_number, line in enumerate(stream, 1):
1040
+ if not line.strip():
1041
+ continue
1042
+ try:
1043
+ record = json.loads(line)
1044
+ request = parse_continuous_request(
1045
+ record, default_max_new_tokens=args.max_new_tokens,
1046
+ default_max_denoising_steps=args.max_denoising_steps,
1047
+ default_seed=args.seed, default_think=args.think,
1048
+ )
1049
+ output.put((request, time.perf_counter()))
1050
+ except Exception as error:
1051
+ output.put(({"event": "request_error", "line": line_number,
1052
+ "error": str(error)}, time.perf_counter()))
1053
+ except Exception as error:
1054
+ output.put(({"event": "request_error", "error": str(error)},
1055
+ time.perf_counter()))
1056
+ finally:
1057
+ if stream is not None and stream is not sys.stdin:
1058
+ stream.close()
1059
+ output.put((None, time.perf_counter()))
1060
+
1061
+
1062
+ def main(argv: list[str] | None = None) -> int:
1063
+ parser = _parser()
1064
+ args = parser.parse_args(argv)
1065
+ _validate_args(parser, args)
1066
+ runtime = load_runtime(args.checkpoint, canvas_length=args.canvas_length,
1067
+ max_new_tokens=args.max_new_tokens,
1068
+ max_denoising_steps=args.max_denoising_steps,
1069
+ repetition_penalty=args.repetition_penalty,
1070
+ commit_failure_budget=args.commit_failure_budget,
1071
+ commit_top_k=args.commit_top_k,
1072
+ commit_min_p=args.commit_min_p,
1073
+ commit_target_confidence=args.commit_target_confidence)
1074
+ if args.max_denoising_steps is None:
1075
+ args.max_denoising_steps = runtime.generation.max_denoising_steps
1076
+
1077
+ streamed_token_ids: dict[str, list[int]] = {}
1078
+ streamed_text: dict[str, str] = {}
1079
+
1080
+ def emit(record: dict[str, Any]) -> None:
1081
+ if args.stream:
1082
+ event = record.get("event")
1083
+ request_id = str(record.get("request_id", "single"))
1084
+ if event == "token":
1085
+ ids = streamed_token_ids.setdefault(request_id, [])
1086
+ ids.extend(int(token_id) for token_id in record["token_ids"])
1087
+ current = runtime.tokenizer.decode(ids, skip_special_tokens=True)
1088
+ previous = streamed_text.get(request_id, "")
1089
+ if current.startswith(previous):
1090
+ delta = current[len(previous):]
1091
+ else:
1092
+ # Keep output append-only if a tokenizer revises a prior decode.
1093
+ common = 0
1094
+ for old_char, new_char in zip(previous, current):
1095
+ if old_char != new_char:
1096
+ break
1097
+ common += 1
1098
+ delta = current[common:]
1099
+ if delta:
1100
+ sys.stdout.write(delta)
1101
+ sys.stdout.flush()
1102
+ streamed_text[request_id] = current
1103
+ elif event == "generation_result":
1104
+ if record.get("status") == "failed":
1105
+ sys.stderr.write(f"\nInference failed: {record.get('error', 'unknown error')}\n")
1106
+ sys.stderr.flush()
1107
+ else:
1108
+ sys.stdout.write("\n")
1109
+ sys.stdout.flush()
1110
+ elif event == "request_error":
1111
+ sys.stderr.write(f"\nInference error: {record.get('error', 'unknown error')}\n")
1112
+ sys.stderr.flush()
1113
+ return
1114
+ sys.stdout.write(json.dumps(record, ensure_ascii=False) + "\n")
1115
+ sys.stdout.flush()
1116
+
1117
+ emit({"event": "runtime_loaded", "checkpoint": str(runtime.checkpoint),
1118
+ "step": runtime.step, "restored_trainable_tensors": runtime.tensor_count,
1119
+ "backend": "mlx", "canvas_length": runtime.config.canvas_length,
1120
+ "pipeline_depth": args.pipeline_depth,
1121
+ "execution": "layer_interleaved_singleton",
1122
+ "commit_failure_budget": runtime.config.commit_failure_budget,
1123
+ "commit_target_confidence": getattr(runtime.config, "commit_target_confidence", None),
1124
+ "commit_top_k": getattr(runtime.config, "commit_top_k", None),
1125
+ "commit_min_p": getattr(runtime.config, "commit_min_p", None)})
1126
+ engine = MLXContinuousEngine(
1127
+ runtime, prefix_mib=args.prefix_cache_mib,
1128
+ prefill_chunk=args.prefill_chunk_size,
1129
+ vocab_chunk=args.vocab_chunk_size,
1130
+ max_batch_rows=args.batch_size,
1131
+ max_batch_tokens=args.max_batch_tokens, emit=emit, pipeline_depth=args.pipeline_depth,
1132
+ )
1133
+ incoming: queue.Queue[Any] = queue.Queue(maxsize=max(2, args.max_queue_size))
1134
+ if args.requests_jsonl is None:
1135
+ request = ContinuousRequest("single", [{"role": "user", "content": args.prompt}],
1136
+ args.max_new_tokens, args.max_denoising_steps,
1137
+ args.seed, args.think, args.prompt)
1138
+ incoming.put((request, time.perf_counter()))
1139
+ incoming.put((None, time.perf_counter()))
1140
+ else:
1141
+ threading.Thread(target=_producer, args=(args.requests_jsonl, incoming, args),
1142
+ daemon=True).start()
1143
+ return _run_scheduler(engine, incoming, args.max_queue_size)
1144
+
1145
+
1146
+ def load_model(
1147
+ model: str = "Modilify/Modilify-Mk2-preview-mlx", *,
1148
+ canvas_length: int | None = None,
1149
+ ) -> MLXRuntime:
1150
+ """Load the complete local model or its Hugging Face snapshot once."""
1151
+ return load_runtime(model, canvas_length=canvas_length,
1152
+ max_new_tokens=256, max_denoising_steps=None,
1153
+ repetition_penalty=1.0)
1154
+
1155
+
1156
+ def generate(
1157
+ runtime: MLXRuntime, prompt: str | None = None, *,
1158
+ messages: list[dict[str, Any]] | None = None,
1159
+ max_new_tokens: int = 256, max_denoising_steps: int | None = None,
1160
+ seed: int = 42, think: bool = True,
1161
+ ) -> dict[str, Any]:
1162
+ """Generate one independent response, returning text, tokens, and metrics."""
1163
+ record = {"request_id": "single", "max_new_tokens": max_new_tokens,
1164
+ "max_denoising_steps": max_denoising_steps, "seed": seed, "think": think}
1165
+ if prompt is not None:
1166
+ record["prompt"] = prompt
1167
+ if messages is not None:
1168
+ record["messages"] = messages
1169
+ request = parse_continuous_request(
1170
+ record, default_max_new_tokens=256,
1171
+ default_max_denoising_steps=runtime.generation.max_denoising_steps,
1172
+ default_seed=42, default_think=True,
1173
+ )
1174
+ engine = MLXContinuousEngine(
1175
+ runtime, prefix_mib=0, prefill_chunk=256, vocab_chunk=4096,
1176
+ max_batch_rows=1, max_batch_tokens=int(runtime.config.canvas_length),
1177
+ emit=lambda event: None, pipeline_depth=1, token_events=False,
1178
+ )
1179
+ row = engine.admit(request, time.perf_counter())
1180
+ while not engine.step([row]):
1181
+ pass
1182
+ return engine.result(row)
1183
+
1184
+
1185
+ if __name__ == "__main__":
1186
+ raise SystemExit(main())
model-00001.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:374380ce2183b488367e840c546c1f4c7d6251fd2740f8d19a89a32f257dbffb
3
+ size 776579654
model-00002.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3b11844cc3f237a2b8008c1a7df7ffe2d6e78b4e4789af9e8dc756d619c72cf6
3
+ size 1991115265
model-00003.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e4ca40905b0e8da647a20a735659268fce793827a64ccdf8de747b06fb604313
3
+ size 1645281286
model-00004.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:8909471307dfe90a1b5d908f89d62ee9456754aa3fc1f4c5a1d36f2fb7b0ca6f
3
+ size 1645281307
model-00005.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:f213f94cb031abe72a258c325145aeb6c0a6fecf9c6d17e46298207141d27cc7
3
+ size 1645281376
model-00006.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:7d4997e1fdc157712d55572f1f427400bdb176888e6ad420bbb8567d6e50fbde
3
+ size 1674191579
model-00007.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:ec97a3868a1d77d35ff51e38e718db99d353ef44e0d320136966ec437f3376da
3
+ size 1645281438
model-00008.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c8aeeb93591a35251e6669fcb56d10101db3714468b5f4d1eecb27ae69a7cb98
3
+ size 1645281390
model-00009.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:43b686872a4680f4d021f2ddaf71a592601972645cc6f97136708a2b529dc4c2
3
+ size 1645281404
model-00010.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:946aa5129578785476e9e241757b64a0f4ec5777c873b576b25bbdc67d2f5f08
3
+ size 1645281370
model-00011.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:adb4df6d408098ec73a1ed3a6f4fbad0202fd93d4af150828996720341c141ec
3
+ size 1645281424
model-00012.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:86c9845b281f09b894ae62b3dde3574cd779417a8906b52a471b601c6ae19890
3
+ size 1674191611
model-00013.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:b2042c1d4bdb49b8a4ed0423a86c4694f2c0cf345e6d247022bc919b818db9e4
3
+ size 1645281386
model-00014.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:0233e8d5a91a79edd030479e287621af3bd962f080d7b1717dd03ab6fb417c97
3
+ size 1645281393
model-00015.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:75bab4540f1aa39d31271ec3d36e414e24fa500d3bce0dc5a00da0c1a8a25582
3
+ size 1645281363
model-00016.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:37e7c3bbf2fc81bd14f95dd5e14d49e7e6cda481b8478d64562a8be0d2c989fa
3
+ size 1645281392
model-00017.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:82dc6c75a7ca3a36ddb16733beead912507fcb5fbc236d06b9c249b9eb22ac09
3
+ size 1645281346
model-00018.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:b49d0e174f0d8422d737d7023c283a555f359c994a07e593aae5d07eb291c0ae
3
+ size 1645281342
model-00019.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e8e66177d4ad899a85ca712eb15fce721a7db756121ce3123d13b16eb3c0b0a8
3
+ size 1674191627
model-00020.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d87e3d52c22cbec2e5681d7c048ddbe5109df7f8c7af99c713c00dfe351b2ae6
3
+ size 1645281442
model-00021.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:44aecb5538ee528dcc90f49733268c67795a44a71b0fece8eb1b51b84cca59ea
3
+ size 1645281412
model-00022.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:70f3124fee705d9cb5221936c59f39820ab9c37a2c70bad89176e309e1c1ca07
3
+ size 1645281408
model-00023.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:499f422a0073dce797799000ab7ea6df9e77f845aad9cff999ac9d404a47b8a0
3
+ size 1645281368
model-00024.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e7cd9ffad7914610f2cf54be55ac46a5d2082d5faee7f3d7d2854fa46315d46e
3
+ size 1645281350
model-00025.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:234e835f2a3abdcd98cdbba245f057948573dd480301c94c00921f6d5c91bd96
3
+ size 1674191592
model-00026.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:57581d8a8c3efc5f19ac487469272cbf8d9fa68cb978c5d8f605701f3dca7fa2
3
+ size 1645281350
model-00027.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:a2d814a172600dd0169ceac390fde8d8683926d72be9bc025c368c9c9bf5e5e9
3
+ size 1645281304
model-00028.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:a7ee3b8626e6b76a82e0fbd78508a66151504ad02e5983f2895c735f3c41e704
3
+ size 1674191574
model-00029.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:51e757bfbf302e147a926faf4ad7ab2baa97355b5e8ffd688b06a62f74b95f6b
3
+ size 1645281378
model-00030.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:2dcc2427d1c1bd6d143283f4d2cede73940db3c093b59b45c1c1b879b5406631
3
+ size 1645281282
model-00031.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:df89b8419ce01b7d78f761aac8e324df22004cfeba4c9ae81c627dde94a73535
3
+ size 1645281334
model-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c62b07e428849f2db92bf286a58d2657c685deb2128692ef0de49d5df071f226
3
+ size 1166260864
model.safetensors.index.json ADDED
The diff for this file is too large to render. See raw diff
 
modilify_mk2/__init__.py ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ """Modilify Mk2 native MLX inference package."""
2
+
3
+ from .configuration_modilify_mk2 import ModilifyMk2Config, ModilifyMk2TextConfig
4
+ from transformers import AutoConfig
5
+
6
+ AutoConfig.register(ModilifyMk2Config.model_type, ModilifyMk2Config, exist_ok=True)
7
+
8
+ __version__ = "1.0.0"
9
+ __all__ = ["ModilifyMk2Config", "ModilifyMk2TextConfig"]
modilify_mk2/chat.py ADDED
@@ -0,0 +1,287 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Modilify chat-template rendering for MLX inference."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import hashlib
6
+ import json
7
+ import re
8
+ from typing import Any
9
+
10
+ GEMMA_THOUGHT_CLOSE = "<channel|>"
11
+
12
+ _THINK_LINE_RE = re.compile(r"^think(?:\r?\n|$)")
13
+
14
+ _CHANNEL_BLOCK_RE = re.compile(
15
+ r"<\|channel>thought\n(.*?)\n?<channel\|>\s*(.*)",
16
+ flags=re.DOTALL,
17
+ )
18
+
19
+ _LITERAL_THINK_RE = re.compile(
20
+ r"\s*<think>(.*?)</think>\s*(.*)",
21
+ flags=re.DOTALL,
22
+ )
23
+
24
+
25
+ def _split_thought_and_content(content: Any) -> tuple[str | None, str]:
26
+ if not isinstance(content, str):
27
+ return None, ""
28
+ text = content.strip()
29
+ if not text:
30
+ return None, ""
31
+ if "�" in text:
32
+ raise ValueError("Assistant target contains a Unicode replacement character.")
33
+ if _THINK_LINE_RE.match(text):
34
+ raise ValueError("Assistant target uses ambiguous literal think without channel markers.")
35
+ channel_match = _CHANNEL_BLOCK_RE.fullmatch(text)
36
+ literal_match = _LITERAL_THINK_RE.fullmatch(text)
37
+ has_channel_token = "<|channel>" in text or GEMMA_THOUGHT_CLOSE in text
38
+ has_literal_token = "<think>" in text or "</think>" in text
39
+ if has_channel_token and channel_match is None:
40
+ raise ValueError("Malformed Gemma channel in assistant target.")
41
+ if has_literal_token and literal_match is None:
42
+ raise ValueError("Malformed `<think>` block in assistant target.")
43
+ if channel_match is not None:
44
+ thought, answer = channel_match.groups()
45
+ elif literal_match is not None:
46
+ thought, answer = literal_match.groups()
47
+ else:
48
+ return None, text
49
+ thought = thought.strip()
50
+ return (thought or None), answer.strip()
51
+
52
+
53
+ def _assistant_thought_and_content(message: dict[str, Any]) -> tuple[str | None, str]:
54
+ thought, content = _split_thought_and_content(message.get("content"))
55
+ explicit = message.get("reasoning") or message.get("reasoning_content")
56
+ if isinstance(explicit, str) and explicit.strip():
57
+ return explicit.strip(), content
58
+ return thought, content
59
+
60
+
61
+ def _deserialize_tool_call_arguments(arguments: Any) -> dict[str, Any] | None:
62
+ """Convert OpenAI-style JSON argument strings into the mapping Gemma's template requires."""
63
+ if arguments is None or isinstance(arguments, dict):
64
+ return arguments
65
+ if not isinstance(arguments, str):
66
+ raise ValueError(
67
+ "chat_template: tool_calls[].function.arguments must be a JSON object "
68
+ f"(mapping), not a {type(arguments).__name__}."
69
+ )
70
+ text = arguments.strip()
71
+ if not text:
72
+ return {}
73
+ try:
74
+ parsed = json.loads(text)
75
+ except json.JSONDecodeError as error:
76
+ raise ValueError(
77
+ "chat_template: tool_calls[].function.arguments must be a JSON object "
78
+ "(mapping), not a string. Deserialize arguments before passing to "
79
+ f"the template: {error}"
80
+ ) from error
81
+ if parsed is None or isinstance(parsed, dict):
82
+ return parsed
83
+ raise ValueError(
84
+ "chat_template: tool_calls[].function.arguments must be a JSON object "
85
+ f"(mapping), not a {type(parsed).__name__}."
86
+ )
87
+
88
+
89
+ def _stable_tool_call_id(tool_call: dict[str, Any], index: int) -> str:
90
+ """Create a deterministic id for traces that omitted OpenAI call ids."""
91
+ payload = json.dumps(
92
+ tool_call,
93
+ ensure_ascii=False,
94
+ sort_keys=True,
95
+ separators=(",", ":"),
96
+ default=str,
97
+ ).encode("utf-8")
98
+ return f"call_modilify_mk2_{index}_{hashlib.sha1(payload).hexdigest()[:16]}"
99
+
100
+
101
+ def _normalize_message_tool_calls(message: dict[str, Any]) -> dict[str, Any]:
102
+ tool_calls = message.get("tool_calls")
103
+ if not isinstance(tool_calls, list) or not tool_calls:
104
+ return message
105
+ updated_calls = list(tool_calls)
106
+ changed = False
107
+ for index, tool_call in enumerate(tool_calls):
108
+ if not isinstance(tool_call, dict):
109
+ continue
110
+ function = tool_call.get("function")
111
+ # Some agent traces use the compact {name, arguments} shape instead
112
+ # of OpenAI's {function: {name, arguments}} wrapper.
113
+ if not isinstance(function, dict):
114
+ name = tool_call.get("name")
115
+ if not isinstance(name, str) or not name.strip():
116
+ continue
117
+ function = {
118
+ "name": name,
119
+ "arguments": tool_call.get(
120
+ "arguments", tool_call.get("input", {})
121
+ ),
122
+ }
123
+ changed = True
124
+ arguments = function.get("arguments")
125
+ parsed = (
126
+ arguments
127
+ if arguments is None or isinstance(arguments, dict)
128
+ else _deserialize_tool_call_arguments(arguments)
129
+ )
130
+ new_function = dict(function)
131
+ if parsed is not arguments:
132
+ changed = True
133
+ new_function["arguments"] = parsed
134
+ new_call = dict(tool_call)
135
+ if not isinstance(new_call.get("id"), str) or not new_call["id"]:
136
+ new_call["id"] = _stable_tool_call_id(tool_call, index)
137
+ changed = True
138
+ new_call.setdefault("type", "function")
139
+ new_call["function"] = new_function
140
+ updated_calls[index] = new_call
141
+ if not changed:
142
+ return message
143
+ updated = dict(message)
144
+ updated["tool_calls"] = updated_calls
145
+ return updated
146
+
147
+
148
+ def _normalize_assistant_message(message: dict[str, Any]) -> dict[str, Any]:
149
+ """Lift think/channel text into ``reasoning`` and deserialize tool arguments."""
150
+ updated = _normalize_message_tool_calls(message)
151
+ if updated.get("role") != "assistant":
152
+ return updated
153
+ thought, content = _assistant_thought_and_content(updated)
154
+ content_changed = content != (updated.get("content") or "")
155
+ reasoning = updated.get("reasoning")
156
+ needs_reasoning = bool(thought) and reasoning != thought
157
+ if not content_changed and not needs_reasoning:
158
+ return updated
159
+ if updated is message:
160
+ updated = dict(message)
161
+ else:
162
+ updated = dict(updated)
163
+ if thought:
164
+ updated["reasoning"] = thought
165
+ updated["content"] = content
166
+ return updated
167
+
168
+
169
+ def normalize_chat_template_messages(messages: Any) -> Any:
170
+ """Copy conversations into the official Gemma chat-template message schema."""
171
+ if not isinstance(messages, list) or not messages:
172
+ return messages
173
+ if isinstance(messages[0], list):
174
+ normalized_batch = None
175
+ for index, conversation in enumerate(messages):
176
+ normalized = normalize_chat_template_messages(conversation)
177
+ if normalized is conversation:
178
+ continue
179
+ if normalized_batch is None:
180
+ normalized_batch = list(messages)
181
+ normalized_batch[index] = normalized
182
+ return messages if normalized_batch is None else normalized_batch
183
+ normalized_messages = None
184
+ for index, message in enumerate(messages):
185
+ if not isinstance(message, dict):
186
+ continue
187
+ updated = _normalize_assistant_message(message)
188
+ if updated is message:
189
+ continue
190
+ if normalized_messages is None:
191
+ normalized_messages = list(messages)
192
+ normalized_messages[index] = updated
193
+ return messages if normalized_messages is None else normalized_messages
194
+
195
+
196
+ def normalize_tool_definitions(tools: Any) -> list[dict[str, Any]] | None:
197
+ """Normalize optional tool declarations and ignore trace-only tool metadata.
198
+
199
+ The native template accepts OpenAI declarations only. ``data-new`` also
200
+ contains JSON-encoded declarations and trace metadata shaped like
201
+ ``{name, arguments, tool_call_id}``; the latter are executed calls, not
202
+ declarations, and must not be passed to ``format_function_declaration``.
203
+ """
204
+ if tools is None:
205
+ return None
206
+ pending: list[Any]
207
+ if isinstance(tools, str):
208
+ try:
209
+ parsed = json.loads(tools)
210
+ except json.JSONDecodeError:
211
+ return None
212
+ pending = parsed if isinstance(parsed, list) else [parsed]
213
+ elif isinstance(tools, dict):
214
+ pending = [tools]
215
+ elif isinstance(tools, list):
216
+ pending = list(tools)
217
+ else:
218
+ return None
219
+
220
+ normalized: list[dict[str, Any]] = []
221
+ for item in pending:
222
+ if isinstance(item, str):
223
+ try:
224
+ item = json.loads(item)
225
+ except json.JSONDecodeError:
226
+ continue
227
+ if isinstance(item, list):
228
+ pending.extend(item)
229
+ continue
230
+ if not isinstance(item, dict):
231
+ continue
232
+ function = item.get("function")
233
+ if isinstance(function, dict):
234
+ name = function.get("name")
235
+ if not isinstance(name, str) or not name.strip():
236
+ continue
237
+ declaration = dict(function)
238
+ declaration["description"] = declaration.get("description", "")
239
+ declaration["parameters"] = declaration.get("parameters") or {}
240
+ normalized.append({
241
+ "type": "function",
242
+ "function": declaration,
243
+ })
244
+ continue
245
+ # Accept the common Anthropic/tool-schema spelling when it really is
246
+ # a declaration. Execution records with only `arguments` are skipped.
247
+ name = item.get("name")
248
+ parameters = item.get("parameters", item.get("input_schema"))
249
+ if isinstance(name, str) and name.strip() and isinstance(parameters, dict):
250
+ normalized.append({
251
+ "type": "function",
252
+ "function": {
253
+ "name": name,
254
+ "description": item.get("description", ""),
255
+ "parameters": parameters,
256
+ },
257
+ })
258
+ return normalized or None
259
+
260
+
261
+ def apply_chat_template(
262
+ processor: Any,
263
+ messages: Any,
264
+ *,
265
+ think: bool,
266
+ return_tensors: str | None = None,
267
+ padding: bool | str = False,
268
+ tools: Any = None,
269
+ ) -> Any:
270
+ """Render conversations with the Modilify tokenizer template."""
271
+
272
+ template_kwargs: dict[str, Any] = {
273
+ "tokenize": True,
274
+ "add_generation_prompt": True,
275
+ "enable_thinking": think,
276
+ "return_dict": True,
277
+ }
278
+ if return_tensors is not None:
279
+ template_kwargs["return_tensors"] = return_tensors
280
+ if padding:
281
+ template_kwargs["padding"] = padding
282
+ tools = normalize_tool_definitions(tools)
283
+ if tools:
284
+ template_kwargs["tools"] = tools
285
+ messages = normalize_chat_template_messages(messages)
286
+ encoded = processor.apply_chat_template(messages, **template_kwargs)
287
+ return encoded
modilify_mk2/configuration_modilify_mk2.py ADDED
@@ -0,0 +1,442 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Strict text-only configuration for the schema25 GDN2 protocol."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import math
7
+ from pathlib import Path
8
+ from collections.abc import Mapping, Sequence
9
+ from typing import Any
10
+
11
+ from transformers.configuration_utils import PreTrainedConfig
12
+
13
+ from transformers.models.diffusion_gemma import DiffusionGemmaTextConfig
14
+
15
+
16
+ STATE_SCHEMA_VERSION = 25
17
+ # Keep topology stable; CE scope and normalization have separate objective tags.
18
+ TRAINING_SCHEME = (
19
+ "gold_prefix_shared_commit_0_256_committed_ce_calibration_"
20
+ "causal_throughput_terminal_sft"
21
+ )
22
+ MEMORY_SCHEME = "dual_timescale_gdn2_trajectory_memory"
23
+ MEMORY_ARCHITECTURE = "compact_gdn2_v2"
24
+ CONFIG_PROTOCOL_ERROR = "ModilifyMk2 requires schema25 compact_gdn2_v2; older memory topologies require a new run."
25
+ HISTORY_VIEWS = 4
26
+ COMMIT_SEQUENCE_LAYERS = 2
27
+ DENOISE_TEMPERATURE = 0.8
28
+ VOCAB_CHUNK_SIZE = 32_768
29
+ COMMIT_READINESS_OBJECTIVE = "frontier_prefix_budget_v2"
30
+ TOKEN_CE_SUPERVISION = "valid_canvas_v1"
31
+ TOKEN_CE_NORMALIZATION = "per_sample_exposure_v1"
32
+
33
+
34
+ def require_current_checkpoint_protocol(metadata: Mapping[str, Any]) -> None:
35
+ """Reject checkpoints that do not implement the GDN2 state topology."""
36
+
37
+ if (
38
+ metadata.get("state_schema_version") != STATE_SCHEMA_VERSION
39
+ or metadata.get("training_scheme") != TRAINING_SCHEME
40
+ or metadata.get("memory_scheme") != MEMORY_SCHEME
41
+ or metadata.get("memory_architecture") != MEMORY_ARCHITECTURE
42
+ ):
43
+ raise RuntimeError(
44
+ f"ModilifyMk2 checkpoint requires schema{STATE_SCHEMA_VERSION} {MEMORY_ARCHITECTURE}; "
45
+ "older memory topologies cannot be restored. Start from the base model."
46
+ )
47
+
48
+
49
+ class ModilifyMk2TextConfig(DiffusionGemmaTextConfig):
50
+ model_type = "modilify_mk2_text"
51
+ vocab_size: int = 262_144
52
+ hidden_size: int = 2816
53
+ intermediate_size: int = 2112
54
+ num_hidden_layers: int = 30
55
+ num_attention_heads: int = 16
56
+ num_key_value_heads: int = 8
57
+ head_dim: int = 256
58
+ max_position_embeddings: int = 262_144
59
+ sliding_window: int = 1024
60
+ use_bidirectional_attention: str | None = None
61
+ num_global_key_value_heads: int | None = 2
62
+ global_head_dim: int = 512
63
+ num_experts: int | None = 128
64
+ top_k_experts: int | None = 8
65
+ moe_intermediate_size: int | None = 704
66
+
67
+
68
+ class ModilifyMk2Config(PreTrainedConfig):
69
+ """Configuration with recurrent GDN2 trajectory memory."""
70
+
71
+ model_type = "modilify_mk2"
72
+ sub_configs = {"text_config": ModilifyMk2TextConfig}
73
+
74
+ def __init__(
75
+ self,
76
+ text_config: ModilifyMk2TextConfig | dict[str, Any] | None = None,
77
+ *,
78
+ canvas_length: int = 256,
79
+ initializer_range: float = 0.02,
80
+ tie_word_embeddings: bool = True,
81
+ state_schema_version: int = STATE_SCHEMA_VERSION,
82
+ training_scheme: str = TRAINING_SCHEME,
83
+ memory_scheme: str = MEMORY_SCHEME,
84
+ memory_architecture: str = MEMORY_ARCHITECTURE,
85
+ latent_dim: int = 2816,
86
+ latent_ffn_dim: int = 7168,
87
+ latent_memory_slots: int = 256,
88
+ latent_num_layers: int = 4,
89
+ latent_num_heads: int = 16,
90
+ latent_local_attention_window: int = 128,
91
+ latent_history_length: int = 16,
92
+ latent_tape_probes: int = 1,
93
+ latent_tape_scheme: str = "gdn2_spatial_probe_v1",
94
+ latent_history_views: int = HISTORY_VIEWS,
95
+ latent_history_kv_rank: int | None = None,
96
+ latent_working_last_block_global: bool = True,
97
+ working_memory_bus: bool = True,
98
+ persistent_memory_bus: bool = True,
99
+ persistent_memory_write: str = "commit_only_transformer",
100
+ commit_sequence_layers: int = COMMIT_SEQUENCE_LAYERS,
101
+ commit_sequence_dim: int | None = None,
102
+ writer_slot_gate: str = "per_slot",
103
+ latent_working_bus_unfreeze_steps: int = 0,
104
+ latent_persistent_bus_unfreeze_steps: int = 0,
105
+ training_bptt_steps: int = 16,
106
+ kv_cache_bucket_size: int = 128,
107
+ turn_end_token_id: int = 106,
108
+ terminal_token_ids: Sequence[int] | None = None,
109
+ channel_end_token_id: int = 101,
110
+ eos_token_id: int = 1,
111
+ commit_failure_budget: float = 0.2,
112
+ commit_top_k: int | None = 40,
113
+ commit_min_p: float | None = 0.05,
114
+ commit_target_confidence: float | None = 0.5,
115
+ commit_entropy_weight: float = 1.0,
116
+ commit_confidence_power: float = 1.0,
117
+ commit_gold_alpha: float = 0.4,
118
+ commit_gold_weight: float = 1.0,
119
+ commit_readiness_failure_weight: float = 1.0,
120
+ token_loss_weight: float = 1.0,
121
+ token_ce_supervision: str = TOKEN_CE_SUPERVISION,
122
+ token_ce_normalization: str = TOKEN_CE_NORMALIZATION,
123
+ confidence_calibration_loss_weight: float = 0.1,
124
+ commit_readiness_objective: str = COMMIT_READINESS_OBJECTIVE,
125
+ commit_readiness_target_tokens: int = 16,
126
+ commit_readiness_loss_weight: float = 0.1,
127
+ commit_readiness_budget_margin: float = 0.02,
128
+ commit_readiness_beta: float = 0.02,
129
+ terminal_stop_loss_weight: float = 1.0,
130
+ terminal_stop_target_probability: float = 0.95,
131
+ **kwargs: Any,
132
+ ) -> None:
133
+ if any(key.startswith("commit_throughput_") for key in kwargs):
134
+ raise ValueError("Removed commit-throughput configuration fields.")
135
+ if token_ce_supervision != TOKEN_CE_SUPERVISION:
136
+ raise ValueError("Schema25 requires valid-canvas token CE.")
137
+ if state_schema_version != STATE_SCHEMA_VERSION:
138
+ raise RuntimeError(CONFIG_PROTOCOL_ERROR)
139
+ if training_scheme != TRAINING_SCHEME or memory_scheme != MEMORY_SCHEME:
140
+ raise RuntimeError(CONFIG_PROTOCOL_ERROR)
141
+ if memory_architecture != MEMORY_ARCHITECTURE:
142
+ raise RuntimeError("Schema25 requires compact_gdn2_v2 memory; start a new run.")
143
+ if text_config is None:
144
+ text_config = ModilifyMk2TextConfig()
145
+ elif isinstance(text_config, dict):
146
+ text_config = dict(text_config)
147
+ text_config.pop("model_type", None)
148
+ text_config["use_bidirectional_attention"] = None
149
+ text_config = ModilifyMk2TextConfig(**text_config)
150
+ elif not isinstance(text_config, ModilifyMk2TextConfig):
151
+ payload = text_config.to_dict()
152
+ payload.pop("model_type", None)
153
+ text_config = ModilifyMk2TextConfig(**payload)
154
+
155
+ self.text_config = text_config
156
+ self.canvas_length = canvas_length
157
+ self.initializer_range = initializer_range
158
+ self.state_schema_version = STATE_SCHEMA_VERSION
159
+ self.training_scheme = training_scheme
160
+ self.memory_scheme = memory_scheme
161
+ self.memory_architecture = memory_architecture
162
+ self.latent_dim = latent_dim
163
+ self.latent_ffn_dim = latent_ffn_dim
164
+ self.latent_memory_slots = latent_memory_slots
165
+ self.latent_num_layers = latent_num_layers
166
+ self.latent_num_heads = latent_num_heads
167
+ self.latent_local_attention_window = latent_local_attention_window
168
+ self.latent_history_length = latent_history_length
169
+ self.latent_tape_probes = int(latent_tape_probes)
170
+ self.latent_tape_scheme = str(latent_tape_scheme)
171
+ self.latent_history_views = int(latent_history_views)
172
+ if latent_history_kv_rank is None:
173
+ rank = min(1024, latent_dim)
174
+ rank -= rank % max(latent_num_heads, 1)
175
+ if rank <= 0:
176
+ rank = latent_num_heads
177
+ self.latent_history_kv_rank = rank
178
+ else:
179
+ self.latent_history_kv_rank = latent_history_kv_rank
180
+ self.latent_working_last_block_global = bool(latent_working_last_block_global)
181
+ self.working_memory_bus = bool(working_memory_bus)
182
+ self.persistent_memory_bus = bool(persistent_memory_bus)
183
+ self.persistent_memory_write = str(persistent_memory_write)
184
+ self.commit_sequence_layers = int(commit_sequence_layers)
185
+ if commit_sequence_dim is None:
186
+ self.commit_sequence_dim = self.latent_history_kv_rank
187
+ else:
188
+ self.commit_sequence_dim = int(commit_sequence_dim)
189
+ self.writer_slot_gate = str(writer_slot_gate)
190
+ self.latent_working_bus_unfreeze_steps = int(latent_working_bus_unfreeze_steps)
191
+ self.latent_persistent_bus_unfreeze_steps = int(
192
+ latent_persistent_bus_unfreeze_steps
193
+ )
194
+ self.training_bptt_steps = training_bptt_steps
195
+ self.kv_cache_bucket_size = kv_cache_bucket_size
196
+ self.turn_end_token_id = turn_end_token_id
197
+ if terminal_token_ids is None:
198
+ self.terminal_token_ids = (int(turn_end_token_id),)
199
+ else:
200
+ self.terminal_token_ids = tuple(int(token_id) for token_id in terminal_token_ids)
201
+ self.channel_end_token_id = int(channel_end_token_id)
202
+ self.commit_failure_budget = float(commit_failure_budget)
203
+ self.commit_top_k = int(commit_top_k) if commit_top_k is not None else None
204
+ self.commit_min_p = float(commit_min_p) if commit_min_p is not None else None
205
+ self.commit_target_confidence = float(commit_target_confidence) if commit_target_confidence is not None else None
206
+ self.commit_entropy_weight = float(commit_entropy_weight)
207
+ self.commit_confidence_power = float(commit_confidence_power)
208
+ self.commit_gold_alpha = float(commit_gold_alpha)
209
+ self.commit_gold_weight = float(commit_gold_weight)
210
+ self.commit_readiness_failure_weight = float(commit_readiness_failure_weight)
211
+ self.token_loss_weight = token_loss_weight
212
+ self.token_ce_supervision = str(token_ce_supervision)
213
+ self.token_ce_normalization = str(token_ce_normalization)
214
+ self.confidence_calibration_loss_weight = confidence_calibration_loss_weight
215
+ self.commit_readiness_objective = str(commit_readiness_objective)
216
+ self.commit_readiness_target_tokens = int(commit_readiness_target_tokens)
217
+ self.commit_readiness_loss_weight = float(commit_readiness_loss_weight)
218
+ self.commit_readiness_budget_margin = float(commit_readiness_budget_margin)
219
+ self.commit_readiness_beta = float(commit_readiness_beta)
220
+ self.terminal_stop_loss_weight = terminal_stop_loss_weight
221
+ self.terminal_stop_target_probability = terminal_stop_target_probability
222
+ self.vocab_chunk_size = VOCAB_CHUNK_SIZE
223
+ super().__init__(
224
+ tie_word_embeddings=tie_word_embeddings,
225
+ eos_token_id=eos_token_id,
226
+ **kwargs,
227
+ )
228
+ self._validate()
229
+
230
+ def _validate(self) -> None:
231
+ positive = (
232
+ self.canvas_length, self.latent_dim, self.latent_ffn_dim,
233
+ self.latent_memory_slots,
234
+ self.latent_num_layers, self.latent_num_heads,
235
+ self.latent_local_attention_window, self.latent_history_length,
236
+ self.latent_tape_probes,
237
+ self.latent_history_kv_rank,
238
+ self.commit_sequence_layers, self.commit_sequence_dim,
239
+ self.training_bptt_steps, self.kv_cache_bucket_size,
240
+ )
241
+ if any(value <= 0 for value in positive):
242
+ raise ValueError("All schema25 dimensions and intervals must be positive.")
243
+ if self.canvas_length != 256:
244
+ raise ValueError("Schema25 requires a 256-token canvas.")
245
+ if self.latent_tape_scheme != "gdn2_spatial_probe_v1":
246
+ raise ValueError("Schema25 requires GDN2 spatial probes.")
247
+ if self.latent_tape_probes > self.canvas_length:
248
+ raise ValueError("Spatial probes cannot exceed canvas positions.")
249
+ if self.commit_sequence_dim % self.latent_num_heads:
250
+ raise ValueError("`commit_sequence_dim` must be divisible by `latent_num_heads`.")
251
+ if self.latent_working_bus_unfreeze_steps < 0:
252
+ raise ValueError("`latent_working_bus_unfreeze_steps` must be non-negative.")
253
+ if self.latent_persistent_bus_unfreeze_steps < 0:
254
+ raise ValueError(
255
+ "`latent_persistent_bus_unfreeze_steps` must be non-negative."
256
+ )
257
+ if self.latent_persistent_bus_unfreeze_steps < self.latent_working_bus_unfreeze_steps:
258
+ raise ValueError(
259
+ "Persistent bus must unfreeze no earlier than the working bus."
260
+ )
261
+ if self.latent_history_kv_rank > self.latent_dim:
262
+ raise ValueError("`latent_history_kv_rank` must not exceed `latent_dim`.")
263
+ if self.latent_history_kv_rank % self.latent_num_heads:
264
+ raise ValueError("`latent_history_kv_rank` must be divisible by `latent_num_heads`.")
265
+ if not isinstance(self.eos_token_id, int) or self.eos_token_id < 0:
266
+ raise ValueError("ModilifyMk2 requires one non-negative integer EOS token ID.")
267
+ if not isinstance(self.channel_end_token_id, int) or self.channel_end_token_id < 0:
268
+ raise ValueError("`channel_end_token_id` must be a non-negative integer.")
269
+ if not self.terminal_token_ids:
270
+ raise ValueError("`terminal_token_ids` must not be empty.")
271
+ if any(
272
+ not isinstance(token_id, int) or token_id < 0
273
+ for token_id in self.terminal_token_ids
274
+ ):
275
+ raise ValueError("`terminal_token_ids` must be non-negative integers.")
276
+ if self.latent_dim % self.latent_num_heads:
277
+ raise ValueError("`latent_dim` must be divisible by `latent_num_heads`.")
278
+ if self.commit_readiness_objective != COMMIT_READINESS_OBJECTIVE:
279
+ raise ValueError(
280
+ "Unsupported commit-readiness objective: "
281
+ f"{self.commit_readiness_objective!r}."
282
+ )
283
+ if not all(math.isfinite(v) for v in (
284
+ self.commit_failure_budget, self.commit_entropy_weight,
285
+ self.commit_confidence_power, self.commit_gold_alpha, self.commit_gold_weight,
286
+ self.commit_readiness_failure_weight, self.commit_readiness_loss_weight,
287
+ )):
288
+ raise ValueError("Commit policy parameters must be finite.")
289
+ if self.commit_failure_budget <= 0 or self.commit_entropy_weight < 0:
290
+ raise ValueError("Commit budget must be positive and entropy weight non-negative.")
291
+ if self.commit_top_k is not None and self.commit_top_k <= 0:
292
+ raise ValueError("`commit_top_k` must be a positive integer.")
293
+ if self.commit_min_p is not None and not 0.0 < self.commit_min_p < 1.0:
294
+ raise ValueError("`commit_min_p` must be in (0, 1).")
295
+ if self.commit_target_confidence is not None and not 0.0 < self.commit_target_confidence < 1.0:
296
+ raise ValueError("`commit_target_confidence` must be in (0, 1).")
297
+ if self.commit_confidence_power <= 0 or not 0 < self.commit_gold_alpha < 1:
298
+ raise ValueError("Commit power must be positive and gold alpha in (0, 1).")
299
+ if self.commit_gold_weight <= 0:
300
+ raise ValueError("Commit gold weight must be positive.")
301
+ if self.commit_readiness_failure_weight < 0:
302
+ raise ValueError("Commit readiness failure weight must be non-negative.")
303
+ if self.commit_readiness_target_tokens <= 0:
304
+ raise ValueError("`commit_readiness_target_tokens` must be positive.")
305
+ if self.commit_readiness_loss_weight < 0:
306
+ raise ValueError("`commit_readiness_loss_weight` must be non-negative.")
307
+ if not 0.0 <= self.commit_readiness_budget_margin < self.commit_failure_budget:
308
+ raise ValueError(
309
+ "`commit_readiness_budget_margin` must be in [0, commit_failure_budget)."
310
+ )
311
+ if self.commit_readiness_beta <= 0:
312
+ raise ValueError("`commit_readiness_beta` must be positive.")
313
+ if self.terminal_stop_loss_weight < 0:
314
+ raise ValueError("`terminal_stop_loss_weight` must be non-negative.")
315
+ if not 0.0 < self.terminal_stop_target_probability < 1.0:
316
+ raise ValueError("`terminal_stop_target_probability` must be in (0, 1).")
317
+ loss_weights = (
318
+ self.token_loss_weight,
319
+ self.confidence_calibration_loss_weight,
320
+ self.commit_readiness_loss_weight,
321
+ self.terminal_stop_loss_weight,
322
+ )
323
+ if any(not math.isfinite(weight) or weight < 0 for weight in loss_weights):
324
+ raise ValueError("Schema25 fixed loss weights must be finite and non-negative.")
325
+ if self.token_ce_normalization != TOKEN_CE_NORMALIZATION:
326
+ raise ValueError("Unsupported token CE normalization.")
327
+
328
+ @classmethod
329
+ def from_dict(cls, config_dict: dict[str, Any], **kwargs: Any) -> "ModilifyMk2Config":
330
+ return_unused_kwargs = kwargs.pop("return_unused_kwargs", False)
331
+ payload = dict(config_dict)
332
+ if (payload.get("state_schema_version") != STATE_SCHEMA_VERSION
333
+ or payload.get("memory_architecture") != MEMORY_ARCHITECTURE):
334
+ raise RuntimeError(CONFIG_PROTOCOL_ERROR)
335
+ payload.pop("model_type", None)
336
+ payload.pop("architectures", None)
337
+ for key in tuple(kwargs):
338
+ if key in payload:
339
+ payload[key] = kwargs.pop(key)
340
+ config = cls(**payload)
341
+ for key in tuple(kwargs):
342
+ if hasattr(config, key) or key == "name_or_path":
343
+ setattr(config, key, kwargs.pop(key))
344
+ return (config, kwargs) if return_unused_kwargs else config
345
+
346
+
347
+ __all__ = [
348
+ "ModilifyMk2Config", "ModilifyMk2TextConfig", "CONFIG_PROTOCOL_ERROR",
349
+ "COMMIT_READINESS_OBJECTIVE", "COMMIT_SEQUENCE_LAYERS", "DENOISE_TEMPERATURE",
350
+ "HISTORY_VIEWS", "MEMORY_SCHEME", "MEMORY_ARCHITECTURE", "STATE_SCHEMA_VERSION",
351
+ "TRAINING_SCHEME", "VOCAB_CHUNK_SIZE",
352
+ "TOKEN_CE_SUPERVISION", "TOKEN_CE_NORMALIZATION",
353
+ "require_current_checkpoint_protocol",
354
+ ]
355
+
356
+
357
+ class ModilifyMk2GenerationConfig:
358
+ """Native settings; no autoregressive sampler or Torch runtime is needed."""
359
+
360
+ def __init__(self, **kwargs: Any):
361
+ unsupported = (
362
+ 'sampler_config', 'stability_threshold', 'confidence_threshold',
363
+ 'one_token_per_denoise_step', 'compile_generation', 'sliding_denoise',
364
+ 'adaptive_ponder_budget', 'force_commit_on_max_steps', 'ponder_budget_id',
365
+ )
366
+ configured = [name for name in unsupported if kwargs.pop(name, None) not in (None, False)]
367
+ if configured:
368
+ raise ValueError(f'Unsupported generation fields: {configured}')
369
+ defaults = {
370
+ 'max_new_tokens': None, 'max_denoising_steps': None,
371
+ 'bos_token_id': None, 'pad_token_id': None,
372
+ 'eos_token_id': None, 'turn_end_token_id': None, 'max_ponder_steps': 64,
373
+ 'jump_on_no_progress_after': 12, 'min_trajectory_progress': 0.005,
374
+ 'repetition_penalty': 1.0, 'repetition_penalty_exclude_token_ids': [],
375
+ }
376
+ for name, default in defaults.items():
377
+ setattr(self, name, kwargs.pop(name, default))
378
+ self.repetition_penalty_exclude_token_ids = list(dict.fromkeys(
379
+ int(value) for value in self.repetition_penalty_exclude_token_ids or ()
380
+ ))
381
+
382
+
383
+ def validate(self) -> None:
384
+ for name in ('max_new_tokens', 'max_denoising_steps',
385
+ 'max_ponder_steps', 'jump_on_no_progress_after'):
386
+ value = getattr(self, name)
387
+ if value is not None and (isinstance(value, bool) or not isinstance(value, int) or value <= 0):
388
+ raise ValueError(f'{name} must be a positive integer.')
389
+ if not math.isfinite(self.repetition_penalty) or self.repetition_penalty <= 0:
390
+ raise ValueError('repetition_penalty must be finite and positive.')
391
+ if not math.isfinite(self.min_trajectory_progress) or self.min_trajectory_progress < 0:
392
+ raise ValueError('min_trajectory_progress must be finite and nonnegative.')
393
+ if self.turn_end_token_id is not None and self.turn_end_token_id < 0:
394
+ raise ValueError('turn_end_token_id must be nonnegative.')
395
+ if any(isinstance(value, bool) or not isinstance(value, int) or value < 0
396
+ for value in self.repetition_penalty_exclude_token_ids):
397
+ raise ValueError('Excluded token IDs must be nonnegative integers.')
398
+
399
+
400
+ def load_generation_config(model_dir: str | Path) -> ModilifyMk2GenerationConfig:
401
+ path = Path(model_dir) / 'generation_config.json'
402
+ if not path.is_file():
403
+ raise FileNotFoundError(f'Missing generation configuration: {path}')
404
+ return ModilifyMk2GenerationConfig(**json.loads(path.read_text(encoding='utf-8')))
405
+
406
+
407
+ def configure_generation_config(
408
+ generation_config: ModilifyMk2GenerationConfig, processor: Any, *,
409
+ max_new_tokens: int, max_denoising_steps: int | None,
410
+ repetition_penalty: float | None = None,
411
+ ) -> ModilifyMk2GenerationConfig:
412
+ if generation_config.max_denoising_steps is None:
413
+ generation_config.max_denoising_steps = 48
414
+ generation_config.max_new_tokens = max_new_tokens
415
+ if max_denoising_steps is not None:
416
+ generation_config.max_denoising_steps = max_denoising_steps
417
+ if repetition_penalty is not None:
418
+ generation_config.repetition_penalty = repetition_penalty
419
+ tokenizer = getattr(processor, 'tokenizer', processor)
420
+ if generation_config.bos_token_id is None:
421
+ generation_config.bos_token_id = tokenizer.bos_token_id
422
+ if generation_config.pad_token_id is None:
423
+ generation_config.pad_token_id = tokenizer.pad_token_id
424
+ turn_id = tokenizer.convert_tokens_to_ids('<turn|>')
425
+ if turn_id is None or turn_id == getattr(tokenizer, 'unk_token_id', None):
426
+ raise ValueError('Tokenizer must define <turn|>.')
427
+ generation_config.turn_end_token_id = int(turn_id)
428
+ if generation_config.eos_token_id is None:
429
+ generation_config.eos_token_id = [value for value in [tokenizer.eos_token_id] if value is not None]
430
+ tool_id = tokenizer.convert_tokens_to_ids('<|tool_response>')
431
+ if tool_id is not None and tool_id != getattr(tokenizer, 'unk_token_id', None):
432
+ eos = generation_config.eos_token_id
433
+ eos = [eos] if isinstance(eos, int) else list(eos or ())
434
+ if int(tool_id) not in eos:
435
+ generation_config.eos_token_id = [*eos, int(tool_id)]
436
+ generation_config.repetition_penalty_exclude_token_ids = list(dict.fromkeys(
437
+ int(value) for value in [*generation_config.repetition_penalty_exclude_token_ids,
438
+ *(getattr(tokenizer, 'all_special_ids', None) or ())]
439
+ if value is not None
440
+ ))
441
+ generation_config.validate()
442
+ return generation_config
modilify_mk2/mlx_commit_policy.py ADDED
@@ -0,0 +1,255 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Native MLX schema25 confidence fusion and sample-independent commit policy."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import math
6
+ from collections.abc import Sequence
7
+ from dataclasses import dataclass
8
+
9
+ import mlx.core as mx
10
+
11
+
12
+ JUMP_FAILURE_BUDGET = 2.0
13
+ FUSED_EPS = 1.0e-6
14
+ COMMIT_REASON_NONE = 0
15
+ COMMIT_REASON_NORMAL = 1
16
+ COMMIT_REASON_FORCED_JUMP = 2
17
+ COMMIT_REASON_TERMINAL = 3
18
+
19
+
20
+ def commit_target_confidence_bias(
21
+ target_confidence: float | None,
22
+ failure_budget: float = 0.2,
23
+ budget_safety_ratio: float = 0.85,
24
+ ) -> float:
25
+ if target_confidence is None or target_confidence <= 0.0 or target_confidence >= 1.0:
26
+ return 0.0
27
+ target_failure = min(budget_safety_ratio * failure_budget, 1.0 - target_confidence)
28
+ target_failure = max(target_failure, 1.0e-4)
29
+ target_conf = 1.0 - target_failure
30
+ logit_c = math.log(target_conf / (1.0 - target_conf))
31
+ logit_p = math.log(target_confidence / (1.0 - target_confidence))
32
+ return max(logit_c - logit_p, 0.0)
33
+
34
+
35
+ def fused_commit_confidence(
36
+ proposal_confidence: mx.array,
37
+ token_entropy: mx.array,
38
+ *,
39
+ eps: float = FUSED_EPS,
40
+ entropy_weight: float = 1.0,
41
+ confidence_power: float = 2.0,
42
+ top_k: int | None = None,
43
+ min_p: float | None = None,
44
+ target_confidence: float | None = None,
45
+ failure_budget: float = 0.2,
46
+ ) -> mx.array:
47
+ """Match the schema25 excess-entropy sigmoid and checkpoint power."""
48
+
49
+ p = mx.clip(mx.nan_to_num(proposal_confidence.astype(mx.float32), nan=0.5), eps, 1.0 - eps)
50
+ entropy = mx.maximum(mx.nan_to_num(token_entropy.astype(mx.float32), nan=0.0), 0.0)
51
+ binary_entropy = -p * mx.log(p) - (1.0 - p) * mx.log1p(-p)
52
+ excess = mx.maximum(entropy - binary_entropy, 0.0)
53
+ k_eff = None
54
+ if top_k is not None and top_k > 0:
55
+ k_eff = mx.full(p.shape, float(top_k), dtype=mx.float32)
56
+ if min_p is not None and min_p > 0:
57
+ thresh = mx.maximum(float(min_p) * p, 1.0e-6)
58
+ k_min_p = mx.maximum((1.0 - p) / thresh, 1.0)
59
+ k_eff = k_min_p if k_eff is None else mx.minimum(k_eff, k_min_p)
60
+ if k_eff is not None:
61
+ max_excess = (1.0 - p) * mx.log(k_eff)
62
+ excess = mx.minimum(excess, max_excess)
63
+ bias = commit_target_confidence_bias(target_confidence, failure_budget=failure_budget)
64
+ fused_logit = mx.log(p) - mx.log1p(-p) + bias - entropy_weight * excess
65
+ return mx.clip(mx.sigmoid(fused_logit) ** confidence_power, eps, 1.0 - eps)
66
+
67
+
68
+ def fused_commit_failure_rate(
69
+ proposal_confidence: mx.array, token_entropy: mx.array, **kwargs: object
70
+ ) -> mx.array:
71
+ return 1.0 - fused_commit_confidence(proposal_confidence, token_entropy, **kwargs)
72
+
73
+
74
+ def _terminal_mask(tokens: mx.array, stop_token_id: int | Sequence[int]) -> mx.array:
75
+ ids = (stop_token_id,) if isinstance(stop_token_id, int) else tuple(dict.fromkeys(stop_token_id))
76
+ if not ids:
77
+ raise ValueError("At least one stop token ID is required.")
78
+ matches = tokens == int(ids[0])
79
+ for value in ids[1:]:
80
+ matches = matches | (tokens == int(value))
81
+ return matches
82
+
83
+
84
+ def prefix_failure_commit_lengths(
85
+ failure_rate: mx.array,
86
+ *,
87
+ failure_budget: float,
88
+ valid_mask: mx.array | None = None,
89
+ ) -> mx.array:
90
+ """Longest contiguous valid prefix whose cumulative failure stays below budget."""
91
+
92
+ if failure_rate.ndim != 2:
93
+ raise ValueError("Failure rate must have shape [batch, canvas].")
94
+ if failure_budget <= 0:
95
+ raise ValueError("Commit failure budget must be positive.")
96
+ if valid_mask is None:
97
+ valid_mask = mx.ones(failure_rate.shape, mx.bool_)
98
+ if valid_mask.shape != failure_rate.shape:
99
+ raise ValueError("Commit validity mask must match failure rate.")
100
+ risk = mx.clip(failure_rate.astype(mx.float32), 0.0, 1.0) * valid_mask.astype(mx.float32)
101
+ allowed = (mx.cumsum(risk, axis=-1) < failure_budget) & (
102
+ mx.cumprod(valid_mask.astype(mx.int32), axis=-1).astype(mx.bool_)
103
+ )
104
+ return mx.sum(mx.cumprod(allowed.astype(mx.int32), axis=-1), axis=-1)
105
+
106
+
107
+ def first_committed_token_lengths(
108
+ proposal: mx.array,
109
+ commit_lengths: mx.array,
110
+ token_id: int | Sequence[int],
111
+ *,
112
+ positions: mx.array | None = None,
113
+ ) -> mx.array:
114
+ if proposal.ndim != 2 or commit_lengths.shape != proposal.shape[:1]:
115
+ raise ValueError("Proposal and commit lengths must share a batch dimension.")
116
+ if positions is None:
117
+ positions = mx.arange(proposal.shape[1])[None, :]
118
+ elif positions.shape != (1, proposal.shape[1]):
119
+ raise ValueError("Commit positions must have shape [1, canvas].")
120
+ matches = _terminal_mask(proposal, token_id) & (positions < commit_lengths[:, None])
121
+ first = mx.min(mx.where(matches, positions, proposal.shape[1]), axis=-1)
122
+ return mx.minimum(mx.where(first < proposal.shape[1], first + 1, commit_lengths), commit_lengths)
123
+
124
+
125
+ def bounded_prefix_failure_commit_lengths(
126
+ committed_token_ids: mx.array,
127
+ failure_rate: mx.array,
128
+ *,
129
+ failure_budget: float,
130
+ remaining_lengths: mx.array,
131
+ stop_token_id: int | Sequence[int],
132
+ valid_mask: mx.array | None = None,
133
+ positions: mx.array | None = None,
134
+ ) -> mx.array:
135
+ if committed_token_ids.shape != failure_rate.shape:
136
+ raise ValueError("Committed token IDs and failure rate must share [batch, canvas].")
137
+ if remaining_lengths.shape != committed_token_ids.shape[:1]:
138
+ raise ValueError("Remaining lengths must have shape [batch].")
139
+ lengths = prefix_failure_commit_lengths(
140
+ failure_rate, failure_budget=failure_budget, valid_mask=valid_mask
141
+ )
142
+ lengths = mx.minimum(lengths, mx.maximum(remaining_lengths, 0))
143
+ return first_committed_token_lengths(
144
+ committed_token_ids, lengths, stop_token_id, positions=positions
145
+ )
146
+
147
+
148
+ @dataclass(frozen=True)
149
+ class MLXCommitPolicyDecision:
150
+ normal_lengths: mx.array
151
+ commit_lengths: mx.array
152
+ commit_token_ids: mx.array
153
+ jump_rows: mx.array
154
+ ponder_steps: mx.array
155
+ stagnation_steps: mx.array
156
+
157
+
158
+ def select_commit_lengths(
159
+ sampled_token_ids: mx.array,
160
+ normal_failure_rate: mx.array,
161
+ previous_failure_rate: mx.array,
162
+ greedy_token_ids: mx.array,
163
+ jump_failure_rate: mx.array,
164
+ *,
165
+ ponder_steps: mx.array,
166
+ stagnation_steps: mx.array,
167
+ active_rows: mx.array,
168
+ remaining_lengths: mx.array,
169
+ failure_budget: float,
170
+ stop_token_id: int | Sequence[int],
171
+ stagnation_threshold: int,
172
+ min_progress: float,
173
+ max_ponder_steps: int | None = None,
174
+ valid_mask: mx.array | None = None,
175
+ ) -> MLXCommitPolicyDecision:
176
+ """Select normal commits or bounded greedy JUMP independently per row."""
177
+
178
+ if not (sampled_token_ids.shape == normal_failure_rate.shape
179
+ == previous_failure_rate.shape == greedy_token_ids.shape
180
+ == jump_failure_rate.shape):
181
+ raise ValueError("Sampled and greedy statistics must share [batch, canvas].")
182
+ if not (ponder_steps.shape == stagnation_steps.shape == active_rows.shape
183
+ == remaining_lengths.shape == sampled_token_ids.shape[:1]):
184
+ raise ValueError("Commit row inputs must share [batch].")
185
+ if min_progress < 0:
186
+ raise ValueError("Minimum progress must be nonnegative.")
187
+ canvas = normal_failure_rate.shape[1]
188
+ positions = mx.arange(canvas)[None, :]
189
+ normal = bounded_prefix_failure_commit_lengths(
190
+ sampled_token_ids, normal_failure_rate,
191
+ failure_budget=failure_budget, remaining_lengths=remaining_lengths,
192
+ stop_token_id=stop_token_id, valid_mask=valid_mask, positions=positions,
193
+ )
194
+ previous = prefix_failure_commit_lengths(
195
+ previous_failure_rate, failure_budget=failure_budget, valid_mask=valid_mask
196
+ )
197
+ frontier = mx.maximum(previous, normal) + 1
198
+ valid_lengths = (mx.sum(valid_mask.astype(mx.int32), axis=-1) if valid_mask is not None
199
+ else mx.full(frontier.shape, canvas, mx.int32))
200
+ frontier = mx.minimum(frontier, valid_lengths)
201
+ progress_mask = (positions < frontier[:, None]) & active_rows[:, None]
202
+ if valid_mask is not None:
203
+ progress_mask = progress_mask & valid_mask
204
+ weights = progress_mask.astype(mx.float32)
205
+ progress = mx.sum((previous_failure_rate.astype(mx.float32)
206
+ - normal_failure_rate.astype(mx.float32)) * weights, axis=-1) / mx.maximum(
207
+ mx.sum(weights, axis=-1), 1.0)
208
+ waiting = active_rows & (normal == 0)
209
+ next_ponder = mx.where(normal > 0, 0, ponder_steps + waiting.astype(mx.int32))
210
+ next_stagnation = mx.where(normal > 0, 0, stagnation_steps + waiting.astype(mx.int32))
211
+ jump = (normal == 0) & active_rows & (next_stagnation >= stagnation_threshold)
212
+ jump = jump & (progress <= min_progress)
213
+ if max_ponder_steps is not None and max_ponder_steps > 0:
214
+ jump = jump | ((normal == 0) & active_rows & (next_ponder >= max_ponder_steps))
215
+ jump_commit = bounded_prefix_failure_commit_lengths(
216
+ greedy_token_ids, jump_failure_rate,
217
+ failure_budget=JUMP_FAILURE_BUDGET, remaining_lengths=remaining_lengths,
218
+ stop_token_id=stop_token_id, valid_mask=valid_mask, positions=positions,
219
+ )
220
+ commit_token_ids = mx.where(jump[:, None], greedy_token_ids, sampled_token_ids)
221
+ committed = mx.where(active_rows, mx.where(jump, jump_commit, normal), 0)
222
+ jump = jump & (committed > 0)
223
+ next_ponder = mx.where(committed > 0, 0, next_ponder).astype(mx.int32)
224
+ next_stagnation = mx.where(committed > 0, 0, next_stagnation).astype(mx.int32)
225
+ return MLXCommitPolicyDecision(
226
+ normal_lengths=normal,
227
+ commit_lengths=committed,
228
+ commit_token_ids=commit_token_ids,
229
+ jump_rows=jump,
230
+ ponder_steps=next_ponder,
231
+ stagnation_steps=next_stagnation,
232
+ )
233
+
234
+
235
+ def infer_commit_reason(
236
+ commit_lengths: mx.array,
237
+ *,
238
+ jump_rows: mx.array | None = None,
239
+ commit_token_ids: mx.array | None = None,
240
+ terminal_token_ids: Sequence[int] = (),
241
+ ) -> mx.array:
242
+ """Reason codes consumed by the persistent writer at commit only."""
243
+
244
+ committed = commit_lengths > 0
245
+ reasons = mx.where(committed, COMMIT_REASON_NORMAL, COMMIT_REASON_NONE)
246
+ if jump_rows is not None:
247
+ reasons = mx.where(committed & jump_rows, COMMIT_REASON_FORCED_JUMP, reasons)
248
+ if commit_token_ids is not None and terminal_token_ids:
249
+ positions = mx.arange(commit_token_ids.shape[1])[None, :]
250
+ terminal = mx.any(
251
+ _terminal_mask(commit_token_ids, terminal_token_ids)
252
+ & (positions < commit_lengths[:, None]), axis=-1,
253
+ )
254
+ reasons = mx.where(committed & terminal, COMMIT_REASON_TERMINAL, reasons)
255
+ return reasons.astype(mx.int32)
modilify_mk2/mlx_gdn2_memory.py ADDED
@@ -0,0 +1,117 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """MLX counterpart of the FP32 GDN2 trajectory recurrence."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import mlx.core as mx
6
+ from mlx import nn
7
+
8
+
9
+ def _l2(value: mx.array) -> mx.array:
10
+ return value * mx.rsqrt(mx.maximum(mx.sum(mx.square(value), axis=-1, keepdims=True), 1e-12))
11
+
12
+
13
+ class _GateProjection(nn.Module):
14
+ def __init__(self, source: int, target: int, rank: int) -> None:
15
+ super().__init__()
16
+ self.down = nn.Linear(source, rank, bias=False)
17
+ self.up = nn.Linear(rank, target, bias=False)
18
+
19
+ def __call__(self, source: mx.array) -> mx.array:
20
+ return self.up(self.down(source))
21
+
22
+
23
+ class GDN2Memory(nn.Module):
24
+ def __init__(self, input_dim: int, heads: int, key_dim: int, value_dim: int,
25
+ *, observation_dim: int | None = None) -> None:
26
+ super().__init__()
27
+ if min(input_dim, heads, key_dim, value_dim) <= 0:
28
+ raise ValueError("GDN2 dimensions must be positive.")
29
+ self.heads, self.key_dim, self.value_dim = heads, key_dim, value_dim
30
+ observation_dim = input_dim if observation_dim is None else observation_dim
31
+ if observation_dim <= 0:
32
+ raise ValueError("GDN2 observation width must be positive.")
33
+ self.q_proj = nn.Linear(input_dim, heads * key_dim, bias=False)
34
+ self.k_proj = nn.Linear(observation_dim, heads * key_dim, bias=False)
35
+ self.v_proj = nn.Linear(observation_dim, heads * value_dim, bias=False)
36
+ self.f_proj = _GateProjection(observation_dim, heads * key_dim, min(observation_dim, key_dim))
37
+ self.b_proj = nn.Linear(observation_dim, heads * key_dim, bias=False)
38
+ self.w_proj = nn.Linear(observation_dim, heads * value_dim, bias=False)
39
+ self.g_proj = _GateProjection(input_dim, heads * value_dim, min(input_dim, value_dim))
40
+ self.o_proj = nn.Linear(heads * value_dim, input_dim, bias=False)
41
+ self.a_log = mx.zeros((heads,), mx.float32)
42
+ self.dt_bias = mx.full((heads, key_dim), -6.906255, mx.float32)
43
+
44
+ def empty(self, *leading: int) -> mx.array:
45
+ return mx.zeros((*leading, self.heads, self.key_dim, self.value_dim), mx.float32)
46
+
47
+ def _projections(self, source: mx.array):
48
+ normalized = source.astype(mx.float32)
49
+ source = (normalized * mx.rsqrt(mx.mean(mx.square(normalized), axis=-1, keepdims=True) + 1e-6)).astype(source.dtype)
50
+ shape = source.shape[:-1]
51
+ key_shape = (*shape, self.heads, self.key_dim)
52
+ value_shape = (*shape, self.heads, self.value_dim)
53
+ k = _l2(nn.silu(self.k_proj(source).astype(mx.float32)).reshape(key_shape))
54
+ v = nn.silu(self.v_proj(source).astype(mx.float32)).reshape(value_shape)
55
+ head_rate = mx.exp(self.a_log.astype(mx.float32)).reshape(
56
+ *((1,) * (source.ndim - 1)), self.heads, 1)
57
+ decay = mx.exp(-head_rate * nn.softplus(
58
+ self.f_proj(source).astype(mx.float32).reshape(key_shape) + self.dt_bias.astype(mx.float32)))
59
+ erase = mx.sigmoid(self.b_proj(source).astype(mx.float32).reshape(key_shape))
60
+ write = mx.sigmoid(self.w_proj(source).astype(mx.float32).reshape(value_shape))
61
+ return k, v, decay, erase, write
62
+
63
+ def _query(self, source: mx.array) -> mx.array:
64
+ return _l2(nn.silu(self.q_proj(source).astype(mx.float32)).reshape(
65
+ *source.shape[:-1], self.heads, self.key_dim))
66
+
67
+ def _output(self, value: mx.array, source: mx.array) -> mx.array:
68
+ gate = nn.silu(self.g_proj(source).astype(mx.float32).reshape(value.shape))
69
+ value = value * mx.rsqrt(mx.mean(mx.square(value), axis=-1, keepdims=True) + 1e-6) * gate
70
+ return self.o_proj(value.reshape(*source.shape[:-1], -1).astype(source.dtype))
71
+
72
+ def read(self, state: mx.array, source: mx.array) -> mx.array:
73
+ if state.shape != (*source.shape[:-1], self.heads, self.key_dim, self.value_dim):
74
+ raise ValueError("GDN2 state and query leading dimensions differ.")
75
+ value = mx.sum(self._query(source)[..., None] * state.astype(mx.float32), axis=-2)
76
+ return self._output(value, source)
77
+
78
+ def read_shared(self, state: mx.array, source: mx.array) -> mx.array:
79
+ if source.ndim != 3 or state.shape != (source.shape[0], self.heads, self.key_dim, self.value_dim):
80
+ raise ValueError("Shared GDN2 read requires [batch, canvas, width] queries.")
81
+ query = self._query(source).transpose(0, 2, 1, 3)
82
+ value = (query @ state.astype(mx.float32)).transpose(0, 2, 1, 3)
83
+ return self._output(value, source)
84
+
85
+ @staticmethod
86
+ def _transition(state, k, v, decay, erase, write, valid):
87
+ decayed = state.astype(mx.float32) * decay[..., None]
88
+ old = mx.sum((erase * k)[..., None] * decayed, axis=-2)
89
+ candidate = decayed + k[..., None] * (write * v - old)[..., None, :]
90
+ return candidate if valid is None else mx.where(
91
+ valid[..., None, None, None], candidate, state.astype(mx.float32))
92
+
93
+ def transition(self, state: mx.array, source: mx.array,
94
+ valid: mx.array | None = None) -> mx.array:
95
+ # Do not override nn.Module.update: it installs parameters for optimizers,
96
+ # dtype conversion, and autodiff. State evolution is a separate operation.
97
+ if state.shape != (*source.shape[:-1], self.heads, self.key_dim, self.value_dim):
98
+ raise ValueError("GDN2 state and observation leading dimensions differ.")
99
+ k, v, decay, erase, write = self._projections(source)
100
+ if valid is not None:
101
+ if valid.shape != source.shape[:-1]:
102
+ raise ValueError("GDN2 valid mask must match observation rows.")
103
+ return self._transition(state, k, v, decay, erase, write, valid)
104
+
105
+ def write_sequence(self, state: mx.array, source: mx.array,
106
+ valid: mx.array) -> mx.array:
107
+ if source.ndim != 3 or valid.shape != source.shape[:2]:
108
+ raise ValueError("GDN2 sequence and mask must share [batch, length].")
109
+ if state.shape != (source.shape[0], self.heads, self.key_dim, self.value_dim):
110
+ raise ValueError("GDN2 sequence state shape differs.")
111
+ projected = self._projections(source)
112
+ for index in range(source.shape[1]):
113
+ state = self._transition(state, *(part[:, index] for part in projected), valid[:, index])
114
+ return state
115
+
116
+
117
+ __all__ = ["GDN2Memory"]
modilify_mk2/mlx_gdn2_trajectory.py ADDED
@@ -0,0 +1,87 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """MLX dual-timescale matrix state with denoise and commit lifetimes."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dataclasses import dataclass
6
+
7
+ import mlx.core as mx
8
+ from mlx import nn
9
+
10
+ from .mlx_gdn2_memory import GDN2Memory
11
+
12
+
13
+ @dataclass(frozen=True)
14
+ class GDN2TrajectoryState:
15
+ cells: mx.array
16
+ row: mx.array
17
+ persistent: mx.array
18
+ seen: mx.array
19
+
20
+
21
+ def clear_refill(self, lengths: mx.array, head: mx.array) -> "GDN2TrajectoryState":
22
+ batch, canvas = self.seen.shape
23
+ if lengths.shape != (batch,) or head.shape != (batch,):
24
+ raise ValueError("Commit lengths and canvas heads must be per row.")
25
+ physical = mx.arange(canvas)[None, :]
26
+ overwritten = ((physical - head[:, None]) % canvas) < lengths[:, None]
27
+ return GDN2TrajectoryState(
28
+ mx.where(overwritten[..., None, None, None], 0.0, self.cells),
29
+ self.row, self.persistent, self.seen & ~overwritten,
30
+ )
31
+
32
+
33
+ class GDN2TrajectoryMemory(nn.Module):
34
+ def __init__(self, width: int, *, probes: int = 4,
35
+ working_heads: int = 16, working_key: int = 64,
36
+ working_value: int = 64, persistent_heads: int = 16,
37
+ persistent_key: int = 128, persistent_value: int = 128,
38
+ persistent_observation_dim: int | None = None) -> None:
39
+ super().__init__()
40
+ if probes <= 0:
41
+ raise ValueError("Probe count must be positive.")
42
+ self.probes = probes
43
+ self.cell = GDN2Memory(width, working_heads, working_key, working_value)
44
+ self.row = GDN2Memory(width, working_heads, working_key, working_value)
45
+ self.persistent = GDN2Memory(width, persistent_heads, persistent_key, persistent_value,
46
+ observation_dim=persistent_observation_dim)
47
+ self.probe_embed = nn.Embedding(probes, width)
48
+ self.probe_embed.weight = mx.random.normal((probes, width)) * 0.02
49
+
50
+ def empty(self, batch: int, canvas: int) -> GDN2TrajectoryState:
51
+ return GDN2TrajectoryState(
52
+ self.cell.empty(batch, canvas), self.row.empty(batch),
53
+ self.persistent.empty(batch), mx.zeros((batch, canvas), mx.bool_),
54
+ )
55
+
56
+ def read(self, state: GDN2TrajectoryState, query: mx.array) -> mx.array:
57
+ result = (self.cell.read(state.cells, query)
58
+ + self.row.read_shared(state.row, query)
59
+ + self.persistent.read_shared(state.persistent, query))
60
+ return mx.where(state.seen[..., None], result, 0.0)
61
+
62
+ def observe(self, state: GDN2TrajectoryState, observation: mx.array,
63
+ live: mx.array, head: mx.array) -> GDN2TrajectoryState:
64
+ batch, canvas, width = observation.shape
65
+ if live.shape != (batch, canvas) or head.shape != (batch,):
66
+ raise ValueError("Observation mask and head have incorrect shapes.")
67
+ cells = self.cell.transition(state.cells, observation, live)
68
+ seen = state.seen | live
69
+ logical_idx = (head[:, None] + mx.arange(canvas)[None, :]) % canvas
70
+ logical = mx.take_along_axis(observation,
71
+ mx.broadcast_to(logical_idx[..., None], observation.shape), axis=1)
72
+ logical_live = mx.take_along_axis(live, mx.stop_gradient(logical_idx), axis=1)
73
+ row_state = state.row
74
+ for probe in range(self.probes):
75
+ lo = canvas * probe // self.probes
76
+ hi = canvas * (probe + 1) // self.probes
77
+ selected = logical_live[:, lo:hi]
78
+ count = mx.sum(selected.astype(mx.float32), axis=1, keepdims=True)
79
+ pooled = mx.sum(logical[:, lo:hi].astype(mx.float32)
80
+ * selected[..., None], axis=1)
81
+ pooled = (pooled / mx.maximum(count, 1.0)).astype(observation.dtype)
82
+ pooled = pooled + self.probe_embed.weight[probe].astype(observation.dtype)
83
+ row_state = self.row.transition(row_state, pooled, count[:, 0] > 0)
84
+ return GDN2TrajectoryState(cells, row_state, state.persistent, seen)
85
+
86
+
87
+ __all__ = ["GDN2TrajectoryMemory", "GDN2TrajectoryState"]