ydy9038074 commited on
Commit
53d5244
·
verified ·
1 Parent(s): 59e207f

Publish Modilify Mk2 Preview

Browse files
.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
 
 
 
 
 
 
 
 
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,264 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Modilify Mk2 — Notices
2
+
3
+ Copyright 2026 Modilify
4
+
5
+ This distribution is derived from `google/diffusiongemma-26B-A4B-it`, published
6
+ by Google DeepMind under Apache License 2.0. It retains the upstream multimodal
7
+ encoder, vision tower, vision projection, tokenizer, processor assets, and base
8
+ language-model parameters. Modilify added dual-timescale latent deliberation
9
+ and an excess-entropy confidence-and-entropy commit policy. The inference graph
10
+ restores the Gemma 4 vision tower on the encoder so text, image, and video
11
+ prompts share one rolling-canvas decoder.
12
+
13
+ The Modilify Open Model License 1.0 applies to Modilify's distribution and
14
+ original contributions. It does not erase, narrow, or replace rights and notices
15
+ applicable to upstream components. Users remain responsible for complying with
16
+ all applicable upstream terms.
17
+
18
+ - Upstream model: https://huggingface.co/google/diffusiongemma-26B-A4B-it
19
+ - Transformers project: https://github.com/huggingface/transformers
20
+
21
+ The remote model implementation subclasses public DiffusionGemma interfaces in
22
+ Hugging Face Transformers, which is also distributed under Apache License 2.0.
23
+
24
+ ## Derivative Model Impact Statement Template
25
+
26
+ When distributing a derivative of Modilify Mk2, include a public impact
27
+ statement covering the following items. No separate submission to Modilify is
28
+ required.
29
+
30
+ ### Identity and modifications
31
+
32
+ - Model name, version, publisher, and contact.
33
+ - Base version.
34
+ - Material modifications, data sources, merges, quantization, or adaptation.
35
+
36
+ ### Intended and excluded uses
37
+
38
+ - Intended users and use cases.
39
+ - Explicitly excluded uses.
40
+ - Deployment context and degree of human oversight.
41
+
42
+ ### Evaluation scope
43
+
44
+ - Evaluated capabilities and datasets.
45
+ - Languages, modalities, populations, or contexts not evaluated.
46
+ - Hardware and software used.
47
+
48
+ ### Known limitations and foreseeable risks
49
+
50
+ - Reliability limitations.
51
+ - Safety, bias, privacy, security, and misuse risks.
52
+ - High-risk decisions the model must not make autonomously.
53
+
54
+ ### Mitigations and monitoring
55
+
56
+ - Technical and organizational safeguards.
57
+ - Human review, appeal, and correction mechanisms.
58
+ - Monitoring, incident response, and update policy.
59
+
60
+ ## Apache License 2.0
61
+
62
+ The complete license text applicable to the upstream components follows.
63
+
64
+ Apache License
65
+ Version 2.0, January 2004
66
+ http://www.apache.org/licenses/
67
+
68
+ TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
69
+
70
+ 1. Definitions.
71
+
72
+ "License" shall mean the terms and conditions for use, reproduction,
73
+ and distribution as defined by Sections 1 through 9 of this document.
74
+
75
+ "Licensor" shall mean the copyright owner or entity authorized by
76
+ the copyright owner that is granting the License.
77
+
78
+ "Legal Entity" shall mean the union of the acting entity and all
79
+ other entities that control, are controlled by, or are under common
80
+ control with that entity. For the purposes of this definition,
81
+ "control" means (i) the power, direct or indirect, to cause the
82
+ direction or management of such entity, whether by contract or
83
+ otherwise, or (ii) ownership of fifty percent (50%) or more of the
84
+ outstanding shares, or (iii) beneficial ownership of such entity.
85
+
86
+ "You" (or "Your") shall mean an individual or Legal Entity
87
+ exercising permissions granted by this License.
88
+
89
+ "Source" form shall mean the preferred form for making modifications,
90
+ including but not limited to software source code, documentation
91
+ source, and configuration files.
92
+
93
+ "Object" form shall mean any form resulting from mechanical
94
+ transformation or translation of a Source form, including but
95
+ not limited to compiled object code, generated documentation,
96
+ and conversions to other media types.
97
+
98
+ "Work" shall mean the work of authorship, whether in Source or
99
+ Object form, made available under the License, as indicated by a
100
+ copyright notice that is included in or attached to the work
101
+ (an example is provided in the Appendix below).
102
+
103
+ "Derivative Works" shall mean any work, whether in Source or Object
104
+ form, that is based on (or derived from) the Work and for which the
105
+ editorial revisions, annotations, elaborations, or other modifications
106
+ represent, as a whole, an original work of authorship. For the purposes
107
+ of this License, Derivative Works shall not include works that remain
108
+ separable from, or merely link (or bind by name) to the interfaces of,
109
+ the Work and Derivative Works thereof.
110
+
111
+ "Contribution" shall mean any work of authorship, including
112
+ the original version of the Work and any modifications or additions
113
+ to that Work or Derivative Works thereof, that is intentionally
114
+ submitted to Licensor for inclusion in the Work by the copyright owner
115
+ or by an individual or Legal Entity authorized to submit on behalf of
116
+ the copyright owner. For the purposes of this definition, "submitted"
117
+ means any form of electronic, verbal, or written communication sent
118
+ to the Licensor or its representatives, including but not limited to
119
+ communication on electronic mailing lists, source code control systems,
120
+ and issue tracking systems that are managed by, or on behalf of, the
121
+ Licensor for the purpose of discussing and improving the Work, but
122
+ excluding communication that is conspicuously marked or otherwise
123
+ designated in writing by the copyright owner as "Not a Contribution."
124
+
125
+ "Contributor" shall mean Licensor and any individual or Legal Entity
126
+ on behalf of whom a Contribution has been received by Licensor and
127
+ subsequently incorporated within the Work.
128
+
129
+ 2. Grant of Copyright License. Subject to the terms and conditions of
130
+ this License, each Contributor hereby grants to You a perpetual,
131
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
132
+ copyright license to reproduce, prepare Derivative Works of,
133
+ publicly display, publicly perform, sublicense, and distribute the
134
+ Work and such Derivative Works in Source or Object form.
135
+
136
+ 3. Grant of Patent License. Subject to the terms and conditions of
137
+ this License, each Contributor hereby grants to You a perpetual,
138
+ worldwide, non-exclusive, no-charge, royalty-free, irrevocable
139
+ (except as stated in this section) patent license to make, have made,
140
+ use, offer to sell, sell, import, and otherwise transfer the Work,
141
+ where such license applies only to those patent claims licensable
142
+ by such Contributor that are necessarily infringed by their
143
+ Contribution(s) alone or by combination of their Contribution(s)
144
+ with the Work to which such Contribution(s) was submitted. If You
145
+ institute patent litigation against any entity (including a
146
+ cross-claim or counterclaim in a lawsuit) alleging that the Work
147
+ or a Contribution incorporated within the Work constitutes direct
148
+ or contributory patent infringement, then any patent licenses
149
+ granted to You under this License for that Work shall terminate
150
+ as of the date such litigation is filed.
151
+
152
+ 4. Redistribution. You may reproduce and distribute copies of the
153
+ Work or Derivative Works thereof in any medium, with or without
154
+ modifications, and in Source or Object form, provided that You
155
+ meet the following conditions:
156
+
157
+ (a) You must give any other recipients of the Work or
158
+ Derivative Works a copy of this License; and
159
+
160
+ (b) You must cause any modified files to carry prominent notices
161
+ stating that You changed the files; and
162
+
163
+ (c) You must retain, in the Source form of any Derivative Works
164
+ that You distribute, all copyright, patent, trademark, and
165
+ attribution notices from the Source form of the Work,
166
+ excluding those notices that do not pertain to any part of
167
+ the Derivative Works; and
168
+
169
+ (d) If the Work includes a "NOTICE" text file as part of its
170
+ distribution, then any Derivative Works that You distribute must
171
+ include a readable copy of the attribution notices contained
172
+ within such NOTICE file, excluding those notices that do not
173
+ pertain to any part of the Derivative Works, in at least one
174
+ of the following places: within a NOTICE text file distributed
175
+ as part of the Derivative Works; within the Source form or
176
+ documentation, if provided along with the Derivative Works; or,
177
+ within a display generated by the Derivative Works, if and
178
+ wherever such third-party notices normally appear. The contents
179
+ of the NOTICE file are for informational purposes only and
180
+ do not modify the License. You may add Your own attribution
181
+ notices within Derivative Works that You distribute, alongside
182
+ or as an addendum to the NOTICE text from the Work, provided
183
+ that such additional attribution notices cannot be construed
184
+ as modifying the License.
185
+
186
+ You may add Your own copyright statement to Your modifications and
187
+ may provide additional or different license terms and conditions
188
+ for use, reproduction, or distribution of Your modifications, or
189
+ for any such Derivative Works as a whole, provided Your use,
190
+ reproduction, and distribution of the Work otherwise complies with
191
+ the conditions stated in this License.
192
+
193
+ 5. Submission of Contributions. Unless You explicitly state otherwise,
194
+ any Contribution intentionally submitted for inclusion in the Work
195
+ by You to the Licensor shall be under the terms and conditions of
196
+ this License, without any additional terms or conditions.
197
+ Notwithstanding the above, nothing herein shall supersede or modify
198
+ the terms of any separate license agreement you may have executed
199
+ with Licensor regarding such Contributions.
200
+
201
+ 6. Trademarks. This License does not grant permission to use the trade
202
+ names, trademarks, service marks, or product names of the Licensor,
203
+ except as required for reasonable and customary use in describing the
204
+ origin of the Work and reproducing the content of the NOTICE file.
205
+
206
+ 7. Disclaimer of Warranty. Unless required by applicable law or
207
+ agreed to in writing, Licensor provides the Work (and each
208
+ Contributor provides its Contributions) on an "AS IS" BASIS,
209
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
210
+ implied, including, without limitation, any warranties or conditions
211
+ of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
212
+ PARTICULAR PURPOSE. You are solely responsible for determining the
213
+ appropriateness of using or redistributing the Work and assume any
214
+ risks associated with Your exercise of permissions under this License.
215
+
216
+ 8. Limitation of Liability. In no event and under no legal theory,
217
+ whether in tort (including negligence), contract, or otherwise,
218
+ unless required by applicable law (such as deliberate and grossly
219
+ negligent acts) or agreed to in writing, shall any Contributor be
220
+ liable to You for damages, including any direct, indirect, special,
221
+ incidental, or consequential damages of any character arising as a
222
+ result of this License or out of the use or inability to use the
223
+ Work (including but not limited to damages for loss of goodwill,
224
+ work stoppage, computer failure or malfunction, or any and all
225
+ other commercial damages or losses), even if such Contributor
226
+ has been advised of the possibility of such damages.
227
+
228
+ 9. Accepting Warranty or Additional Liability. While redistributing
229
+ the Work or Derivative Works thereof, You may choose to offer,
230
+ and charge a fee for, acceptance of support, warranty, indemnity,
231
+ or other liability obligations and/or rights consistent with this
232
+ License. However, in accepting such obligations, You may act only
233
+ on Your own behalf and on Your sole responsibility, not on behalf
234
+ of any other Contributor, and only if You agree to indemnify,
235
+ defend, and hold each Contributor harmless for any liability
236
+ incurred by, or claims asserted against, such Contributor by reason
237
+ of your accepting any such warranty or additional liability.
238
+
239
+ END OF TERMS AND CONDITIONS
240
+
241
+ APPENDIX: How to apply the Apache License to your work.
242
+
243
+ To apply the Apache License to your work, attach the following
244
+ boilerplate notice, with the fields enclosed by brackets "[]"
245
+ replaced with your own identifying information. (Don't include
246
+ the brackets!) The text should be enclosed in the appropriate
247
+ comment syntax for the file format. We also recommend that a
248
+ file or class name and description of purpose be included on the
249
+ same "printed page" as the copyright notice for easier
250
+ identification within third-party archives.
251
+
252
+ Copyright [yyyy] [name of copyright owner]
253
+
254
+ Licensed under the Apache License, Version 2.0 (the "License");
255
+ you may not use this file except in compliance with the License.
256
+ You may obtain a copy of the License at
257
+
258
+ http://www.apache.org/licenses/LICENSE-2.0
259
+
260
+ Unless required by applicable law or agreed to in writing, software
261
+ distributed under the License is distributed on an "AS IS" BASIS,
262
+ WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
263
+ See the License for the specific language governing permissions and
264
+ limitations under the License.
README.md ADDED
@@ -0,0 +1,238 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: other
3
+ license_name: modilify-open-model-license-1.0
4
+ license_link: LICENSE
5
+ library_name: transformers
6
+ pipeline_tag: image-text-to-text
7
+ tags:
8
+ - diffusion
9
+ - multimodal
10
+ - image-text-to-text
11
+ - mixture-of-experts
12
+ - trust-remote-code
13
+ ---
14
+
15
+ ![LOGO](assets/01-LOGO.jpg)
16
+
17
+ # Modilify Mk2 Preview
18
+
19
+ A 26B-A4B multimodal block-diffusion model with dual-timescale latent deliberation.
20
+
21
+ Mk2 does not dump chain-of-thought into extra visible tokens. Each heavy denoise runs a latent Transformer over a packed trajectory history and a denoise-time tape, then writes persistent memory only when the canvas actually commits. Working state is recomputed every step. Persistent slots survive the rolling window. The exclusive excess-entropy commit formula still decides how many tokens lock in; temperature and failure budget are first-class inference knobs.
22
+
23
+ This repository is the first public Mk2 preview checkpoint: merged BF16 weights, remote code, processor, and tokenizer. The text trunk and dual-timescale latent stack come from schema23 training step 900 (~12.4 million adaptation tokens). The Gemma 4 vision tower is restored from DiffusionGemma so text, image, and video share one decoder. This is not a full benchmark release.
24
+
25
+ ## Architecture
26
+
27
+ The heavy trunk is DiffusionGemma 26B-A4B. Inside every denoise, a 4-layer latent Transformer reads the noisy 256-token canvas plus:
28
+
29
+ 1. **Packed history** (T=16). Four views of each canvas position's recent deliberation, projected into rank-1024 space. New canvas positions start empty; they do not inherit the previous token's thought.
30
+ 2. **Denoise tape**. Row-level probes written on a time ring. They do not shift when tokens commit.
31
+ 3. **Persistent slots** (256 × 2816). Updated only at commit by a Transformer writer. Full-attention decoder layers read working and persistent buses.
32
+
33
+ Visible tokens are the product of that loop, not the workspace. Easy prompts commit a long prefix. Hard prompts keep pondering.
34
+
35
+ The encoder is the official Gemma 4 multimodal encoder. Image and video tokens condition prefix KV the same way text does; the rolling canvas and latent stack stay text-side.
36
+
37
+ ## Model Summary
38
+
39
+ | | |
40
+ | --- | ---: |
41
+ | Architecture | Mixture-of-Experts block diffusion + dual-timescale latent Transformer |
42
+ | Total Parameters | 26.139B text trunk + 569.550M vision encoder + latent stack |
43
+ | Text Heavy-Denoise Activated Parameters | 4.159B |
44
+ | Vision Encoder | Gemma 4 Vision, 569.550M |
45
+ | Layers | 30 |
46
+ | Number of Experts | 128 |
47
+ | Selected Experts per Token | 8 |
48
+ | Vocabulary Size | 262,144 |
49
+ | Context Length | 262,144 tokens |
50
+ | Sliding Window | 1024 |
51
+ | Canvas Length | 256 |
52
+ | Latent Width | 2,816 |
53
+ | Latent Memory | 256 slots × 2,816-d, 4 layers |
54
+ | Trajectory History | 16 frames, 4 views, rank 1,024 |
55
+ | Denoise Tape | 16 probes |
56
+ | Modality | Text, Image, Video |
57
+ | Preview checkpoint | schema23 step 900 |
58
+ | Adaptation tokens | ~12.4 million |
59
+
60
+ ## Getting Started
61
+
62
+ Transformers 5.14.1 is the minimum supported version.
63
+
64
+ ```shell
65
+ pip install -U transformers torch accelerate
66
+ ```
67
+
68
+ ### Text generation
69
+
70
+ ```python
71
+ import torch
72
+ from transformers import AutoModelForMultimodalLM, AutoProcessor
73
+
74
+ model_id = "modilify/Modilify-Mk2-preview"
75
+ processor = AutoProcessor.from_pretrained(model_id, trust_remote_code=True)
76
+ model = AutoModelForMultimodalLM.from_pretrained(
77
+ model_id,
78
+ trust_remote_code=True,
79
+ dtype=torch.bfloat16,
80
+ device_map="auto",
81
+ )
82
+
83
+ messages = [{"role": "user", "content": "Explain why the sky is blue."}]
84
+ inputs = processor.apply_chat_template(
85
+ messages,
86
+ tokenize=True,
87
+ add_generation_prompt=True,
88
+ enable_thinking=False,
89
+ return_dict=True,
90
+ return_tensors="pt",
91
+ ).to(model.device)
92
+
93
+ output = model.generate(
94
+ **inputs,
95
+ max_new_tokens=256,
96
+ denoise_temperature=0.8,
97
+ commit_failure_budget=0.2,
98
+ )
99
+ new_tokens = output.sequences[:, inputs["input_ids"].shape[1]:]
100
+ print(processor.batch_decode(new_tokens, skip_special_tokens=False)[0])
101
+ ```
102
+
103
+ ### Image input
104
+
105
+ ```python
106
+ from PIL import Image
107
+
108
+ image = Image.open("example.jpg").convert("RGB")
109
+ messages = [{
110
+ "role": "user",
111
+ "content": [
112
+ {"type": "image", "image": image},
113
+ {"type": "text", "text": "Describe the image and identify uncertainty."},
114
+ ],
115
+ }]
116
+ inputs = processor.apply_chat_template(
117
+ messages,
118
+ tokenize=True,
119
+ add_generation_prompt=True,
120
+ enable_thinking=True,
121
+ return_dict=True,
122
+ return_tensors="pt",
123
+ ).to(model.device)
124
+ output = model.generate(**inputs, max_new_tokens=256)
125
+ ```
126
+
127
+ ### Video-frame input
128
+
129
+ The processor represents video as a sampled sequence of frames. The following example uses PyAV to decode a short local clip and samples at most 32 RGB frames.
130
+
131
+ ```python
132
+ import av
133
+ from PIL import Image
134
+
135
+ container = av.open("short_clip.mp4")
136
+ decoded = [Image.fromarray(frame.to_rgb().to_ndarray()) for frame in container.decode(video=0)]
137
+ stride = max(1, len(decoded) // 32)
138
+ frames = decoded[::stride][:32]
139
+
140
+ messages = [{
141
+ "role": "user",
142
+ "content": [
143
+ {"type": "video", "video": frames},
144
+ {"type": "text", "text": "Summarize the main visual events in order."},
145
+ ],
146
+ }]
147
+ inputs = processor.apply_chat_template(
148
+ messages,
149
+ tokenize=True,
150
+ add_generation_prompt=True,
151
+ enable_thinking=True,
152
+ return_dict=True,
153
+ return_tensors="pt",
154
+ ).to(model.device)
155
+ output = model.generate(**inputs, max_new_tokens=256)
156
+ ```
157
+
158
+ ## Thinking mode
159
+
160
+ The chat template controls the prompt, not the model's first generated tokens.
161
+
162
+ - `enable_thinking=True` inserts a system turn that contains `<|think|>` and still ends the prompt at `<|turn>model`.
163
+ - `enable_thinking=False` does **not** inject an empty thought channel. The prompt ends at `<|turn>model`.
164
+
165
+ The model may still open `<|channel>thought` on its own. That is generation, not a template artifact. Applications should not assume hidden reasoning is complete, correct, or appropriate to expose to end users.
166
+
167
+ ## Configurable inference
168
+
169
+ The two primary knobs are sampling temperature and the prefix failure budget. Both default to the values used in Mk2 training (`0.8` and `0.2`) and can be changed per call or on the config object.
170
+
171
+ | Parameter | Default | Meaning |
172
+ | --- | ---: | --- |
173
+ | `denoise_temperature` | 0.8 | Sampling temperature for every canvas step |
174
+ | `commit_failure_budget` | 0.2 | Cumulative prefix risk limit for normal commits |
175
+ | `jump_failure_budget` | 2.0 | Cumulative risk limit for forced jumps |
176
+ | `jump_on_no_progress_after` | 12 | Stagnation steps before a forced jump |
177
+ | `max_ponder_steps` | 64 | Watchdog multiplier per requested token |
178
+ | `min_trajectory_progress` | 0.005 | Minimum fused-risk improvement counted as progress |
179
+ | `canvas_length` | 256 | Rolling diffusion canvas length |
180
+ | `repetition_penalty` | 1.0 | Transformers-style repetition penalty |
181
+ | `turn_end_token_id` | 106 | Gemma turn terminator |
182
+
183
+ Call-site override:
184
+
185
+ ```python
186
+ output = model.generate(
187
+ **inputs,
188
+ max_new_tokens=256,
189
+ denoise_temperature=0.4,
190
+ commit_failure_budget=0.05,
191
+ )
192
+ ```
193
+
194
+ Load-time override:
195
+
196
+ ```python
197
+ from transformers import AutoConfig
198
+
199
+ config = AutoConfig.from_pretrained(model_id, trust_remote_code=True)
200
+ config.denoise_temperature = 0.4
201
+ config.commit_failure_budget = 0.05
202
+ model = AutoModelForMultimodalLM.from_pretrained(
203
+ model_id,
204
+ config=config,
205
+ trust_remote_code=True,
206
+ dtype=torch.bfloat16,
207
+ device_map="auto",
208
+ )
209
+ ```
210
+
211
+ Lower temperature and a tighter budget make the model more cautious and usually slower. Higher temperature and a looser budget commit more tokens per denoise. These are compute-control decisions, not guarantees of correctness.
212
+
213
+ Generation supports left-padded batches with independent stopping. Batch prompts of similar lengths together for the best throughput. Streaming and caller-supplied KV caches remain limited to batch size 1.
214
+
215
+ ## Evaluation status, limitations, and risks
216
+
217
+ This is a preview. It does not include a complete accuracy, robustness, calibration, fairness, or safety evaluation. Structural export checks and a small graduate-level qualitative probe, if present, do not establish fitness for use.
218
+
219
+ The model can hallucinate facts, citations, visual details, or temporal relationships; reproduce bias, unsafe content, personal information, or copyrighted material; and consume substantial time and memory during long iterative generation. Confidence-based commits control compute. They do not certify that a prefix is true. Visual performance can degrade with poor resolution, motion, occlusion, unusual aspect ratios, or domain shift.
220
+
221
+ Evaluate the exact deployment on representative, adversarial, and out-of-distribution inputs. Use layered safeguards, monitoring, incident response, and qualified human review. Never delegate autonomous high-risk medical, legal, financial, employment, housing, education, critical-infrastructure, or safety decisions to the model.
222
+
223
+ ## License
224
+
225
+ Released under the [Modilify Open Model License 1.0](LICENSE), subject to its responsible-use and derivative-impact terms. Upstream rights, attribution, Apache-2.0 text, and the impact-statement template are retained in [NOTICE.md](NOTICE.md).
226
+
227
+ ## Citation
228
+
229
+ ```bibtex
230
+ @software{modilify_mk2_preview_2026,
231
+ title = {Modilify Mk2 Preview},
232
+ author = {Modilify},
233
+ year = {2026},
234
+ note = {A multimodal dual-timescale latent-deliberation derivative of DiffusionGemma}
235
+ }
236
+ ```
237
+
238
+ Also cite the upstream DiffusionGemma release as requested by Google DeepMind.
assets/01-LOGO.jpg ADDED

Git LFS Details

  • SHA256: 4d2316525989c32b9028374439f97af3b90be87f0a7ff0666a6c35d283f83358
  • Pointer size: 131 Bytes
  • Size of remote file: 431 kB
chat_template.jinja ADDED
@@ -0,0 +1,387 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {#
2
+ Template: Modilify Canonical Chat Template
3
+ Author: Modilify
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 -#}
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 -%}
commit_policy.py ADDED
@@ -0,0 +1,305 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 Modilify
2
+ # SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.0
3
+ """Confidence-and-entropy commit policy for inference."""
4
+
5
+ from __future__ import annotations
6
+
7
+ from collections.abc import Sequence
8
+ from dataclasses import dataclass
9
+ import math
10
+
11
+ import torch
12
+
13
+ from .latent_deliberation import (
14
+ advance_trajectory_clocks,
15
+ should_force_trajectory_jump,
16
+ )
17
+
18
+ JUMP_FAILURE_BUDGET = 2.0
19
+
20
+ FUSED_EPS = 1e-6
21
+
22
+
23
+ def fused_commit_confidence(
24
+ proposal_confidence: torch.Tensor,
25
+ token_entropy: torch.Tensor,
26
+ *,
27
+ vocab_size: int = 256000,
28
+ eps: float = FUSED_EPS,
29
+ ) -> torch.Tensor:
30
+ """Fuse proposal confidence with token entropy into effective commit confidence.
31
+
32
+ Effective confidence is defined using an excess-entropy sigmoid:
33
+
34
+ p = clamp(proposal_confidence, eps, 1 - eps)
35
+ h2 = -p * log(p) - (1 - p) * log(1 - p) # binary entropy of p
36
+ excess = max(token_entropy - h2, 0)
37
+ base_fused = sigmoid(logit(p) - excess)
38
+ fused = base_fused.square()
39
+
40
+ When the token entropy equals the binary entropy implied by p, the fused
41
+ confidence equals p^2. Entropy *above* h2 pulls fused below p^2.
42
+ """
43
+
44
+ p = proposal_confidence.float().nan_to_num(0.5).clamp(min=eps, max=1.0 - eps)
45
+ H = token_entropy.float().nan_to_num(0.0).clamp(min=0.0)
46
+
47
+ # Binary entropy of p, using log1p for the (1-p) term.
48
+ h2 = -p * torch.log(p) - (1.0 - p) * torch.log1p(-p)
49
+
50
+ excess = (H - h2).clamp(min=0.0)
51
+
52
+ # logit(p) = log(p / (1-p)) = log(p) - log1p(-p)
53
+ logit_p = torch.log(p) - torch.log1p(-p)
54
+
55
+ base_fused = torch.sigmoid(logit_p - excess)
56
+ fused = base_fused.square()
57
+ return fused.clamp(min=eps, max=1.0 - eps)
58
+
59
+
60
+ def fused_commit_failure_rate(
61
+ proposal_confidence: torch.Tensor,
62
+ token_entropy: torch.Tensor,
63
+ **kwargs: object,
64
+ ) -> torch.Tensor:
65
+ """Return ``1 - effective_confidence`` from the shared fusion helper."""
66
+
67
+ return 1.0 - fused_commit_confidence(
68
+ proposal_confidence, token_entropy, **kwargs
69
+ )
70
+
71
+
72
+ @dataclass(frozen=True)
73
+ class CommitPolicyDecision:
74
+ """One inference transition from proposal to committed prefix."""
75
+
76
+ normal_lengths: torch.LongTensor
77
+ commit_lengths: torch.LongTensor
78
+ commit_token_ids: torch.LongTensor
79
+ jump_rows: torch.BoolTensor
80
+ ponder_steps: torch.IntTensor
81
+ stagnation_steps: torch.IntTensor
82
+
83
+
84
+ def prefix_failure_commit_lengths(
85
+ failure_rate: torch.Tensor,
86
+ *,
87
+ failure_budget: float,
88
+ valid_mask: torch.BoolTensor | None = None,
89
+ ) -> torch.LongTensor:
90
+ """Return the longest valid prefix satisfying ``cumsum(failure_rate) < 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 = torch.ones_like(failure_rate, dtype=torch.bool)
98
+ if valid_mask.shape != failure_rate.shape:
99
+ raise ValueError("Commit validity mask must match failure rate.")
100
+
101
+ risk = failure_rate.float().clamp(0.0, 1.0) * valid_mask.to(torch.float32)
102
+ cumulative_risk = risk.cumsum(dim=-1)
103
+ contiguous_valid = valid_mask.long().cumprod(dim=-1).bool()
104
+ allowed = cumulative_risk.lt(float(failure_budget)) & contiguous_valid
105
+ return allowed.long().cumprod(dim=-1).sum(dim=-1)
106
+
107
+
108
+ def first_committed_token_lengths(
109
+ proposal: torch.LongTensor,
110
+ commit_lengths: torch.LongTensor,
111
+ token_id: int | Sequence[int],
112
+ *,
113
+ positions: torch.LongTensor | None = None,
114
+ ) -> torch.LongTensor:
115
+ """Clip each committed prefix immediately after its first matching stop token."""
116
+
117
+ if proposal.ndim != 2 or commit_lengths.shape != proposal.shape[:1]:
118
+ raise ValueError("Proposal and commit lengths must share a batch dimension.")
119
+ if positions is None:
120
+ positions = torch.arange(proposal.shape[1], device=proposal.device).unsqueeze(0)
121
+ elif positions.shape != (1, proposal.shape[1]):
122
+ raise ValueError("Commit positions must have shape [1, canvas].")
123
+ committed = positions.lt(commit_lengths[:, None])
124
+ stop_token_ids = (
125
+ (int(token_id),)
126
+ if isinstance(token_id, int)
127
+ else tuple(dict.fromkeys(int(value) for value in token_id))
128
+ )
129
+ if not stop_token_ids:
130
+ raise ValueError("At least one stop token ID is required.")
131
+ matches = proposal.eq(stop_token_ids[0])
132
+ for value in stop_token_ids[1:]:
133
+ matches |= proposal.eq(value)
134
+ matches &= committed
135
+ sentinel = torch.full_like(positions, proposal.shape[1])
136
+ first = torch.where(matches, positions, sentinel).min(dim=-1).values
137
+ clipped = torch.where(first.lt(proposal.shape[1]), first + 1, commit_lengths)
138
+ return torch.minimum(clipped, commit_lengths)
139
+
140
+
141
+ def bounded_prefix_failure_commit_lengths(
142
+ committed_token_ids: torch.LongTensor,
143
+ failure_rate: torch.Tensor,
144
+ *,
145
+ failure_budget: float,
146
+ remaining_lengths: torch.LongTensor,
147
+ stop_token_id: int | Sequence[int],
148
+ valid_mask: torch.BoolTensor | None = None,
149
+ positions: torch.LongTensor | None = None,
150
+ ) -> torch.LongTensor:
151
+ """Apply length and stop-token bounds to the shared failure-rate policy."""
152
+
153
+ if committed_token_ids.shape != failure_rate.shape:
154
+ raise ValueError("Committed token IDs and failure rate must share [batch, canvas].")
155
+ if remaining_lengths.shape != committed_token_ids.shape[:1]:
156
+ raise ValueError("Remaining lengths must have shape [batch].")
157
+ commit_lengths = prefix_failure_commit_lengths(
158
+ failure_rate,
159
+ failure_budget=failure_budget,
160
+ valid_mask=valid_mask,
161
+ )
162
+ commit_lengths = torch.minimum(commit_lengths, remaining_lengths.clamp_min(0))
163
+ return first_committed_token_lengths(
164
+ committed_token_ids,
165
+ commit_lengths,
166
+ stop_token_id,
167
+ positions=positions,
168
+ )
169
+
170
+
171
+ def select_commit_lengths(
172
+ sampled_token_ids: torch.LongTensor,
173
+ normal_failure_rate: torch.Tensor,
174
+ previous_failure_rate: torch.Tensor,
175
+ greedy_token_ids: torch.LongTensor,
176
+ jump_failure_rate: torch.Tensor,
177
+ *,
178
+ ponder_steps: torch.Tensor,
179
+ stagnation_steps: torch.Tensor,
180
+ active_rows: torch.BoolTensor,
181
+ remaining_lengths: torch.LongTensor,
182
+ failure_budget: float,
183
+ stop_token_id: int | Sequence[int],
184
+ stagnation_threshold: int,
185
+ min_progress: float,
186
+ max_ponder_steps: int | None = None,
187
+ jump_failure_budget: float | None = None,
188
+ valid_mask: torch.BoolTensor | None = None,
189
+ ) -> CommitPolicyDecision:
190
+ """Use normal sampled commits and a fixed-budget greedy JUMP.
191
+
192
+ Progress is measured from the signed change in fused failure rate over the
193
+ frontier region (the union of the previous and current commit prefixes plus
194
+ one blocking position), not from raw confidence/entropy deltas.
195
+ """
196
+
197
+ if not (
198
+ sampled_token_ids.shape
199
+ == normal_failure_rate.shape
200
+ == previous_failure_rate.shape
201
+ == greedy_token_ids.shape
202
+ == jump_failure_rate.shape
203
+ ):
204
+ raise ValueError("Sampled and greedy statistics must share [batch, canvas].")
205
+
206
+ canvas_length = normal_failure_rate.shape[1]
207
+ positions = torch.arange(canvas_length, device=normal_failure_rate.device)[None, :]
208
+ normal = bounded_prefix_failure_commit_lengths(
209
+ sampled_token_ids,
210
+ normal_failure_rate,
211
+ failure_budget=failure_budget,
212
+ remaining_lengths=remaining_lengths,
213
+ stop_token_id=stop_token_id,
214
+ valid_mask=valid_mask,
215
+ positions=positions,
216
+ )
217
+ previous_prefix_length = prefix_failure_commit_lengths(
218
+ previous_failure_rate,
219
+ failure_budget=failure_budget,
220
+ valid_mask=valid_mask,
221
+ )
222
+ frontier_length = torch.maximum(previous_prefix_length, normal) + 1
223
+ valid_lengths = (
224
+ valid_mask.long().sum(dim=-1)
225
+ if valid_mask is not None
226
+ else torch.full_like(frontier_length, canvas_length)
227
+ )
228
+ frontier_length = torch.minimum(frontier_length, valid_lengths)
229
+ progress_mask = positions < frontier_length[:, None]
230
+ if valid_mask is not None:
231
+ progress_mask &= valid_mask
232
+ progress_mask &= active_rows[:, None]
233
+ signed_improvement = (
234
+ previous_failure_rate.float() - normal_failure_rate.float()
235
+ )
236
+ weights = progress_mask.float()
237
+ progress = (
238
+ signed_improvement * weights
239
+ ).sum(dim=-1) / weights.sum(dim=-1).clamp_min(1.0)
240
+ next_ponder, next_stagnation = advance_trajectory_clocks(
241
+ ponder_steps,
242
+ stagnation_steps,
243
+ commit_lengths=normal,
244
+ active_rows=active_rows,
245
+ progress_scores=progress,
246
+ min_progress=min_progress,
247
+ )
248
+ jump_rows = normal.eq(0) & active_rows & should_force_trajectory_jump(
249
+ next_stagnation,
250
+ progress_scores=progress,
251
+ min_progress=min_progress,
252
+ stagnation_threshold=stagnation_threshold,
253
+ ponder_steps=next_ponder,
254
+ max_ponder_steps=max_ponder_steps,
255
+ )
256
+ jump_commit = bounded_prefix_failure_commit_lengths(
257
+ greedy_token_ids,
258
+ jump_failure_rate,
259
+ failure_budget=(
260
+ JUMP_FAILURE_BUDGET if jump_failure_budget is None else float(jump_failure_budget)
261
+ ),
262
+ remaining_lengths=remaining_lengths,
263
+ stop_token_id=stop_token_id,
264
+ valid_mask=valid_mask,
265
+ positions=positions,
266
+ )
267
+ committed = torch.where(jump_rows, jump_commit, normal)
268
+ commit_token_ids = torch.where(
269
+ jump_rows[:, None],
270
+ greedy_token_ids,
271
+ sampled_token_ids,
272
+ )
273
+ committed = torch.where(active_rows, committed, 0)
274
+ committed_rows = committed.gt(0)
275
+ jump_rows &= committed_rows
276
+ next_ponder = torch.where(
277
+ committed_rows,
278
+ 0,
279
+ next_ponder,
280
+ ).to(torch.int32)
281
+ next_stagnation = torch.where(
282
+ committed_rows,
283
+ 0,
284
+ next_stagnation,
285
+ ).to(torch.int32)
286
+ return CommitPolicyDecision(
287
+ normal_lengths=normal,
288
+ commit_lengths=committed,
289
+ commit_token_ids=commit_token_ids,
290
+ jump_rows=jump_rows,
291
+ ponder_steps=next_ponder,
292
+ stagnation_steps=next_stagnation,
293
+ )
294
+
295
+
296
+ __all__ = [
297
+ "CommitPolicyDecision",
298
+ "JUMP_FAILURE_BUDGET",
299
+ "bounded_prefix_failure_commit_lengths",
300
+ "first_committed_token_lengths",
301
+ "fused_commit_confidence",
302
+ "fused_commit_failure_rate",
303
+ "prefix_failure_commit_lengths",
304
+ "select_commit_lengths",
305
+ ]
config.json ADDED
@@ -0,0 +1,177 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "ModilifyMk2ForBlockDiffusion"
4
+ ],
5
+ "auto_map": {
6
+ "AutoConfig": "configuration_modilify_mk2.ModilifyMk2Config",
7
+ "AutoModel": "modeling_modilify_mk2.ModilifyMk2Model",
8
+ "AutoModelForCausalLM": "modeling_modilify_mk2.ModilifyMk2ForBlockDiffusion",
9
+ "AutoModelForMultimodalLM": "modeling_modilify_mk2.ModilifyMk2ForBlockDiffusion"
10
+ },
11
+ "boi_token_id": 255999,
12
+ "bos_token_id": 2,
13
+ "canvas_length": 256,
14
+ "channel_end_token_id": 101,
15
+ "commit_failure_budget": 0.2,
16
+ "commit_sequence_dim": 1024,
17
+ "commit_sequence_layers": 2,
18
+ "denoise_temperature": 0.8,
19
+ "dtype": "bfloat16",
20
+ "eoi_token_id": 258882,
21
+ "eos_token_id": [
22
+ 1,
23
+ 106
24
+ ],
25
+ "experience_roles": 3,
26
+ "image_token_id": 258880,
27
+ "initializer_range": 0.02,
28
+ "jump_failure_budget": 2.0,
29
+ "jump_on_no_progress_after": 12,
30
+ "kv_cache_bucket_size": 128,
31
+ "latent_dim": 2816,
32
+ "latent_dropout": 0.0,
33
+ "latent_ffn_dim": 7168,
34
+ "latent_history_kv_rank": 1024,
35
+ "latent_history_length": 16,
36
+ "latent_history_views": 4,
37
+ "latent_local_attention_window": 128,
38
+ "latent_memory_slots": 256,
39
+ "latent_num_heads": 16,
40
+ "latent_num_layers": 4,
41
+ "latent_tape_probes": 16,
42
+ "latent_working_last_block_global": true,
43
+ "max_ponder_steps": 64,
44
+ "memory_scheme": "dual_timescale_transformer_trajectory_memory",
45
+ "min_trajectory_progress": 0.005,
46
+ "model_type": "modilify_mk2",
47
+ "pad_token_id": 0,
48
+ "persistent_memory_bus": true,
49
+ "persistent_memory_write": "commit_only_transformer",
50
+ "repetition_penalty": 1.0,
51
+ "state_schema_version": 23,
52
+ "terminal_token_ids": [
53
+ 106,
54
+ 50
55
+ ],
56
+ "text_config": {
57
+ "attention_bias": false,
58
+ "attention_dropout": 0.0,
59
+ "bos_token_id": 2,
60
+ "dtype": "bfloat16",
61
+ "eos_token_id": 1,
62
+ "final_logit_softcapping": 30.0,
63
+ "global_head_dim": 512,
64
+ "head_dim": 256,
65
+ "hidden_activation": "gelu_pytorch_tanh",
66
+ "hidden_size": 2816,
67
+ "initializer_range": 0.02,
68
+ "intermediate_size": 2112,
69
+ "layer_types": [
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
+ "sliding_attention",
95
+ "sliding_attention",
96
+ "sliding_attention",
97
+ "sliding_attention",
98
+ "sliding_attention",
99
+ "full_attention"
100
+ ],
101
+ "max_position_embeddings": 262144,
102
+ "model_type": "modilify_mk2_text",
103
+ "moe_intermediate_size": 704,
104
+ "num_attention_heads": 16,
105
+ "num_experts": 128,
106
+ "num_global_key_value_heads": 2,
107
+ "num_hidden_layers": 30,
108
+ "num_key_value_heads": 8,
109
+ "pad_token_id": 0,
110
+ "rms_norm_eps": 1e-06,
111
+ "rope_parameters": {
112
+ "full_attention": {
113
+ "partial_rotary_factor": 0.25,
114
+ "rope_theta": 1000000.0,
115
+ "rope_type": "proportional"
116
+ },
117
+ "sliding_attention": {
118
+ "rope_theta": 10000.0,
119
+ "rope_type": "default"
120
+ }
121
+ },
122
+ "sliding_window": 1024,
123
+ "tie_word_embeddings": true,
124
+ "top_k_experts": 8,
125
+ "use_bidirectional_attention": "vision",
126
+ "vocab_size": 262144
127
+ },
128
+ "tie_word_embeddings": true,
129
+ "transformers_version": "5.14.1",
130
+ "turn_end_token_id": 106,
131
+ "vision_config": {
132
+ "_name_or_path": "",
133
+ "architectures": null,
134
+ "attention_bias": false,
135
+ "attention_dropout": 0.0,
136
+ "chunk_size_feed_forward": 0,
137
+ "default_output_length": 280,
138
+ "dtype": "bfloat16",
139
+ "global_head_dim": 72,
140
+ "head_dim": 72,
141
+ "hidden_activation": "gelu_pytorch_tanh",
142
+ "hidden_size": 1152,
143
+ "id2label": {
144
+ "0": "LABEL_0",
145
+ "1": "LABEL_1"
146
+ },
147
+ "initializer_range": 0.02,
148
+ "intermediate_size": 4304,
149
+ "is_encoder_decoder": false,
150
+ "label2id": {
151
+ "LABEL_0": 0,
152
+ "LABEL_1": 1
153
+ },
154
+ "max_position_embeddings": 131072,
155
+ "model_type": "gemma4_vision",
156
+ "num_attention_heads": 16,
157
+ "num_hidden_layers": 27,
158
+ "num_key_value_heads": 16,
159
+ "output_attentions": false,
160
+ "output_hidden_states": false,
161
+ "patch_size": 16,
162
+ "pooling_kernel_size": 3,
163
+ "position_embedding_size": 10240,
164
+ "problem_type": null,
165
+ "return_dict": true,
166
+ "rms_norm_eps": 1e-06,
167
+ "rope_parameters": {
168
+ "rope_theta": 100.0,
169
+ "rope_type": "default"
170
+ },
171
+ "standardize": true,
172
+ "use_clipped_linears": false
173
+ },
174
+ "vocab_chunk_size": 32768,
175
+ "working_memory_bus": true,
176
+ "writer_slot_gate": "per_slot"
177
+ }
configuration_modilify_mk2.py ADDED
@@ -0,0 +1,323 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 Modilify
2
+ # SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.0
3
+ """Inference configuration for Modilify Mk2."""
4
+
5
+ from __future__ import annotations
6
+
7
+ import math
8
+ from collections.abc import Sequence
9
+ from typing import Any
10
+
11
+ from transformers.models.diffusion_gemma import (
12
+ DiffusionGemmaConfig,
13
+ DiffusionGemmaTextConfig,
14
+ )
15
+
16
+
17
+ MEMORY_SCHEME = "dual_timescale_transformer_trajectory_memory"
18
+ HISTORY_VIEWS = 4
19
+ EXPERIENCE_ROLES = 3
20
+ COMMIT_SEQUENCE_LAYERS = 2
21
+ DENOISE_TEMPERATURE = 0.8
22
+ COMMIT_FAILURE_BUDGET = 0.2
23
+ JUMP_FAILURE_BUDGET = 2.0
24
+ VOCAB_CHUNK_SIZE = 32_768
25
+
26
+ _DROPPED_TRAINING_FIELDS = frozenset(
27
+ {
28
+ "training_scheme",
29
+ "training_bptt_steps",
30
+ "training_prefix_cache",
31
+ "token_loss_weight",
32
+ "jump_token_loss_weight",
33
+ "confidence_calibration_loss_weight",
34
+ "denoise_improvement_loss_weight",
35
+ "commit_throughput_softness",
36
+ "commit_throughput_target_tpd",
37
+ "commit_throughput_loss_weight",
38
+ "terminal_stop_loss_weight",
39
+ "terminal_stop_target_probability",
40
+ "latent_working_bus_unfreeze_steps",
41
+ "latent_persistent_bus_unfreeze_steps",
42
+ "virtual_commit_chunk_sizes",
43
+ "virtual_commit_chunk_probs",
44
+ "decoder_checkpoint_policy",
45
+ "fused_entropy_weight",
46
+ "schema_version",
47
+ "sampler_entropy_bound",
48
+ "loss_transition_steps",
49
+ "initial_jump_token_loss_weight",
50
+ "initial_confidence_calibration_loss_weight",
51
+ "initial_denoise_improvement_loss_weight",
52
+ "readiness_loss_weight",
53
+ "commit_risk_budget",
54
+ "router_load_balance_weight",
55
+ "router_z_loss_weight",
56
+ "latent_refinement_steps",
57
+ "latent_memory_bus_unfreeze_steps",
58
+ }
59
+ )
60
+
61
+
62
+ class ModilifyMk2TextConfig(DiffusionGemmaTextConfig):
63
+ """Text configuration for the Modilify Mk2 decoder."""
64
+
65
+ model_type = "modilify_mk2_text"
66
+
67
+
68
+ class ModilifyMk2Config(DiffusionGemmaConfig):
69
+ """Multimodal inference configuration for dual-timescale Modilify Mk2.
70
+
71
+ Args:
72
+ text_config: DiffusionGemma text configuration or its serialized form.
73
+ vision_config: Gemma 4 vision configuration or its serialized form.
74
+ denoise_temperature: Sampling temperature used at every denoising step.
75
+ commit_failure_budget: Maximum cumulative failure risk for normal commits.
76
+ jump_failure_budget: Maximum cumulative failure risk for forced jumps.
77
+ canvas_length: Rolling diffusion canvas length.
78
+ latent_dim: Width of the working trajectory state.
79
+ latent_ffn_dim: Feed-forward width inside the latent Transformer.
80
+ latent_memory_slots: Number of persistent latent memory slots.
81
+ latent_num_layers: Number of latent Transformer blocks.
82
+ latent_num_heads: Number of latent attention heads.
83
+ latent_local_attention_window: Local token-attention radius.
84
+ latent_dropout: Latent Transformer dropout probability.
85
+ latent_history_length: Packed per-token trajectory history length.
86
+ latent_tape_probes: Denoise-time tape probes per frame.
87
+ jump_on_no_progress_after: Stagnation steps before a forced jump.
88
+ max_ponder_steps: Maximum denoising iterations per requested token.
89
+ min_trajectory_progress: Minimum fused-risk improvement counted as progress.
90
+ turn_end_token_id: Native Gemma turn terminator.
91
+ repetition_penalty: Transformers-style repetition penalty. ``1.0`` disables
92
+ it.
93
+ kwargs: Standard DiffusionGemma configuration values.
94
+ """
95
+
96
+ model_type = "modilify_mk2"
97
+ sub_configs = {
98
+ "text_config": ModilifyMk2TextConfig,
99
+ **{
100
+ key: value
101
+ for key, value in DiffusionGemmaConfig.sub_configs.items()
102
+ if key != "text_config"
103
+ },
104
+ }
105
+
106
+ def __init__(
107
+ self,
108
+ text_config: (
109
+ ModilifyMk2TextConfig
110
+ | DiffusionGemmaTextConfig
111
+ | dict[str, Any]
112
+ | None
113
+ ) = None,
114
+ vision_config: Any | dict[str, Any] | None = None,
115
+ *,
116
+ denoise_temperature: float = DENOISE_TEMPERATURE,
117
+ commit_failure_budget: float = COMMIT_FAILURE_BUDGET,
118
+ jump_failure_budget: float = JUMP_FAILURE_BUDGET,
119
+ canvas_length: int = 256,
120
+ initializer_range: float = 0.02,
121
+ memory_scheme: str = MEMORY_SCHEME,
122
+ latent_dim: int = 2816,
123
+ latent_ffn_dim: int = 7168,
124
+ latent_memory_slots: int = 256,
125
+ latent_num_layers: int = 4,
126
+ latent_num_heads: int = 16,
127
+ latent_local_attention_window: int = 128,
128
+ latent_dropout: float = 0.0,
129
+ latent_history_length: int = 16,
130
+ latent_tape_probes: int = 16,
131
+ latent_history_views: int = HISTORY_VIEWS,
132
+ latent_history_kv_rank: int | None = None,
133
+ latent_working_last_block_global: bool = True,
134
+ working_memory_bus: bool = True,
135
+ persistent_memory_bus: bool = True,
136
+ persistent_memory_write: str = "commit_only_transformer",
137
+ experience_roles: int = EXPERIENCE_ROLES,
138
+ commit_sequence_layers: int = COMMIT_SEQUENCE_LAYERS,
139
+ commit_sequence_dim: int | None = None,
140
+ writer_slot_gate: str = "per_slot",
141
+ kv_cache_bucket_size: int = 128,
142
+ turn_end_token_id: int = 106,
143
+ terminal_token_ids: Sequence[int] | None = None,
144
+ channel_end_token_id: int = 101,
145
+ vocab_chunk_size: int = VOCAB_CHUNK_SIZE,
146
+ jump_on_no_progress_after: int = 12,
147
+ max_ponder_steps: int = 64,
148
+ min_trajectory_progress: float = 0.005,
149
+ repetition_penalty: float = 1.0,
150
+ state_schema_version: int = 23,
151
+ **kwargs: Any,
152
+ ) -> None:
153
+ kwargs.pop("model_type", None)
154
+ for name in _DROPPED_TRAINING_FIELDS:
155
+ kwargs.pop(name, None)
156
+ if isinstance(text_config, DiffusionGemmaTextConfig):
157
+ text_payload = text_config.to_dict()
158
+ text_payload.pop("model_type", None)
159
+ if text_payload.get("use_bidirectional_attention") in (None, False):
160
+ text_payload["use_bidirectional_attention"] = "vision"
161
+ text_config = ModilifyMk2TextConfig(**text_payload)
162
+ elif isinstance(text_config, dict):
163
+ text_payload = dict(text_config)
164
+ text_payload.pop("model_type", None)
165
+ if text_payload.get("use_bidirectional_attention") in (None, False):
166
+ text_payload["use_bidirectional_attention"] = "vision"
167
+ text_config = ModilifyMk2TextConfig(**text_payload)
168
+ elif text_config is None:
169
+ text_config = ModilifyMk2TextConfig(use_bidirectional_attention="vision")
170
+
171
+ self.denoise_temperature = float(denoise_temperature)
172
+ self.commit_failure_budget = float(commit_failure_budget)
173
+ self.jump_failure_budget = float(jump_failure_budget)
174
+ self.canvas_length = int(canvas_length)
175
+ self.memory_scheme = str(memory_scheme)
176
+ self.latent_dim = int(latent_dim)
177
+ self.latent_ffn_dim = int(latent_ffn_dim)
178
+ self.latent_memory_slots = int(latent_memory_slots)
179
+ self.latent_num_layers = int(latent_num_layers)
180
+ self.latent_num_heads = int(latent_num_heads)
181
+ self.latent_local_attention_window = int(latent_local_attention_window)
182
+ self.latent_dropout = float(latent_dropout)
183
+ self.latent_history_length = int(latent_history_length)
184
+ self.latent_tape_probes = int(latent_tape_probes)
185
+ self.latent_history_views = int(latent_history_views)
186
+ if latent_history_kv_rank is None:
187
+ rank = min(1024, self.latent_dim)
188
+ rank -= rank % max(self.latent_num_heads, 1)
189
+ if rank <= 0:
190
+ rank = self.latent_num_heads
191
+ self.latent_history_kv_rank = rank
192
+ else:
193
+ self.latent_history_kv_rank = int(latent_history_kv_rank)
194
+ self.latent_working_last_block_global = bool(latent_working_last_block_global)
195
+ self.working_memory_bus = bool(working_memory_bus)
196
+ self.persistent_memory_bus = bool(persistent_memory_bus)
197
+ self.persistent_memory_write = str(persistent_memory_write)
198
+ self.experience_roles = int(experience_roles)
199
+ self.commit_sequence_layers = int(commit_sequence_layers)
200
+ if commit_sequence_dim is None:
201
+ self.commit_sequence_dim = self.latent_history_kv_rank
202
+ else:
203
+ self.commit_sequence_dim = int(commit_sequence_dim)
204
+ self.writer_slot_gate = str(writer_slot_gate)
205
+ self.kv_cache_bucket_size = int(kv_cache_bucket_size)
206
+ self.turn_end_token_id = int(turn_end_token_id)
207
+ if terminal_token_ids is None:
208
+ self.terminal_token_ids = (int(self.turn_end_token_id),)
209
+ else:
210
+ self.terminal_token_ids = tuple(int(token_id) for token_id in terminal_token_ids)
211
+ self.channel_end_token_id = int(channel_end_token_id)
212
+ self.vocab_chunk_size = int(vocab_chunk_size)
213
+ self.jump_on_no_progress_after = int(jump_on_no_progress_after)
214
+ self.max_ponder_steps = int(max_ponder_steps)
215
+ self.min_trajectory_progress = float(min_trajectory_progress)
216
+ self.repetition_penalty = float(repetition_penalty)
217
+ self.state_schema_version = int(state_schema_version)
218
+ super().__init__(
219
+ text_config=text_config,
220
+ vision_config=vision_config,
221
+ initializer_range=initializer_range,
222
+ **kwargs,
223
+ )
224
+ self.model_type = type(self).model_type
225
+ if not hasattr(self, "eos_token_id"):
226
+ self.eos_token_id = self.text_config.eos_token_id
227
+ if not hasattr(self, "pad_token_id"):
228
+ self.pad_token_id = self.text_config.pad_token_id
229
+ if not hasattr(self, "bos_token_id"):
230
+ self.bos_token_id = self.text_config.bos_token_id
231
+ self._validate_modilify()
232
+
233
+ def _validate_modilify(self) -> None:
234
+ """Validate inference architecture and policy values."""
235
+
236
+ if self.memory_scheme != MEMORY_SCHEME:
237
+ raise ValueError(
238
+ f"`memory_scheme` must be {MEMORY_SCHEME!r}."
239
+ )
240
+ policy_values = (
241
+ self.denoise_temperature,
242
+ self.commit_failure_budget,
243
+ self.jump_failure_budget,
244
+ self.min_trajectory_progress,
245
+ self.repetition_penalty,
246
+ )
247
+ if any(not math.isfinite(value) for value in policy_values):
248
+ raise ValueError("Modilify Mk2 policy values must be finite.")
249
+ positive = (
250
+ self.denoise_temperature,
251
+ self.commit_failure_budget,
252
+ self.jump_failure_budget,
253
+ self.canvas_length,
254
+ self.latent_dim,
255
+ self.latent_ffn_dim,
256
+ self.latent_memory_slots,
257
+ self.latent_num_layers,
258
+ self.latent_num_heads,
259
+ self.latent_local_attention_window,
260
+ self.latent_history_length,
261
+ self.latent_tape_probes,
262
+ self.latent_history_kv_rank,
263
+ self.commit_sequence_layers,
264
+ self.commit_sequence_dim,
265
+ self.kv_cache_bucket_size,
266
+ self.vocab_chunk_size,
267
+ self.jump_on_no_progress_after,
268
+ self.max_ponder_steps,
269
+ self.repetition_penalty,
270
+ )
271
+ if any(value <= 0 for value in positive):
272
+ raise ValueError(
273
+ "Modilify Mk2 dimensions, budgets, intervals, and "
274
+ "`repetition_penalty` must be positive."
275
+ )
276
+ if self.latent_history_views != HISTORY_VIEWS:
277
+ raise ValueError(f"`latent_history_views` must be {HISTORY_VIEWS}.")
278
+ if self.experience_roles != EXPERIENCE_ROLES:
279
+ raise ValueError(f"`experience_roles` must be {EXPERIENCE_ROLES}.")
280
+ if self.commit_sequence_dim % self.latent_num_heads:
281
+ raise ValueError("`commit_sequence_dim` must be divisible by `latent_num_heads`.")
282
+ if self.commit_sequence_dim != self.latent_history_kv_rank:
283
+ raise ValueError("`commit_sequence_dim` must equal `latent_history_kv_rank`.")
284
+ if self.persistent_memory_write != "commit_only_transformer":
285
+ raise ValueError("`persistent_memory_write` must be commit_only_transformer.")
286
+ if self.writer_slot_gate != "per_slot":
287
+ raise ValueError("`writer_slot_gate` must be per_slot.")
288
+ if self.latent_history_kv_rank > self.latent_dim:
289
+ raise ValueError("`latent_history_kv_rank` must not exceed `latent_dim`.")
290
+ if self.latent_history_kv_rank % self.latent_num_heads:
291
+ raise ValueError("`latent_history_kv_rank` must be divisible by `latent_num_heads`.")
292
+ if not isinstance(self.channel_end_token_id, int) or self.channel_end_token_id < 0:
293
+ raise ValueError("`channel_end_token_id` must be a non-negative integer.")
294
+ if not self.terminal_token_ids:
295
+ raise ValueError("`terminal_token_ids` must not be empty.")
296
+ if any(
297
+ not isinstance(token_id, int) or token_id < 0
298
+ for token_id in self.terminal_token_ids
299
+ ):
300
+ raise ValueError("`terminal_token_ids` must be non-negative integers.")
301
+ if self.latent_dim % self.latent_num_heads:
302
+ raise ValueError("`latent_dim` must be divisible by `latent_num_heads`.")
303
+ if not 0.0 <= self.latent_dropout < 1.0:
304
+ raise ValueError("`latent_dropout` must be in [0, 1).")
305
+ if self.min_trajectory_progress < 0:
306
+ raise ValueError("`min_trajectory_progress` must be non-negative.")
307
+
308
+
309
+ ModilifyMk2Config.register_for_auto_class()
310
+
311
+
312
+ __all__ = [
313
+ "COMMIT_FAILURE_BUDGET",
314
+ "COMMIT_SEQUENCE_LAYERS",
315
+ "DENOISE_TEMPERATURE",
316
+ "EXPERIENCE_ROLES",
317
+ "HISTORY_VIEWS",
318
+ "JUMP_FAILURE_BUDGET",
319
+ "MEMORY_SCHEME",
320
+ "ModilifyMk2Config",
321
+ "ModilifyMk2TextConfig",
322
+ "VOCAB_CHUNK_SIZE",
323
+ ]
continuous_batching.py ADDED
@@ -0,0 +1,1726 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 Modilify
2
+ # SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.0
3
+ """Continuous batching for Modilify Mk2 behind the Transformers public API shape.
4
+
5
+ The upstream continuous runner is autoregressive: it persists every query in a
6
+ paged cache and emits exactly one token per request and step. ModilifyMk2 instead
7
+ denoises a transient bidirectional canvas and may accept a ragged token chunk.
8
+ This module consequently owns the request runner while preserving the public
9
+ manager lifecycle and ``GenerationOutput`` contract.
10
+
11
+ Accepted prefix K/V is stored per request without padding. Every heavy denoise
12
+ step creates a temporary, left-padded batched cache view. The decoder only
13
+ reads that view, so padding can never become persistent or evict real tokens
14
+ from a sliding-window cache.
15
+ """
16
+
17
+ from __future__ import annotations
18
+
19
+ import asyncio
20
+ import copy
21
+ import hashlib
22
+ import math
23
+ import os
24
+ import queue
25
+ import threading
26
+ import time
27
+ import uuid
28
+ import warnings
29
+ from collections import defaultdict, deque
30
+ from collections.abc import Callable, Generator, Sequence
31
+ from dataclasses import asdict, dataclass, field, is_dataclass, replace
32
+ from typing import Any
33
+
34
+ import torch
35
+ from transformers.cache_utils import Cache, DynamicCache
36
+ from transformers.generation.configuration_utils import ContinuousBatchingConfig
37
+ from transformers.generation.continuous_batching.requests import (
38
+ GenerationOutput,
39
+ RequestStatus,
40
+ )
41
+
42
+ from .commit_policy import fused_commit_failure_rate, select_commit_lengths
43
+ from .generation_modilify_mk2 import (
44
+ ModilifyMk2GenerationConfig,
45
+ ModilifyMk2GenerationOutput,
46
+ ModilifyMk2RollingState,
47
+ NoiseCanvasSampler,
48
+ _add_repetition_history,
49
+ _flatten_token_ids,
50
+ build_denoise_trace_event,
51
+ deterministic_episode_iteration_bound,
52
+ )
53
+ from .latent_deliberation import (
54
+ LatentDeliberationState,
55
+ TrajectoryHistory,
56
+ cat_latent_states,
57
+ cat_trajectory_history,
58
+ cat_trajectory_tape,
59
+ empty_trajectory_tape,
60
+ infer_commit_reason,
61
+ slice_latent_state,
62
+ slice_trajectory_history,
63
+ slice_trajectory_tape,
64
+ )
65
+
66
+
67
+ _TERMINAL_REASONS = frozenset(
68
+ {
69
+ "turn_end",
70
+ "eos",
71
+ "max_new_tokens",
72
+ "max_denoising_steps",
73
+ "episode_watchdog",
74
+ "cancelled",
75
+ "error",
76
+ }
77
+ )
78
+
79
+
80
+ def continuous_config_fingerprint(
81
+ generation_config: Any,
82
+ continuous_batching_config: ContinuousBatchingConfig | None,
83
+ ) -> str:
84
+ """Return a stable-enough in-process fingerprint for persistent reuse."""
85
+
86
+ generation_payload = (
87
+ generation_config.to_dict()
88
+ if hasattr(generation_config, "to_dict")
89
+ else vars(generation_config)
90
+ )
91
+ batching = continuous_batching_config or ContinuousBatchingConfig()
92
+ batching_payload = asdict(batching) if is_dataclass(batching) else vars(batching)
93
+ return repr(
94
+ (
95
+ sorted(generation_payload.items(), key=lambda item: item[0]),
96
+ sorted(batching_payload.items(), key=lambda item: item[0]),
97
+ )
98
+ )
99
+
100
+
101
+ @dataclass
102
+ class ModilifyMk2ContinuousGenerationOutput(GenerationOutput):
103
+ """Official ``GenerationOutput`` plus ModilifyMk2 request-local diagnostics."""
104
+
105
+ stop_reason: str | None = None
106
+ committed_tokens: int = 0
107
+ denoise_steps: int = 0
108
+ no_progress_steps: int = 0
109
+ jump_count: int = 0
110
+ forced_jump_bad_count: int = 0
111
+ heavy_forward_count: int = 0
112
+ latent_context_update_count: int = 0
113
+ average_commit_len: float = 0.0
114
+ tokens_per_forward: float = 0.0
115
+ seed: int | None = None
116
+ scheduler_run_id: str | None = None
117
+ queue_seconds: float = 0.0
118
+ inference_seconds: float = 0.0
119
+ total_seconds: float = 0.0
120
+ last_step_batch_size: int = 0
121
+ is_stream_update: bool = False
122
+ delta_tokens: list[int] = field(default_factory=list)
123
+ state_shift_count: int = 0
124
+ latent_memory_norm: float = 0.0
125
+ state_retention_score: float = 0.0
126
+
127
+ def is_finished(self) -> bool:
128
+ """Treat failed/cancelled requests as terminal for every consumer API."""
129
+
130
+ return self.status in {RequestStatus.FINISHED, RequestStatus.FAILED}
131
+
132
+
133
+ @dataclass
134
+ class ModilifyMk2RequestState:
135
+ """All mutable state required to suspend and re-batch one request."""
136
+
137
+ request_id: str
138
+ prompt_ids: list[int]
139
+ max_new_tokens: int
140
+ eos_token_ids: tuple[int, ...]
141
+ streaming: bool
142
+ record_timestamps: bool
143
+ seed: int
144
+ max_denoising_steps: int | None
145
+ trace_callback: Callable[[dict[str, object]], None] | None = None
146
+ created_time: float = field(default_factory=time.perf_counter)
147
+ status: RequestStatus = RequestStatus.PENDING
148
+ started_time: float = -1.0
149
+ finished_time: float = -1.0
150
+ generated_tokens: list[int] = field(default_factory=list)
151
+ logprobs: list[float] = field(default_factory=list)
152
+ timestamps: list[float] = field(default_factory=list)
153
+ cache: Cache | None = None
154
+ rolling_state: ModilifyMk2RollingState | None = None
155
+ repetition_history: torch.BoolTensor | None = None
156
+ generator: torch.Generator | None = None
157
+ logical_length: int = 0
158
+ max_iterations: int = 0
159
+ reserved_blocks: int = 0
160
+ denoise_steps: int = 0
161
+ jumps: int = 0
162
+ forced_jump_tokens: int = 0
163
+ shifts: int = 0
164
+ stop_reason: str | None = None
165
+ error: str | None = None
166
+ terminal_emitted: bool = False
167
+ last_step_batch_size: int = 0
168
+ last_delta_tokens: list[int] = field(default_factory=list)
169
+
170
+
171
+ def _clone_tensor_row(value: torch.Tensor, row: int) -> torch.Tensor:
172
+ return value[row : row + 1].clone()
173
+
174
+
175
+ def _slice_rolling_state(state: ModilifyMk2RollingState, row: int) -> ModilifyMk2RollingState:
176
+ selected = slice(row, row + 1)
177
+ return ModilifyMk2RollingState(
178
+ canvas=_clone_tensor_row(state.canvas, row),
179
+ confidence=_clone_tensor_row(state.confidence, row),
180
+ entropy=_clone_tensor_row(state.entropy, row),
181
+ age=_clone_tensor_row(state.age, row),
182
+ latent_state=slice_latent_state(state.latent_state, selected),
183
+ history=slice_trajectory_history(state.history, selected),
184
+ tape=slice_trajectory_tape(state.tape, selected),
185
+ )
186
+
187
+
188
+ def _pack_rolling_states(states: Sequence[ModilifyMk2RollingState]) -> ModilifyMk2RollingState:
189
+ return ModilifyMk2RollingState(
190
+ canvas=torch.cat([state.canvas for state in states], dim=0),
191
+ confidence=torch.cat([state.confidence for state in states], dim=0),
192
+ entropy=torch.cat([state.entropy for state in states], dim=0),
193
+ age=torch.cat([state.age for state in states], dim=0),
194
+ latent_state=cat_latent_states([state.latent_state for state in states]),
195
+ history=cat_trajectory_history([state.history for state in states]),
196
+ tape=cat_trajectory_tape([state.tape for state in states]),
197
+ )
198
+
199
+
200
+ class ModilifyMk2LogicalCachePool:
201
+ """Per-request hole-free cache storage with ephemeral batched read views."""
202
+
203
+ def __init__(self, model: Any, *, max_batch_tokens: int | None = None) -> None:
204
+ self.model = model
205
+ self.text_config = model.config.get_text_config(decoder=True)
206
+ self.device = model.model.decoder.embed_tokens.weight.device
207
+ self.max_batch_tokens = max_batch_tokens
208
+
209
+ def new_cache(self) -> DynamicCache:
210
+ return DynamicCache(config=self.text_config)
211
+
212
+ @torch.inference_mode()
213
+ def prefill(self, prompt_ids: Sequence[int]) -> Cache:
214
+ cache = self.new_cache()
215
+ chunk_size = self.max_batch_tokens or len(prompt_ids)
216
+ for start in range(0, len(prompt_ids), chunk_size):
217
+ stop = min(start + chunk_size, len(prompt_ids))
218
+ tokens = torch.tensor(
219
+ [list(prompt_ids[start:stop])], device=self.device, dtype=torch.long
220
+ )
221
+ mask = torch.ones(1, stop, device=self.device, dtype=torch.bool)
222
+ positions = torch.arange(
223
+ start, stop, device=self.device, dtype=torch.int32
224
+ ).unsqueeze(0)
225
+ cache = self.model.model.encoder(
226
+ input_ids=tokens,
227
+ attention_mask=mask,
228
+ past_key_values=cache,
229
+ position_ids=positions,
230
+ ).past_key_values
231
+ return cache
232
+
233
+ @torch.inference_mode()
234
+ def append(self, state: ModilifyMk2RequestState, token_ids: Sequence[int]) -> None:
235
+ if not token_ids:
236
+ return
237
+ if state.cache is None:
238
+ raise RuntimeError("Cannot append tokens before request prefill.")
239
+ tokens = torch.tensor([list(token_ids)], device=self.device, dtype=torch.long)
240
+ positions = torch.arange(
241
+ state.logical_length,
242
+ state.logical_length + tokens.shape[1],
243
+ device=self.device,
244
+ dtype=torch.int32,
245
+ ).unsqueeze(0)
246
+ mask = torch.ones(
247
+ 1,
248
+ state.logical_length + tokens.shape[1],
249
+ device=self.device,
250
+ dtype=torch.bool,
251
+ )
252
+ state.cache = self.model.model.encoder(
253
+ input_ids=tokens,
254
+ attention_mask=mask,
255
+ past_key_values=state.cache,
256
+ position_ids=positions,
257
+ ).past_key_values
258
+
259
+ def pack(
260
+ self, states: Sequence[ModilifyMk2RequestState]
261
+ ) -> tuple[DynamicCache, torch.BoolTensor, torch.LongTensor]:
262
+ if not states or any(state.cache is None for state in states):
263
+ raise ValueError("Every packed request must have an initialized cache.")
264
+ logical_lengths = torch.tensor(
265
+ [state.logical_length for state in states],
266
+ device=self.device,
267
+ dtype=torch.long,
268
+ )
269
+ maximum_length = int(logical_lengths.max())
270
+ attention_mask = torch.arange(
271
+ maximum_length, device=self.device
272
+ )[None, :].ge(maximum_length - logical_lengths[:, None])
273
+
274
+ packed = self.new_cache()
275
+ source_caches = [state.cache for state in states]
276
+ assert all(cache is not None for cache in source_caches)
277
+ if any(len(cache.layers) != len(packed.layers) for cache in source_caches):
278
+ raise RuntimeError("Request cache layer structures differ.")
279
+
280
+ for layer_index, packed_layer in enumerate(packed.layers):
281
+ source_layers = [cache.layers[layer_index] for cache in source_caches]
282
+ if any(not layer.is_initialized for layer in source_layers):
283
+ raise RuntimeError("Request cache contains an uninitialized layer.")
284
+ stored_lengths = [int(layer.keys.shape[-2]) for layer in source_layers]
285
+ maximum_stored = max(stored_lengths)
286
+
287
+ def padded(name: str) -> torch.Tensor:
288
+ values = []
289
+ for layer, stored_length in zip(source_layers, stored_lengths, strict=True):
290
+ value = getattr(layer, name)
291
+ if stored_length < maximum_stored:
292
+ padding = value.new_zeros(
293
+ value.shape[0],
294
+ value.shape[1],
295
+ maximum_stored - stored_length,
296
+ value.shape[3],
297
+ )
298
+ value = torch.cat((padding, value), dim=-2)
299
+ values.append(value)
300
+ return torch.cat(values, dim=0)
301
+
302
+ keys = padded("keys")
303
+ values = padded("values")
304
+ packed_layer.lazy_initialization(keys, values)
305
+ packed_layer.keys = keys
306
+ packed_layer.values = values
307
+ if hasattr(packed_layer, "cumulative_length"):
308
+ packed_layer.cumulative_length = maximum_length
309
+ return packed, attention_mask, logical_lengths
310
+
311
+
312
+ class ModilifyMk2ContinuousBatchingManager:
313
+ """FIFO/prefill-first continuous manager compatible with Transformers APIs."""
314
+
315
+ def __init__(
316
+ self,
317
+ model: Any,
318
+ generation_config: ModilifyMk2GenerationConfig | None,
319
+ continuous_batching_config: ContinuousBatchingConfig | None,
320
+ workload_hints: Any = None,
321
+ ) -> None:
322
+ del workload_hints
323
+ # Generation must not silently mutate the caller's train/eval mode.
324
+ # Inference mode below disables autograd without changing module-local
325
+ # dropout or other training flags.
326
+ self.model = model
327
+ self.generation_config = copy.deepcopy(
328
+ generation_config or getattr(model, "generation_config", None)
329
+ or ModilifyMk2GenerationConfig.from_model_config(model.config)
330
+ )
331
+ if not isinstance(self.generation_config, ModilifyMk2GenerationConfig):
332
+ payload = self.generation_config.to_dict()
333
+ self.generation_config = ModilifyMk2GenerationConfig(**payload)
334
+ self.continuous_batching_config = copy.deepcopy(
335
+ continuous_batching_config or ContinuousBatchingConfig()
336
+ )
337
+ self.config_fingerprint = continuous_config_fingerprint(
338
+ self.generation_config, self.continuous_batching_config
339
+ )
340
+ self._validate_config()
341
+
342
+ self.device = model.model.decoder.embed_tokens.weight.device
343
+ self.dtype = model.model.decoder.embed_tokens.weight.dtype
344
+ self.cache_pool = ModilifyMk2LogicalCachePool(
345
+ model,
346
+ max_batch_tokens=self.continuous_batching_config.max_batch_tokens,
347
+ )
348
+ self.sampler: NoiseCanvasSampler = model._prepare_sampler(
349
+ self.generation_config, model.config.canvas_length
350
+ )
351
+ self.run_id = uuid.uuid4().hex
352
+ self.warmed_up = False
353
+ self.destroyed = False
354
+
355
+ configured_requests = self.continuous_batching_config.max_requests_per_batch
356
+ self.max_requests_per_batch = int(configured_requests or 8)
357
+ max_batch_tokens = self.continuous_batching_config.max_batch_tokens
358
+ if max_batch_tokens is not None:
359
+ token_capacity = int(max_batch_tokens) // int(model.config.canvas_length)
360
+ if token_capacity < 1:
361
+ raise ValueError(
362
+ "`max_batch_tokens` must fit at least one ModilifyMk2 canvas."
363
+ )
364
+ self.max_requests_per_batch = min(
365
+ self.max_requests_per_batch, token_capacity
366
+ )
367
+ self.block_size = int(self.continuous_batching_config.block_size)
368
+ self.block_capacity = self._resolve_block_capacity()
369
+ self._base_seed = (
370
+ int(self.continuous_batching_config.seed)
371
+ if self.continuous_batching_config.seed is not None
372
+ else int(torch.initial_seed())
373
+ )
374
+
375
+ self._condition = threading.Condition(threading.RLock())
376
+ self._pending: deque[ModilifyMk2RequestState] = deque()
377
+ self._active: dict[str, ModilifyMk2RequestState] = {}
378
+ self._known_request_ids: set[str] = set()
379
+ self._cancelled: set[str] = set()
380
+ self._output_queue: queue.Queue[ModilifyMk2ContinuousGenerationOutput] = queue.Queue()
381
+ self._stashed_outputs: dict[
382
+ str, deque[ModilifyMk2ContinuousGenerationOutput]
383
+ ] = defaultdict(deque)
384
+ self._result_handlers: dict[str, tuple[Callable, asyncio.AbstractEventLoop]] = {}
385
+ self._thread: threading.Thread | None = None
386
+ self._finished = threading.Event()
387
+ self.fatal_error: BaseException | None = None
388
+ self._input_closed = False
389
+ self._hard_stop = False
390
+ self._keep_for_next_session = False
391
+ self._request_counter = 0
392
+ self._active_reserved_blocks = 0
393
+ self._stats = {
394
+ "submitted": 0,
395
+ "admitted": 0,
396
+ "completed": 0,
397
+ "failed": 0,
398
+ "cancelled": 0,
399
+ "model_steps": 0,
400
+ "generated_tokens": 0,
401
+ "max_observed_batch_size": 0,
402
+ "peak_reserved_blocks": 0,
403
+ "peak_cache_blocks": 0,
404
+ "active_slot_steps": 0,
405
+ "slot_capacity_steps": 0,
406
+ }
407
+
408
+ turn_end = self.generation_config.turn_end_token_id
409
+ self.turn_end_token_id = int(
410
+ model.config.turn_end_token_id if turn_end is None else turn_end
411
+ )
412
+ self.repetition_penalty = float(self.generation_config.repetition_penalty)
413
+ self.excluded_repetition_token_ids = _flatten_token_ids(
414
+ self.generation_config.repetition_penalty_exclude_token_ids,
415
+ self.generation_config.pad_token_id,
416
+ self.generation_config.bos_token_id,
417
+ self.generation_config.eos_token_id,
418
+ self.generation_config.turn_end_token_id,
419
+ getattr(model.config, "image_token_id", None),
420
+ )
421
+
422
+ def _validate_config(self) -> None:
423
+ config = self.continuous_batching_config
424
+ positive_optional = (
425
+ "num_blocks",
426
+ "max_batch_tokens",
427
+ "max_requests_per_batch",
428
+ )
429
+ if not isinstance(config.block_size, int) or config.block_size < 4:
430
+ raise ValueError("`block_size` must be an integer greater than or equal to 4.")
431
+ for name in positive_optional:
432
+ value = getattr(config, name)
433
+ if value is not None and (not isinstance(value, int) or value <= 0):
434
+ raise ValueError(f"`{name}` must be a positive integer when set.")
435
+ if config.max_blocks_per_request is not None and (
436
+ not isinstance(config.max_blocks_per_request, int)
437
+ or config.max_blocks_per_request < 0
438
+ ):
439
+ raise ValueError("`max_blocks_per_request` must be a non-negative integer.")
440
+ if not isinstance(config.max_queue_size, int) or config.max_queue_size < 0:
441
+ raise ValueError("`max_queue_size` must be a non-negative integer.")
442
+ if config.scheduler_type not in {"fifo", "prefill_first"}:
443
+ raise ValueError("ModilifyMk2 continuous batching supports `fifo` and `prefill_first`.")
444
+ if config.max_memory_percent is not None and not (
445
+ 0.0 < float(config.max_memory_percent) <= 1.0
446
+ ):
447
+ raise ValueError("`max_memory_percent` must be in (0, 1].")
448
+ if config.use_async_batching is True:
449
+ raise ValueError("ModilifyMk2 continuous batching currently uses synchronous model steps.")
450
+ requested_graphs = config.use_cuda_graph
451
+ if requested_graphs is True or (
452
+ isinstance(requested_graphs, tuple) and any(requested_graphs)
453
+ ):
454
+ raise ValueError("CUDA graphs are not supported by the ragged ModilifyMk2 runner.")
455
+ if config.cpu_offload_space is not None and config.cpu_offload_space > 0:
456
+ raise ValueError("CPU cache offload is not supported by the ModilifyMk2 runner.")
457
+ if int(config.default_compile_level or 0) > 0:
458
+ raise ValueError("Continuous ModilifyMk2 compilation is not supported yet.")
459
+ if config.varlen_compile_config is not None or config.decode_compile_config is not None:
460
+ raise ValueError("Continuous ModilifyMk2 compilation is not supported yet.")
461
+ if config.use_default_compile_configs is True:
462
+ raise ValueError("Continuous ModilifyMk2 compilation is not supported yet.")
463
+ if int(config.q_padding_interval_size or 0) > 0 or int(
464
+ config.kv_padding_interval_size or 0
465
+ ) > 0:
466
+ raise ValueError("Compiled continuous padding intervals are not supported.")
467
+ if config.max_cached_graphs is not None:
468
+ raise ValueError("Cached continuous graphs are not supported.")
469
+ if torch.distributed.is_available() and torch.distributed.is_initialized():
470
+ if torch.distributed.get_world_size() > 1:
471
+ raise ValueError(
472
+ "Tensor/distributed parallel continuous batching is not supported."
473
+ )
474
+ if getattr(self.model, "device_mesh", None) is not None or getattr(
475
+ self.model, "_device_mesh", None
476
+ ) is not None:
477
+ raise ValueError("Tensor-parallel continuous batching is not supported.")
478
+ # Prefix sharing would make request ownership and row-local RNG/state
479
+ # ambiguous. Normalize this optimization off rather than silently use it.
480
+ config.allow_block_sharing = False
481
+
482
+ def _available_memory_bytes(self) -> int | None:
483
+ if self.device.type == "cuda" and torch.cuda.is_available():
484
+ free, _ = torch.cuda.mem_get_info(self.device)
485
+ return int(free)
486
+ if self.device.type == "mps" and torch.backends.mps.is_available():
487
+ return max(
488
+ 0,
489
+ int(torch.mps.recommended_max_memory())
490
+ - int(torch.mps.driver_allocated_memory()),
491
+ )
492
+ if self.device.type == "cpu":
493
+ try:
494
+ import psutil
495
+
496
+ return int(psutil.virtual_memory().available)
497
+ except (ImportError, OSError, ValueError):
498
+ pass
499
+ try:
500
+ return int(os.sysconf("SC_AVPHYS_PAGES")) * int(
501
+ os.sysconf("SC_PAGE_SIZE")
502
+ )
503
+ except (OSError, TypeError, ValueError):
504
+ return None
505
+ return None
506
+
507
+ def _estimated_block_bytes(self) -> int:
508
+ config = self.model.config.text_config
509
+ layer_types = list(config.layer_types)
510
+ local_heads = int(config.num_key_value_heads)
511
+ local_dim = int(config.head_dim)
512
+ global_heads = int(
513
+ getattr(config, "num_global_key_value_heads", None) or local_heads
514
+ )
515
+ global_dim = int(getattr(config, "global_head_dim", None) or local_dim)
516
+ per_token = 0
517
+ for layer_type in layer_types:
518
+ if layer_type == "full_attention":
519
+ heads, dimension = global_heads, global_dim
520
+ else:
521
+ heads, dimension = local_heads, local_dim
522
+ per_token += 2 * heads * dimension * torch.empty((), dtype=self.dtype).element_size()
523
+ return max(1, per_token * int(self.continuous_batching_config.block_size))
524
+
525
+ def _resolve_block_capacity(self) -> int | None:
526
+ capacity = self.continuous_batching_config.num_blocks
527
+ percent = self.continuous_batching_config.max_memory_percent
528
+ available = self._available_memory_bytes()
529
+ if percent is None and capacity is None:
530
+ # Never make the default cache silently unbounded. This fraction is
531
+ # applied to currently available device/host memory after model load.
532
+ percent = 0.8
533
+ if percent is not None and available is None:
534
+ raise RuntimeError(
535
+ "Cannot infer available cache memory on this device; set `num_blocks` "
536
+ "explicitly instead of `max_memory_percent`."
537
+ )
538
+ if percent is not None and available is not None:
539
+ memory_blocks = int(
540
+ available * float(percent) / self._estimated_block_bytes()
541
+ )
542
+ capacity = memory_blocks if capacity is None else min(int(capacity), memory_blocks)
543
+ return None if capacity is None else max(0, int(capacity))
544
+
545
+ @staticmethod
546
+ def _block_footprint(reservations: Sequence[int]) -> int:
547
+ """Return persistent plus temporary packed-cache block equivalents."""
548
+
549
+ if not reservations:
550
+ return 0
551
+ return sum(reservations) + len(reservations) * max(reservations)
552
+
553
+ def _current_block_footprint(self) -> int:
554
+ return self._block_footprint(
555
+ [state.reserved_blocks for state in self._active.values()]
556
+ )
557
+
558
+ def _derive_seed(self, request_id: str) -> int:
559
+ digest = hashlib.sha256(
560
+ str(self._base_seed).encode("ascii")
561
+ + b"\0"
562
+ + request_id.encode("utf-8")
563
+ ).digest()
564
+ return int.from_bytes(digest[:8], "big") & ((1 << 63) - 1)
565
+
566
+ @property
567
+ def stats(self) -> dict[str, Any]:
568
+ with self._condition:
569
+ capacity_steps = int(self._stats["slot_capacity_steps"])
570
+ return {
571
+ "scheduler_run_id": self.run_id,
572
+ **self._stats,
573
+ "slot_utilization": (
574
+ float(self._stats["active_slot_steps"]) / capacity_steps
575
+ if capacity_steps
576
+ else 0.0
577
+ ),
578
+ "active_requests": len(self._active),
579
+ "pending_requests": len(self._pending),
580
+ "max_requests_per_batch": self.max_requests_per_batch,
581
+ "block_capacity": -1 if self.block_capacity is None else self.block_capacity,
582
+ "reserved_blocks": self._active_reserved_blocks,
583
+ "cache_blocks": self._current_block_footprint(),
584
+ }
585
+
586
+ def is_running(self) -> bool:
587
+ return self._thread is not None and self._thread.is_alive()
588
+
589
+ def warmup(self) -> None:
590
+ if self.destroyed:
591
+ raise RuntimeError("Cannot warm up a destroyed manager.")
592
+ # CUDA graphs and static-shape compilation are intentionally unsupported;
593
+ # normal eager kernels warm naturally on the first real batch.
594
+ self.warmed_up = True
595
+
596
+ def start(self) -> None:
597
+ if self._keep_for_next_session:
598
+ self._prepare_for_next_session()
599
+ with self._condition:
600
+ if self.destroyed:
601
+ raise RuntimeError("Cannot start a destroyed manager.")
602
+ if self.is_running():
603
+ return
604
+ self._finished.clear()
605
+ self.fatal_error = None
606
+ self._hard_stop = False
607
+ self._thread = threading.Thread(
608
+ target=self._run_generation_loop,
609
+ name=f"modilify_mk2-continuous-{self.run_id[:8]}",
610
+ daemon=True,
611
+ )
612
+ self._thread.start()
613
+
614
+ def join(
615
+ self,
616
+ stop_trigger_time: float | None = None,
617
+ timeout: float | None = None,
618
+ ) -> None:
619
+ """Wait for the current worker, matching the official manager lifecycle."""
620
+
621
+ del stop_trigger_time
622
+ with self._condition:
623
+ thread = self._thread
624
+ if thread is None or thread is threading.current_thread():
625
+ return
626
+ thread.join(timeout=timeout)
627
+ if thread.is_alive():
628
+ raise TimeoutError("Timed out waiting for continuous generation to stop.")
629
+
630
+ def _prepare_for_next_session(self) -> None:
631
+ """Finish an asynchronous prior stop and reopen a cached manager safely."""
632
+
633
+ with self._condition:
634
+ if not self._keep_for_next_session:
635
+ return
636
+ thread = self._thread
637
+ if thread is not None and thread.is_alive():
638
+ thread.join()
639
+ with self._condition:
640
+ if self.destroyed:
641
+ raise RuntimeError("Cannot reuse a destroyed manager.")
642
+ if self._pending or self._active:
643
+ raise RuntimeError("Cannot reuse a manager with unfinished requests.")
644
+ self._input_closed = False
645
+ self._hard_stop = False
646
+ self._keep_for_next_session = False
647
+ self.fatal_error = None
648
+ self._cancelled.clear()
649
+ self._condition.notify_all()
650
+
651
+ def close_input(self) -> None:
652
+ """Stop accepting requests and let the iterator drain all submitted work."""
653
+
654
+ with self._condition:
655
+ self._input_closed = True
656
+ self._condition.notify_all()
657
+
658
+ def stop(
659
+ self,
660
+ block: bool = True,
661
+ timeout: float | None = None,
662
+ keep_for_next_session: bool = False,
663
+ hard_stop: bool = False,
664
+ ) -> None:
665
+ with self._condition:
666
+ self._input_closed = True
667
+ self._hard_stop = bool(hard_stop)
668
+ self._keep_for_next_session = bool(keep_for_next_session)
669
+ if hard_stop:
670
+ self._cancelled.update(self._known_request_ids)
671
+ self._condition.notify_all()
672
+ thread = self._thread
673
+ if hard_stop and (thread is None or not thread.is_alive()):
674
+ self._apply_cancellations()
675
+ if block and thread is not None:
676
+ self.join(timeout=timeout)
677
+ if keep_for_next_session and not self.is_running():
678
+ with self._condition:
679
+ self._input_closed = False
680
+ self._hard_stop = False
681
+ self._keep_for_next_session = False
682
+ self.fatal_error = None
683
+
684
+ def destroy(self) -> None:
685
+ if self.destroyed:
686
+ return
687
+ self.stop(block=True, hard_stop=True)
688
+ self.destroyed = True
689
+ with self._condition:
690
+ self._pending.clear()
691
+ self._active.clear()
692
+ self._condition.notify_all()
693
+
694
+ def add_request(
695
+ self,
696
+ input_ids: list[int],
697
+ request_id: str | None = None,
698
+ max_new_tokens: int | None = None,
699
+ streaming: bool = False,
700
+ record_timestamps: bool = False,
701
+ eos_token_id: int | list[int] | None = None,
702
+ **request_kwargs: Any,
703
+ ) -> str:
704
+ if not input_ids or any(
705
+ not isinstance(token_id, int) or isinstance(token_id, bool)
706
+ for token_id in input_ids
707
+ ):
708
+ raise ValueError("`input_ids` must be a non-empty list of integer token IDs.")
709
+ seed = request_kwargs.pop("seed", None)
710
+ trace_callback = request_kwargs.pop("denoise_trace_callback", None)
711
+ max_denoising_steps = request_kwargs.pop(
712
+ "max_denoising_steps", self.generation_config.max_denoising_steps
713
+ )
714
+ if request_kwargs:
715
+ unsupported = ", ".join(sorted(request_kwargs))
716
+ raise ValueError(f"Unsupported per-request generation options: {unsupported}")
717
+ if trace_callback is not None and not callable(trace_callback):
718
+ raise TypeError("`denoise_trace_callback` must be callable.")
719
+ if trace_callback is not None and self.max_requests_per_batch > 1:
720
+ raise ValueError(
721
+ "ModilifyMk2 denoise tracing remains a batch-size-1 interface; "
722
+ "set `max_requests_per_batch=1`."
723
+ )
724
+ limit = self.generation_config.max_new_tokens if max_new_tokens is None else max_new_tokens
725
+ if not isinstance(limit, int) or limit <= 0:
726
+ raise ValueError("`max_new_tokens` must be a positive integer.")
727
+ if max_denoising_steps is not None and (
728
+ not isinstance(max_denoising_steps, int) or max_denoising_steps <= 0
729
+ ):
730
+ raise ValueError("`max_denoising_steps` must be a positive integer when set.")
731
+
732
+ with self._condition:
733
+ if self.destroyed or self._input_closed:
734
+ raise RuntimeError("Continuous batching manager is not accepting requests.")
735
+ if self.fatal_error is not None:
736
+ raise RuntimeError("Continuous batching manager has failed.") from self.fatal_error
737
+ if request_id is None:
738
+ request_id = f"req_{self._request_counter}"
739
+ self._request_counter += 1
740
+ if request_id in self._known_request_ids:
741
+ raise ValueError(f"Duplicate continuous request ID: {request_id}")
742
+ queue_limit = int(self.continuous_batching_config.max_queue_size)
743
+ deadline = time.monotonic() + 10.0
744
+ while queue_limit and len(self._pending) >= queue_limit:
745
+ if not self.is_running():
746
+ raise queue.Full(
747
+ "Continuous request queue is full; start the manager before "
748
+ "submitting more requests."
749
+ )
750
+ remaining = deadline - time.monotonic()
751
+ if remaining <= 0:
752
+ raise queue.Full("Continuous request queue remained full for 10 seconds.")
753
+ self._condition.wait(timeout=remaining)
754
+ if self.destroyed or self._input_closed:
755
+ raise RuntimeError(
756
+ "Continuous batching manager stopped while waiting for queue space."
757
+ )
758
+ if self.fatal_error is not None:
759
+ raise RuntimeError("Continuous batching manager has failed.") from self.fatal_error
760
+ # The worker can close/fail the manager while this producer is
761
+ # asleep. Recheck under the same lock immediately before append.
762
+ if self.destroyed or self._input_closed:
763
+ raise RuntimeError("Continuous batching manager is not accepting requests.")
764
+ if self.fatal_error is not None:
765
+ raise RuntimeError("Continuous batching manager has failed.") from self.fatal_error
766
+ if request_id in self._known_request_ids:
767
+ raise ValueError(f"Duplicate continuous request ID: {request_id}")
768
+ configured_eos = self.generation_config.eos_token_id if eos_token_id is None else eos_token_id
769
+ if configured_eos is None:
770
+ configured_eos = self.model.config.eos_token_id
771
+ eos_values = (
772
+ [configured_eos]
773
+ if isinstance(configured_eos, int)
774
+ else list(configured_eos or [])
775
+ )
776
+ stop_ids = tuple(
777
+ dict.fromkeys(
778
+ [self.turn_end_token_id, *(int(value) for value in eos_values if int(value) >= 0)]
779
+ )
780
+ )
781
+ resolved_seed = self._derive_seed(request_id) if seed is None else int(seed)
782
+ state = ModilifyMk2RequestState(
783
+ request_id=request_id,
784
+ prompt_ids=list(input_ids),
785
+ max_new_tokens=int(limit),
786
+ eos_token_ids=stop_ids,
787
+ streaming=bool(streaming),
788
+ record_timestamps=bool(record_timestamps),
789
+ seed=resolved_seed & ((1 << 63) - 1),
790
+ max_denoising_steps=max_denoising_steps,
791
+ trace_callback=trace_callback,
792
+ )
793
+ state.reserved_blocks = math.ceil(
794
+ (len(state.prompt_ids) + state.max_new_tokens) / self.block_size
795
+ )
796
+ self._pending.append(state)
797
+ self._known_request_ids.add(request_id)
798
+ self._stats["submitted"] += 1
799
+ self._condition.notify_all()
800
+ return request_id
801
+
802
+ def add_requests(
803
+ self,
804
+ inputs: list[list[int]],
805
+ max_new_tokens: int | None = None,
806
+ streaming: bool = False,
807
+ record_timestamps: bool = False,
808
+ **request_kwargs: Any,
809
+ ) -> list[str]:
810
+ request_ids = request_kwargs.pop("request_ids", None)
811
+ seeds = request_kwargs.pop("seeds", None)
812
+ if request_ids is not None and len(request_ids) != len(inputs):
813
+ raise ValueError("`request_ids` must contain one ID per request.")
814
+ if seeds is not None and len(seeds) != len(inputs):
815
+ raise ValueError("`seeds` must contain one seed per request.")
816
+ result = []
817
+ for index, input_ids in enumerate(inputs):
818
+ per_request = dict(request_kwargs)
819
+ if seeds is not None:
820
+ per_request["seed"] = seeds[index]
821
+ result.append(
822
+ self.add_request(
823
+ input_ids=input_ids,
824
+ request_id=None if request_ids is None else request_ids[index],
825
+ max_new_tokens=max_new_tokens,
826
+ streaming=streaming,
827
+ record_timestamps=record_timestamps,
828
+ **per_request,
829
+ )
830
+ )
831
+ return result
832
+
833
+ def cancel_request(self, request_id: str) -> None:
834
+ with self._condition:
835
+ if request_id in self._known_request_ids:
836
+ self._cancelled.add(request_id)
837
+ self._condition.notify_all()
838
+
839
+ def register_result_handler(self, request_id: str, callback: Callable) -> None:
840
+ loop = asyncio.get_running_loop()
841
+ with self._condition:
842
+ self._result_handlers[request_id] = (callback, loop)
843
+
844
+ def _pop_stashed(self, request_id: str | None):
845
+ with self._condition:
846
+ if request_id is not None:
847
+ values = self._stashed_outputs.get(request_id)
848
+ if values:
849
+ return values.popleft()
850
+ return None
851
+ for values in self._stashed_outputs.values():
852
+ if values:
853
+ return values.popleft()
854
+ return None
855
+
856
+ def _has_stashed_outputs(self) -> bool:
857
+ with self._condition:
858
+ return any(values for values in self._stashed_outputs.values())
859
+
860
+ def get_result(
861
+ self, request_id: str | None = None, timeout: float | None = None
862
+ ) -> ModilifyMk2ContinuousGenerationOutput | None:
863
+ stashed = self._pop_stashed(request_id)
864
+ if stashed is not None:
865
+ return stashed
866
+ if not self.is_running() and self._output_queue.empty():
867
+ return None
868
+ deadline = None if timeout is None else time.monotonic() + timeout
869
+ while True:
870
+ remaining = None if deadline is None else max(0.0, deadline - time.monotonic())
871
+ if remaining == 0.0:
872
+ return None
873
+ try:
874
+ output = self._output_queue.get(timeout=remaining)
875
+ except queue.Empty:
876
+ return None
877
+ if request_id is None or output.request_id == request_id:
878
+ return output
879
+ with self._condition:
880
+ self._stashed_outputs[output.request_id].append(output)
881
+
882
+ def __iter__(self) -> Generator[ModilifyMk2ContinuousGenerationOutput, None, None]:
883
+ while True:
884
+ output = self.get_result(timeout=0.05)
885
+ if output is not None:
886
+ yield output
887
+ continue
888
+ if self._finished.is_set() and self._output_queue.empty():
889
+ if not self._has_stashed_outputs():
890
+ return
891
+
892
+ def request_id_iter(
893
+ self, request_id: str
894
+ ) -> Generator[ModilifyMk2ContinuousGenerationOutput, None, None]:
895
+ while True:
896
+ output = self.get_result(request_id=request_id, timeout=0.05)
897
+ if output is not None:
898
+ yield output
899
+ if output.is_finished():
900
+ return
901
+ elif self._finished.is_set():
902
+ return
903
+
904
+ def _deliver(self, output: ModilifyMk2ContinuousGenerationOutput) -> None:
905
+ handler = None
906
+ with self._condition:
907
+ handler = self._result_handlers.get(output.request_id)
908
+ if output.is_finished():
909
+ self._result_handlers.pop(output.request_id, None)
910
+ if handler is None:
911
+ self._output_queue.put(output)
912
+ else:
913
+ callback, loop = handler
914
+ try:
915
+ loop.call_soon_threadsafe(callback, output)
916
+ except RuntimeError as error:
917
+ # A callback owner may close its event loop while a terminal
918
+ # event is in flight. Preserve the result for pull consumers
919
+ # instead of turning that client race into a worker fatality.
920
+ warnings.warn(
921
+ f"Result callback loop closed for {output.request_id}: {error!r}",
922
+ stacklevel=2,
923
+ )
924
+ self._output_queue.put(output)
925
+
926
+ def _output_for(
927
+ self,
928
+ state: ModilifyMk2RequestState,
929
+ *,
930
+ stream_update: bool = False,
931
+ delta_tokens: Sequence[int] | None = None,
932
+ ) -> ModilifyMk2ContinuousGenerationOutput:
933
+ now = time.perf_counter()
934
+ finished = state.status in {RequestStatus.FINISHED, RequestStatus.FAILED}
935
+ end = state.finished_time if finished else -1.0
936
+ shifts = max(1, state.shifts)
937
+ steps = max(1, state.denoise_steps)
938
+ return ModilifyMk2ContinuousGenerationOutput(
939
+ request_id=state.request_id,
940
+ prompt_ids=list(state.prompt_ids),
941
+ generated_tokens=list(state.generated_tokens),
942
+ logprobs=list(state.logprobs),
943
+ error=state.error,
944
+ status=state.status,
945
+ created_time=state.created_time,
946
+ lifespan=(state.started_time, end),
947
+ timestamps=(list(state.timestamps) if state.record_timestamps else None),
948
+ stop_reason=state.stop_reason,
949
+ committed_tokens=len(state.generated_tokens),
950
+ denoise_steps=state.denoise_steps,
951
+ no_progress_steps=(
952
+ 0
953
+ if state.rolling_state is None
954
+ else int(state.rolling_state.latent_state.stagnation_steps[0])
955
+ ),
956
+ jump_count=state.jumps,
957
+ forced_jump_bad_count=state.forced_jump_tokens,
958
+ heavy_forward_count=state.denoise_steps,
959
+ latent_context_update_count=state.denoise_steps,
960
+ average_commit_len=len(state.generated_tokens) / shifts,
961
+ tokens_per_forward=len(state.generated_tokens) / steps,
962
+ seed=state.seed,
963
+ scheduler_run_id=self.run_id,
964
+ queue_seconds=max(0.0, state.started_time - state.created_time),
965
+ inference_seconds=(
966
+ max(0.0, (end if finished else now) - state.started_time)
967
+ if state.started_time >= 0
968
+ else 0.0
969
+ ),
970
+ total_seconds=max(0.0, (end if finished else now) - state.created_time),
971
+ last_step_batch_size=state.last_step_batch_size,
972
+ is_stream_update=stream_update,
973
+ delta_tokens=list(
974
+ state.last_delta_tokens if delta_tokens is None else delta_tokens
975
+ ),
976
+ state_shift_count=state.shifts,
977
+ latent_memory_norm=(
978
+ 0.0
979
+ if state.rolling_state is None
980
+ else float(
981
+ state.rolling_state.latent_state.memory_slots.float()
982
+ .norm(dim=-1)
983
+ .mean()
984
+ )
985
+ ),
986
+ state_retention_score=1.0 if state.shifts else 0.0,
987
+ )
988
+
989
+ def _finish(
990
+ self,
991
+ state: ModilifyMk2RequestState,
992
+ reason: str,
993
+ error: BaseException | str | None = None,
994
+ ) -> None:
995
+ if state.terminal_emitted:
996
+ return
997
+ if reason not in _TERMINAL_REASONS:
998
+ raise ValueError(f"Unknown continuous stop reason: {reason}")
999
+ state.stop_reason = reason
1000
+ state.error = None if error is None else (str(error) if isinstance(error, str) else repr(error))
1001
+ state.status = (
1002
+ RequestStatus.FAILED
1003
+ if reason in {"cancelled", "error"} or error is not None
1004
+ else RequestStatus.FINISHED
1005
+ )
1006
+ state.finished_time = time.perf_counter()
1007
+ state.terminal_emitted = True
1008
+ self._stats["completed"] += 1
1009
+ if reason == "cancelled":
1010
+ self._stats["cancelled"] += 1
1011
+ elif error is not None:
1012
+ self._stats["failed"] += 1
1013
+ self._deliver(self._output_for(state))
1014
+
1015
+ def _fail_all_requests(self, error: BaseException) -> None:
1016
+ """Convert an unexpected worker failure into one terminal result per request."""
1017
+
1018
+ self.fatal_error = error
1019
+ with self._condition:
1020
+ pending = list(self._pending)
1021
+ active = list(self._active.values())
1022
+ self._pending.clear()
1023
+ self._active.clear()
1024
+ self._active_reserved_blocks = 0
1025
+ self._input_closed = True
1026
+ for state in [*active, *pending]:
1027
+ self._finish(state, "error", error)
1028
+ self._condition.notify_all()
1029
+
1030
+ def _request_fits(self, state: ModilifyMk2RequestState) -> bool:
1031
+ per_request_limit = self.continuous_batching_config.max_blocks_per_request
1032
+ if per_request_limit not in (None, 0) and state.reserved_blocks > per_request_limit:
1033
+ return False
1034
+ if self.block_capacity is None:
1035
+ return True
1036
+ reservations = [
1037
+ *(active.reserved_blocks for active in self._active.values()),
1038
+ state.reserved_blocks,
1039
+ ]
1040
+ return self._block_footprint(reservations) <= self.block_capacity
1041
+
1042
+ def _request_can_ever_fit(self, state: ModilifyMk2RequestState) -> bool:
1043
+ per_request_limit = self.continuous_batching_config.max_blocks_per_request
1044
+ if per_request_limit not in (None, 0) and state.reserved_blocks > per_request_limit:
1045
+ return False
1046
+ return (
1047
+ self.block_capacity is None
1048
+ or self._block_footprint([state.reserved_blocks]) <= self.block_capacity
1049
+ )
1050
+
1051
+ def _initialize_request(self, state: ModilifyMk2RequestState) -> None:
1052
+ generator = torch.Generator(device=self.device)
1053
+ generator.manual_seed(state.seed)
1054
+ state.generator = generator
1055
+ try:
1056
+ canvas = self.sampler.initialize_canvas(
1057
+ 1, self.device, generators=[generator]
1058
+ )
1059
+ except TypeError:
1060
+ canvas = self.sampler.initialize_canvas(1, self.device)
1061
+ dtype = self.model.model.decoder.embed_tokens.weight.dtype
1062
+ canvas_length = int(self.model.config.canvas_length)
1063
+ latent = LatentDeliberationState.empty(
1064
+ batch_size=1,
1065
+ canvas_length=canvas_length,
1066
+ latent_dim=self.model.config.latent_dim,
1067
+ memory_slots=self.model.config.latent_memory_slots,
1068
+ device=self.device,
1069
+ dtype=dtype,
1070
+ )
1071
+ state.rolling_state = ModilifyMk2RollingState(
1072
+ canvas=canvas,
1073
+ confidence=torch.zeros(1, canvas_length, device=self.device, dtype=torch.float32),
1074
+ entropy=torch.full(
1075
+ (1, canvas_length),
1076
+ math.log(self.model.config.text_config.vocab_size),
1077
+ device=self.device,
1078
+ dtype=torch.float32,
1079
+ ),
1080
+ age=torch.zeros(1, canvas_length, device=self.device, dtype=torch.int32),
1081
+ latent_state=latent,
1082
+ history=TrajectoryHistory.empty(
1083
+ batch_size=1,
1084
+ canvas_length=canvas_length,
1085
+ hidden_size=self.model.config.text_config.hidden_size,
1086
+ history_length=self.model.config.latent_history_length,
1087
+ device=self.device,
1088
+ dtype=dtype,
1089
+ ),
1090
+ tape=empty_trajectory_tape(
1091
+ batch_size=1,
1092
+ config=self.model.config,
1093
+ device=self.device,
1094
+ dtype=dtype,
1095
+ ),
1096
+ )
1097
+ state.cache = self.cache_pool.prefill(state.prompt_ids)
1098
+ state.logical_length = len(state.prompt_ids)
1099
+ if self.repetition_penalty != 1.0:
1100
+ state.repetition_history = torch.zeros(
1101
+ self.model.config.text_config.vocab_size,
1102
+ device=self.device,
1103
+ dtype=torch.bool,
1104
+ )
1105
+ prompt = torch.tensor([state.prompt_ids], device=self.device, dtype=torch.long)
1106
+ _add_repetition_history(
1107
+ state.repetition_history.unsqueeze(0),
1108
+ prompt,
1109
+ torch.ones_like(prompt, dtype=torch.bool),
1110
+ self.excluded_repetition_token_ids,
1111
+ )
1112
+ state.max_iterations = deterministic_episode_iteration_bound(
1113
+ torch.tensor([state.max_new_tokens]),
1114
+ max_ponder_steps=self.generation_config.max_ponder_steps,
1115
+ )
1116
+ state.started_time = time.perf_counter()
1117
+ state.status = RequestStatus.DECODING
1118
+
1119
+ def _apply_cancellations(self) -> None:
1120
+ with self._condition:
1121
+ cancelled = set(self._cancelled)
1122
+ self._cancelled.clear()
1123
+ if not cancelled:
1124
+ return
1125
+ retained = deque()
1126
+ while self._pending:
1127
+ state = self._pending.popleft()
1128
+ if state.request_id in cancelled:
1129
+ self._finish(state, "cancelled", "request cancelled")
1130
+ else:
1131
+ retained.append(state)
1132
+ self._pending = retained
1133
+ for request_id in cancelled:
1134
+ state = self._active.pop(request_id, None)
1135
+ if state is not None:
1136
+ self._active_reserved_blocks -= state.reserved_blocks
1137
+ self._finish(state, "cancelled", "request cancelled")
1138
+ self._condition.notify_all()
1139
+
1140
+ def _admit_requests(self) -> None:
1141
+ while True:
1142
+ with self._condition:
1143
+ if len(self._active) >= self.max_requests_per_batch or not self._pending:
1144
+ return
1145
+ state = self._pending[0]
1146
+ if not self._request_can_ever_fit(state):
1147
+ self._pending.popleft()
1148
+ self._finish(
1149
+ state,
1150
+ "error",
1151
+ "request exceeds continuous cache block limits",
1152
+ )
1153
+ self._condition.notify_all()
1154
+ continue
1155
+ if not self._request_fits(state):
1156
+ return
1157
+ self._pending.popleft()
1158
+ self._condition.notify_all()
1159
+ try:
1160
+ self._initialize_request(state)
1161
+ except Exception as error:
1162
+ self._finish(state, "error", error)
1163
+ continue
1164
+ with self._condition:
1165
+ if state.request_id in self._cancelled:
1166
+ self._cancelled.remove(state.request_id)
1167
+ self._finish(state, "cancelled", "request cancelled")
1168
+ continue
1169
+ self._active[state.request_id] = state
1170
+ self._active_reserved_blocks += state.reserved_blocks
1171
+ self._stats["admitted"] += 1
1172
+ self._stats["peak_reserved_blocks"] = max(
1173
+ self._stats["peak_reserved_blocks"],
1174
+ self._active_reserved_blocks,
1175
+ )
1176
+ self._stats["peak_cache_blocks"] = max(
1177
+ self._stats["peak_cache_blocks"],
1178
+ self._current_block_footprint(),
1179
+ )
1180
+ self._condition.notify_all()
1181
+
1182
+ def _select_rowwise_policy(
1183
+ self,
1184
+ states: Sequence[ModilifyMk2RequestState],
1185
+ proposal: torch.LongTensor,
1186
+ normal_failure_rate: torch.Tensor,
1187
+ previous_failure_rate: torch.Tensor,
1188
+ greedy_proposal: torch.LongTensor,
1189
+ jump_failure_rate: torch.Tensor,
1190
+ rolling: ModilifyMk2RollingState,
1191
+ ):
1192
+ decisions = []
1193
+ for row, state in enumerate(states):
1194
+ remaining = state.max_new_tokens - len(state.generated_tokens)
1195
+ decisions.append(
1196
+ select_commit_lengths(
1197
+ sampled_token_ids=proposal[row : row + 1],
1198
+ normal_failure_rate=normal_failure_rate[row : row + 1],
1199
+ previous_failure_rate=previous_failure_rate[row : row + 1],
1200
+ greedy_token_ids=greedy_proposal[row : row + 1],
1201
+ jump_failure_rate=jump_failure_rate[row : row + 1],
1202
+ ponder_steps=rolling.latent_state.ponder_steps[row : row + 1],
1203
+ stagnation_steps=rolling.latent_state.stagnation_steps[row : row + 1],
1204
+ active_rows=torch.ones(1, device=self.device, dtype=torch.bool),
1205
+ remaining_lengths=torch.tensor([remaining], device=self.device),
1206
+ failure_budget=self.generation_config.commit_failure_budget,
1207
+ jump_failure_budget=self.generation_config.jump_failure_budget,
1208
+ stop_token_id=state.eos_token_ids,
1209
+ max_ponder_steps=self.generation_config.max_ponder_steps,
1210
+ stagnation_threshold=self.generation_config.jump_on_no_progress_after,
1211
+ min_progress=self.generation_config.min_trajectory_progress,
1212
+ )
1213
+ )
1214
+ return (
1215
+ torch.cat([decision.normal_lengths for decision in decisions]),
1216
+ torch.cat([decision.commit_lengths for decision in decisions]),
1217
+ torch.cat([decision.commit_token_ids for decision in decisions]),
1218
+ torch.cat([decision.jump_rows for decision in decisions]),
1219
+ torch.cat([decision.ponder_steps for decision in decisions]),
1220
+ torch.cat([decision.stagnation_steps for decision in decisions]),
1221
+ )
1222
+
1223
+ @torch.inference_mode()
1224
+ def _run_batch_step(self, states: Sequence[ModilifyMk2RequestState]) -> list[str]:
1225
+ started = time.perf_counter()
1226
+ rolling_states = [state.rolling_state for state in states]
1227
+ if any(state is None for state in rolling_states):
1228
+ raise RuntimeError("Active request has no rolling state.")
1229
+ rolling = _pack_rolling_states(rolling_states) # type: ignore[arg-type]
1230
+ packed_cache, cache_mask, logical_lengths = self.cache_pool.pack(states)
1231
+ batch_size = len(states)
1232
+ canvas_length = int(self.model.config.canvas_length)
1233
+ decoder_positions = (
1234
+ logical_lengths[:, None]
1235
+ + torch.arange(canvas_length, device=self.device)[None, :]
1236
+ ).to(torch.int32)
1237
+ decoder_mask = torch.cat(
1238
+ (
1239
+ cache_mask,
1240
+ torch.ones(
1241
+ batch_size,
1242
+ canvas_length,
1243
+ device=self.device,
1244
+ dtype=torch.bool,
1245
+ ),
1246
+ ),
1247
+ dim=-1,
1248
+ )
1249
+ repetition_history = None
1250
+ if self.repetition_penalty != 1.0:
1251
+ repetition_history = torch.stack(
1252
+ [state.repetition_history for state in states], dim=0 # type: ignore[list-item]
1253
+ )
1254
+ generators = [state.generator for state in states]
1255
+ if any(generator is None for generator in generators):
1256
+ raise RuntimeError("Active request has no sampling generator.")
1257
+
1258
+ output = self.model(
1259
+ input_ids=None,
1260
+ past_key_values=packed_cache,
1261
+ decoder_input_ids=rolling.canvas,
1262
+ previous_confidence=rolling.confidence,
1263
+ previous_entropy=rolling.entropy,
1264
+ token_age=rolling.age,
1265
+ latent_state=rolling.latent_state,
1266
+ history=rolling.history,
1267
+ tape=rolling.tape,
1268
+ decoder_position_ids=decoder_positions,
1269
+ decoder_read_cache=True,
1270
+ decoder_attention_mask=decoder_mask,
1271
+ compact_vocab=True,
1272
+ denoise_temperature=self.generation_config.denoise_temperature,
1273
+ repetition_token_mask=repetition_history,
1274
+ repetition_penalty=self.repetition_penalty,
1275
+ sampling_generators=generators,
1276
+ )
1277
+ required = (
1278
+ output.proposal,
1279
+ output.proposal_confidence,
1280
+ output.token_entropy,
1281
+ output.greedy_proposal,
1282
+ output.greedy_confidence,
1283
+ output.next_latent_state,
1284
+ )
1285
+ if any(value is None for value in required):
1286
+ raise RuntimeError("Compact ModilifyMk2 forward did not return proposal state.")
1287
+ proposal = output.proposal
1288
+ proposal_confidence = output.proposal_confidence
1289
+ token_entropy = output.token_entropy
1290
+ greedy_proposal = output.greedy_proposal
1291
+ greedy_confidence = output.greedy_confidence
1292
+ next_canvas = proposal.clone()
1293
+ next_confidence = proposal_confidence.float()
1294
+ next_latent = replace(
1295
+ output.next_latent_state,
1296
+ confidence=next_confidence.detach().float(),
1297
+ entropy=token_entropy.detach().float(),
1298
+ age=rolling.age + 1,
1299
+ token_changed=next_canvas.ne(rolling.canvas).detach().float(),
1300
+ confidence_delta=next_confidence.detach().float() - rolling.confidence,
1301
+ entropy_delta=token_entropy.detach().float() - rolling.entropy,
1302
+ )
1303
+ live_mask = torch.ones(
1304
+ rolling.canvas.shape, device=rolling.canvas.device, dtype=torch.bool
1305
+ )
1306
+ tape_probes, tape_valid = self.model.latent_deliberation.encode_tape_frame(
1307
+ output.heavy_hidden_state, live_mask
1308
+ )
1309
+ next_state = ModilifyMk2RollingState(
1310
+ canvas=next_canvas,
1311
+ confidence=next_confidence,
1312
+ entropy=token_entropy,
1313
+ age=rolling.age + 1,
1314
+ latent_state=next_latent,
1315
+ history=rolling.history.append(
1316
+ output.heavy_hidden_state,
1317
+ next_confidence,
1318
+ token_entropy,
1319
+ next_canvas.ne(rolling.canvas).detach().float(),
1320
+ live_mask=live_mask,
1321
+ ),
1322
+ tape=rolling.tape.append(tape_probes, tape_valid),
1323
+ )
1324
+ normal_failure_rate = fused_commit_failure_rate(
1325
+ proposal_confidence,
1326
+ token_entropy,
1327
+ vocab_size=self.model.config.text_config.vocab_size,
1328
+ )
1329
+ jump_failure_rate = fused_commit_failure_rate(
1330
+ greedy_confidence,
1331
+ token_entropy,
1332
+ vocab_size=self.model.config.text_config.vocab_size,
1333
+ )
1334
+ previous_failure_rate = fused_commit_failure_rate(
1335
+ rolling.confidence,
1336
+ rolling.entropy,
1337
+ vocab_size=self.model.config.text_config.vocab_size,
1338
+ )
1339
+ (
1340
+ normal_commit,
1341
+ commit_lengths,
1342
+ commit_token_ids,
1343
+ jump_rows,
1344
+ next_ponder,
1345
+ next_stagnation,
1346
+ ) = self._select_rowwise_policy(
1347
+ states,
1348
+ proposal,
1349
+ normal_failure_rate,
1350
+ previous_failure_rate,
1351
+ greedy_proposal,
1352
+ jump_failure_rate,
1353
+ rolling,
1354
+ )
1355
+ positions = torch.arange(canvas_length, device=self.device)[None, :]
1356
+ commit_positions = positions.lt(commit_lengths[:, None])
1357
+ policy_prefix_mask = positions.lt(normal_commit[:, None])
1358
+ if bool(jump_rows.any()):
1359
+ next_state = replace(
1360
+ next_state,
1361
+ canvas=torch.where(
1362
+ commit_positions & jump_rows[:, None],
1363
+ commit_token_ids,
1364
+ next_state.canvas,
1365
+ ),
1366
+ )
1367
+ next_state = replace(
1368
+ next_state,
1369
+ latent_state=replace(
1370
+ next_state.latent_state,
1371
+ ponder_steps=next_ponder,
1372
+ stagnation_steps=next_stagnation,
1373
+ ),
1374
+ )
1375
+ unshifted_trace_states = [
1376
+ _slice_rolling_state(next_state, row) for row in range(batch_size)
1377
+ ]
1378
+ if output.history_projected is None or output.working_state is None:
1379
+ raise RuntimeError("Forward did not return working trajectory features.")
1380
+ next_state = self.model._write_committed_memory(
1381
+ previous_history=rolling.history,
1382
+ next_state=next_state,
1383
+ working_state=output.working_state,
1384
+ history_projected=output.history_projected,
1385
+ heavy_hidden=output.heavy_hidden_state,
1386
+ commit_lengths=commit_lengths,
1387
+ prefix_lengths=logical_lengths,
1388
+ commit_reason=infer_commit_reason(
1389
+ commit_lengths,
1390
+ jump_rows=jump_rows,
1391
+ commit_token_ids=commit_token_ids,
1392
+ terminal_token_ids=getattr(
1393
+ self.model.config, "terminal_token_ids", ()
1394
+ ),
1395
+ ),
1396
+ )
1397
+ shifted = self.model._shift_state_rows(
1398
+ next_state,
1399
+ commit_lengths,
1400
+ self.sampler,
1401
+ generators=generators,
1402
+ )
1403
+ shifted_states = [
1404
+ _slice_rolling_state(shifted, row) for row in range(batch_size)
1405
+ ]
1406
+ selected_confidence = torch.where(
1407
+ jump_rows[:, None], greedy_confidence, proposal_confidence
1408
+ ).float()
1409
+
1410
+ finished_ids = []
1411
+ for row, state in enumerate(states):
1412
+ state.last_step_batch_size = batch_size
1413
+ state.denoise_steps += 1
1414
+ commit_length = int(commit_lengths[row])
1415
+ chunk = commit_token_ids[row, :commit_length].detach().cpu().tolist()
1416
+ state.last_delta_tokens = [int(token_id) for token_id in chunk]
1417
+ before = len(state.generated_tokens)
1418
+ try:
1419
+ self.cache_pool.append(state, chunk)
1420
+ except Exception as error:
1421
+ self._finish(state, "error", error)
1422
+ finished_ids.append(state.request_id)
1423
+ continue
1424
+ state.logical_length += commit_length
1425
+ state.generated_tokens.extend(int(token_id) for token_id in chunk)
1426
+ if self.continuous_batching_config.return_logprobs and commit_length:
1427
+ probabilities = selected_confidence[row, :commit_length].clamp_min(
1428
+ torch.finfo(torch.float32).tiny
1429
+ )
1430
+ state.logprobs.extend(probabilities.log().detach().cpu().tolist())
1431
+ if state.record_timestamps and commit_length:
1432
+ state.timestamps.extend([time.perf_counter()] * commit_length)
1433
+ if state.repetition_history is not None and commit_length:
1434
+ tokens = commit_token_ids[row : row + 1]
1435
+ eligible = commit_positions[row : row + 1]
1436
+ _add_repetition_history(
1437
+ state.repetition_history.unsqueeze(0),
1438
+ tokens,
1439
+ eligible,
1440
+ self.excluded_repetition_token_ids,
1441
+ )
1442
+ state.jumps += int(jump_rows[row])
1443
+ if bool(jump_rows[row]):
1444
+ state.forced_jump_tokens += commit_length
1445
+ if commit_length:
1446
+ state.shifts += 1
1447
+ state.rolling_state = shifted_states[row]
1448
+
1449
+ reason = None
1450
+ if self.turn_end_token_id in chunk:
1451
+ reason = "turn_end"
1452
+ elif any(token_id in state.eos_token_ids for token_id in chunk):
1453
+ reason = "eos"
1454
+ elif len(state.generated_tokens) >= state.max_new_tokens:
1455
+ reason = "max_new_tokens"
1456
+ elif (
1457
+ state.max_denoising_steps is not None
1458
+ and state.denoise_steps >= state.max_denoising_steps
1459
+ ):
1460
+ reason = "max_denoising_steps"
1461
+ elif state.denoise_steps >= state.max_iterations:
1462
+ reason = "episode_watchdog"
1463
+
1464
+ elapsed = time.perf_counter() - started
1465
+ if state.trace_callback is not None:
1466
+ trace = build_denoise_trace_event(
1467
+ denoise_step=state.denoise_steps,
1468
+ prefix_length=state.logical_length - commit_length,
1469
+ committed_before=before,
1470
+ committed_after=len(state.generated_tokens),
1471
+ no_progress_steps=int(next_stagnation[row]),
1472
+ policy_prefix_mask=policy_prefix_mask[row : row + 1],
1473
+ commit_length=commit_length,
1474
+ ponder_fallback=bool(jump_rows[row]),
1475
+ state=unshifted_trace_states[row],
1476
+ proposal=proposal[row : row + 1],
1477
+ committed_token_ids=commit_token_ids[row : row + 1, :commit_length],
1478
+ step_elapsed_seconds=elapsed,
1479
+ latent_residual_diagnostics=None,
1480
+ )
1481
+ trace["request_id"] = state.request_id
1482
+ trace["batch_size"] = batch_size
1483
+ try:
1484
+ state.trace_callback(trace)
1485
+ except Exception as error:
1486
+ warnings.warn(
1487
+ f"Denoise trace callback failed for {state.request_id}: {error!r}",
1488
+ stacklevel=2,
1489
+ )
1490
+
1491
+ if reason is not None:
1492
+ self._finish(state, reason)
1493
+ finished_ids.append(state.request_id)
1494
+ elif state.streaming and commit_length:
1495
+ self._deliver(
1496
+ self._output_for(
1497
+ state,
1498
+ stream_update=True,
1499
+ delta_tokens=state.last_delta_tokens,
1500
+ )
1501
+ )
1502
+
1503
+ self._stats["model_steps"] += 1
1504
+ self._stats["generated_tokens"] += int(commit_lengths.sum())
1505
+ self._stats["max_observed_batch_size"] = max(
1506
+ self._stats["max_observed_batch_size"], batch_size
1507
+ )
1508
+ self._stats["active_slot_steps"] += batch_size
1509
+ self._stats["slot_capacity_steps"] += self.max_requests_per_batch
1510
+ return finished_ids
1511
+
1512
+ def _run_step_with_isolation(self, states: Sequence[ModilifyMk2RequestState]) -> None:
1513
+ generator_states = {
1514
+ state.request_id: state.generator.get_state()
1515
+ for state in states
1516
+ if state.generator is not None
1517
+ }
1518
+ try:
1519
+ finished_ids = self._run_batch_step(states)
1520
+ except Exception as batch_error:
1521
+ for state in states:
1522
+ if state.generator is not None:
1523
+ state.generator.set_state(generator_states[state.request_id])
1524
+ if len(states) == 1:
1525
+ self._finish(states[0], "error", batch_error)
1526
+ finished_ids = [states[0].request_id]
1527
+ else:
1528
+ finished_ids = []
1529
+ for state in states:
1530
+ if state.terminal_emitted:
1531
+ finished_ids.append(state.request_id)
1532
+ continue
1533
+ try:
1534
+ finished_ids.extend(self._run_batch_step([state]))
1535
+ except Exception as request_error:
1536
+ self._finish(state, "error", request_error)
1537
+ finished_ids.append(state.request_id)
1538
+ with self._condition:
1539
+ for request_id in dict.fromkeys(finished_ids):
1540
+ state = self._active.pop(request_id, None)
1541
+ if state is not None:
1542
+ self._active_reserved_blocks -= state.reserved_blocks
1543
+ self._condition.notify_all()
1544
+
1545
+ @torch.inference_mode()
1546
+ def _run_generation_loop(self) -> None:
1547
+ try:
1548
+ while True:
1549
+ self._apply_cancellations()
1550
+ if self._hard_stop:
1551
+ self._apply_cancellations()
1552
+ with self._condition:
1553
+ has_active = bool(self._active)
1554
+ # ``prefill_first`` fills every available slot before the next
1555
+ # denoise step. FIFO lets the already-active cohort take its
1556
+ # next step first, then fills slots released by that step.
1557
+ if (
1558
+ self.continuous_batching_config.scheduler_type == "prefill_first"
1559
+ or not has_active
1560
+ ):
1561
+ self._admit_requests()
1562
+ with self._condition:
1563
+ active = list(self._active.values())
1564
+ should_finish = (
1565
+ self._input_closed and not self._pending and not active
1566
+ )
1567
+ if should_finish:
1568
+ return
1569
+ if not active:
1570
+ self._condition.wait(timeout=0.05)
1571
+ continue
1572
+ self._run_step_with_isolation(active)
1573
+ if self.continuous_batching_config.scheduler_type == "fifo":
1574
+ self._admit_requests()
1575
+ except BaseException as error:
1576
+ self._fail_all_requests(error)
1577
+ finally:
1578
+ self._finished.set()
1579
+ with self._condition:
1580
+ self._condition.notify_all()
1581
+
1582
+
1583
+ @torch.inference_mode()
1584
+ def generate_static_batch_with_logical_cache(
1585
+ model: Any,
1586
+ input_ids: torch.LongTensor,
1587
+ attention_mask: torch.BoolTensor | None,
1588
+ generation_config: ModilifyMk2GenerationConfig,
1589
+ *,
1590
+ seeds: Sequence[int] | None = None,
1591
+ max_new_tokens: Sequence[int] | None = None,
1592
+ ) -> ModilifyMk2GenerationOutput:
1593
+ """Run one fixed cohort through the same hole-free continuous engine."""
1594
+
1595
+ batch_size, input_width = input_ids.shape
1596
+ if attention_mask is None:
1597
+ attention_mask = torch.ones_like(input_ids, dtype=torch.bool)
1598
+ else:
1599
+ attention_mask = attention_mask.to(device=input_ids.device, dtype=torch.bool)
1600
+ if attention_mask.shape != input_ids.shape:
1601
+ raise ValueError("`attention_mask` must have the same shape as `input_ids`.")
1602
+ prompts = [
1603
+ input_ids[row, attention_mask[row]].detach().cpu().tolist()
1604
+ for row in range(batch_size)
1605
+ ]
1606
+ if any(not prompt for prompt in prompts):
1607
+ raise ValueError("Every batched ModilifyMk2 prompt must contain at least one token.")
1608
+ if seeds is not None and len(seeds) != batch_size:
1609
+ raise ValueError("`seeds` must contain one seed per batch row.")
1610
+ if max_new_tokens is None:
1611
+ max_new_tokens = [int(generation_config.max_new_tokens)] * batch_size
1612
+ if len(max_new_tokens) != batch_size or any(
1613
+ not isinstance(limit, int) or isinstance(limit, bool) or limit <= 0
1614
+ for limit in max_new_tokens
1615
+ ):
1616
+ raise ValueError("`max_new_tokens` must contain one positive limit per batch row.")
1617
+
1618
+ batching_config = ContinuousBatchingConfig(
1619
+ block_size=max(4, int(getattr(model.config, "kv_cache_bucket_size", 128))),
1620
+ max_batch_tokens=batch_size * int(model.config.canvas_length),
1621
+ max_requests_per_batch=batch_size,
1622
+ allow_block_sharing=False,
1623
+ scheduler_type="prefill_first",
1624
+ )
1625
+ manager = ModilifyMk2ContinuousBatchingManager(
1626
+ model=model,
1627
+ generation_config=generation_config,
1628
+ continuous_batching_config=batching_config,
1629
+ )
1630
+ try:
1631
+ request_ids = []
1632
+ for row, prompt in enumerate(prompts):
1633
+ request_kwargs = {}
1634
+ if seeds is not None:
1635
+ request_kwargs["seed"] = int(seeds[row])
1636
+ request_ids.append(
1637
+ manager.add_request(
1638
+ prompt,
1639
+ request_id=f"static_{row}",
1640
+ max_new_tokens=int(max_new_tokens[row]),
1641
+ streaming=False,
1642
+ max_denoising_steps=generation_config.max_denoising_steps,
1643
+ eos_token_id=generation_config.eos_token_id,
1644
+ **request_kwargs,
1645
+ )
1646
+ )
1647
+ manager.close_input()
1648
+ # A static batch is one fixed cohort: queue every row before the worker
1649
+ # starts so its first heavy forward necessarily contains the full batch.
1650
+ manager.start()
1651
+ final = {}
1652
+ for output in manager:
1653
+ if output.is_finished():
1654
+ final[output.request_id] = output
1655
+ ordered = [final[request_id] for request_id in request_ids]
1656
+ finally:
1657
+ manager.stop(block=True, hard_stop=True)
1658
+ manager.destroy()
1659
+ failures = [output for output in ordered if output.error is not None]
1660
+ if failures:
1661
+ details = "; ".join(
1662
+ f"{output.request_id}: {output.error}" for output in failures
1663
+ )
1664
+ raise RuntimeError(f"Static ModilifyMk2 batch generation failed: {details}")
1665
+
1666
+ lengths = torch.tensor(
1667
+ [len(output.generated_tokens) for output in ordered],
1668
+ device=input_ids.device,
1669
+ dtype=torch.long,
1670
+ )
1671
+ output_width = int(lengths.max()) if lengths.numel() else 0
1672
+ pad_token_id = generation_config.pad_token_id
1673
+ if isinstance(pad_token_id, (list, tuple)):
1674
+ pad_token_id = pad_token_id[0]
1675
+ pad_token_id = int(0 if pad_token_id is None else pad_token_id)
1676
+ generated = torch.full(
1677
+ (batch_size, output_width),
1678
+ pad_token_id,
1679
+ device=input_ids.device,
1680
+ dtype=input_ids.dtype,
1681
+ )
1682
+ for row, output in enumerate(ordered):
1683
+ if output.generated_tokens:
1684
+ generated[row, : len(output.generated_tokens)] = torch.tensor(
1685
+ output.generated_tokens,
1686
+ device=input_ids.device,
1687
+ dtype=input_ids.dtype,
1688
+ )
1689
+
1690
+ def tensor(name: str, *, dtype: torch.dtype) -> torch.Tensor:
1691
+ return torch.tensor(
1692
+ [getattr(output, name) for output in ordered],
1693
+ device=input_ids.device,
1694
+ dtype=dtype,
1695
+ )
1696
+
1697
+ return ModilifyMk2GenerationOutput(
1698
+ sequences=torch.cat((input_ids, generated), dim=-1),
1699
+ generated_lengths=lengths,
1700
+ tokens_per_forward=tensor("tokens_per_forward", dtype=torch.float32),
1701
+ past_key_values=None,
1702
+ stop_reason=tuple(output.stop_reason for output in ordered),
1703
+ committed_tokens=lengths.clone(),
1704
+ denoise_steps=tensor("denoise_steps", dtype=torch.long),
1705
+ no_progress_steps=tensor("no_progress_steps", dtype=torch.long),
1706
+ jump_count=tensor("jump_count", dtype=torch.long),
1707
+ forced_jump_bad_count=tensor("forced_jump_bad_count", dtype=torch.long),
1708
+ heavy_forward_count=tensor("heavy_forward_count", dtype=torch.long),
1709
+ latent_context_update_count=tensor(
1710
+ "latent_context_update_count", dtype=torch.long
1711
+ ),
1712
+ average_commit_len=tensor("average_commit_len", dtype=torch.float32),
1713
+ state_shift_count=tensor("state_shift_count", dtype=torch.long),
1714
+ latent_memory_norm=tensor("latent_memory_norm", dtype=torch.float32),
1715
+ state_retention_score=tensor("state_retention_score", dtype=torch.float32),
1716
+ )
1717
+
1718
+
1719
+ __all__ = [
1720
+ "ModilifyMk2ContinuousBatchingManager",
1721
+ "ModilifyMk2ContinuousGenerationOutput",
1722
+ "ModilifyMk2LogicalCachePool",
1723
+ "ModilifyMk2RequestState",
1724
+ "continuous_config_fingerprint",
1725
+ "generate_static_batch_with_logical_cache",
1726
+ ]
generation_config.json ADDED
@@ -0,0 +1,24 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "commit_failure_budget": 0.2,
3
+ "confidence_threshold": null,
4
+ "denoise_temperature": 0.8,
5
+ "eos_token_id": [
6
+ 1,
7
+ 106
8
+ ],
9
+ "jump_failure_budget": 2.0,
10
+ "jump_on_no_progress_after": 12,
11
+ "max_denoising_steps": null,
12
+ "max_new_tokens": 256,
13
+ "max_ponder_steps": 64,
14
+ "min_trajectory_progress": 0.005,
15
+ "repetition_penalty": 1.0,
16
+ "repetition_penalty_exclude_token_ids": [],
17
+ "return_dict_in_generate": true,
18
+ "sampler_config": null,
19
+ "stability_threshold": null,
20
+ "t_max": 0.8,
21
+ "t_min": 0.8,
22
+ "transformers_version": "5.14.1",
23
+ "turn_end_token_id": 106
24
+ }
generation_modilify_mk2.py ADDED
@@ -0,0 +1,1455 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 Modilify
2
+ # SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.0
3
+ """Rolling generation for latent-memory Modilify Mk2."""
4
+
5
+ from __future__ import annotations
6
+
7
+ from collections.abc import Callable, Sequence
8
+ from contextlib import contextmanager
9
+ from dataclasses import dataclass, replace
10
+ import math
11
+ import time
12
+ from typing import Any
13
+
14
+ import torch
15
+ from transformers.cache_utils import Cache
16
+ from transformers.generation import LogitsProcessorList
17
+ from transformers.generation.streamers import BaseStreamer
18
+ from transformers.modeling_outputs import ModelOutput
19
+
20
+ from transformers.models.diffusion_gemma import (
21
+ DiffusionGemmaGenerationConfig,
22
+ DiffusionGemmaGenerationMixin,
23
+ )
24
+ from .commit_policy import (
25
+ first_committed_token_lengths,
26
+ fused_commit_failure_rate,
27
+ select_commit_lengths,
28
+ )
29
+ from .configuration_modilify_mk2 import (
30
+ COMMIT_FAILURE_BUDGET,
31
+ DENOISE_TEMPERATURE,
32
+ JUMP_FAILURE_BUDGET,
33
+ )
34
+ from .latent_deliberation import (
35
+ LatentDeliberationState,
36
+ TrajectoryHistory,
37
+ TrajectoryTape,
38
+ choose_trajectory_history,
39
+ choose_trajectory_tape,
40
+ empty_trajectory_tape,
41
+ infer_commit_reason,
42
+ )
43
+
44
+
45
+ _PARENT_GENERATION_KEYS = frozenset({
46
+ "max_new_tokens",
47
+ "max_length",
48
+ "return_dict_in_generate",
49
+ "max_denoising_steps",
50
+ "sampler_config",
51
+ "t_min",
52
+ "t_max",
53
+ "stability_threshold",
54
+ "confidence_threshold",
55
+ "cache_implementation",
56
+ "cache_config",
57
+ "disable_compile",
58
+ "bos_token_id",
59
+ "pad_token_id",
60
+ "eos_token_id",
61
+ "_commit_hash",
62
+ "_from_model_config",
63
+ "transformers_version",
64
+ })
65
+ _IGNORED_GENERATION_KEYS = frozenset({
66
+ "compile_generation",
67
+ "sliding_denoise",
68
+ "one_token_per_denoise_step",
69
+ "adaptive_ponder_budget",
70
+ "force_commit_on_max_steps",
71
+ "ponder_budget_id",
72
+ })
73
+ _FORBIDDEN_COMMIT_FIELDS = (
74
+ "sampler_config",
75
+ "stability_threshold",
76
+ "confidence_threshold",
77
+ "one_token_per_denoise_step",
78
+ )
79
+
80
+
81
+ def _reject_legacy_commit_fields(fields: dict[str, object]) -> None:
82
+ configured = {
83
+ name: value
84
+ for name, value in fields.items()
85
+ if value not in (None, False)
86
+ }
87
+ if configured:
88
+ raise ValueError(
89
+ "ModilifyMk2 accepts only the confidence-prefix commit policy; "
90
+ f"unsupported generation fields: {sorted(configured)}"
91
+ )
92
+
93
+
94
+ def _flatten_token_ids(*values: object) -> set[int]:
95
+ """Normalize scalar and sequence token-ID configuration values."""
96
+
97
+ token_ids: set[int] = set()
98
+ for value in values:
99
+ if value is None:
100
+ continue
101
+ if isinstance(value, int):
102
+ token_ids.add(int(value))
103
+ continue
104
+ if isinstance(value, (list, tuple, set)):
105
+ token_ids.update(int(token_id) for token_id in value if token_id is not None)
106
+ return token_ids
107
+
108
+
109
+ def _add_repetition_history(
110
+ history: torch.BoolTensor,
111
+ token_ids: torch.LongTensor,
112
+ eligible: torch.BoolTensor,
113
+ excluded_token_ids: set[int],
114
+ ) -> None:
115
+ """Add eligible row-local token IDs to a compact [batch, vocab] history."""
116
+
117
+ if token_ids.shape != eligible.shape or token_ids.shape[0] != history.shape[0]:
118
+ raise ValueError("Repetition history token and eligibility shapes must match.")
119
+ eligible = eligible.clone()
120
+ for token_id in excluded_token_ids:
121
+ eligible &= token_ids.ne(token_id)
122
+ if not bool(eligible.any()):
123
+ return
124
+ rows = torch.arange(history.shape[0], device=history.device)[:, None]
125
+ rows = rows.expand_as(token_ids)
126
+ history[rows[eligible], token_ids[eligible]] = True
127
+
128
+
129
+ class ModilifyMk2GenerationConfig(DiffusionGemmaGenerationConfig):
130
+ def __init__(self, **kwargs):
131
+ _reject_legacy_commit_fields({
132
+ name: kwargs.pop(name)
133
+ for name in _FORBIDDEN_COMMIT_FIELDS
134
+ if name in kwargs
135
+ })
136
+ self.turn_end_token_id: int | None = kwargs.pop("turn_end_token_id", None)
137
+ self.denoise_temperature: float = float(
138
+ kwargs.pop("denoise_temperature", DENOISE_TEMPERATURE)
139
+ )
140
+ self.commit_failure_budget: float = float(
141
+ kwargs.pop("commit_failure_budget", COMMIT_FAILURE_BUDGET)
142
+ )
143
+ self.jump_failure_budget: float = float(
144
+ kwargs.pop("jump_failure_budget", JUMP_FAILURE_BUDGET)
145
+ )
146
+ self.max_ponder_steps: int = kwargs.pop("max_ponder_steps", 64)
147
+ self.jump_on_no_progress_after: int = kwargs.pop("jump_on_no_progress_after", 12)
148
+ self.min_trajectory_progress: float = float(kwargs.pop("min_trajectory_progress", 0.005))
149
+ self.repetition_penalty: float = float(kwargs.pop("repetition_penalty", 1.0))
150
+ excluded_token_ids = kwargs.pop("repetition_penalty_exclude_token_ids", ())
151
+ self.repetition_penalty_exclude_token_ids: list[int] = list(
152
+ dict.fromkeys(int(token_id) for token_id in excluded_token_ids or ())
153
+ )
154
+ for name in _IGNORED_GENERATION_KEYS:
155
+ kwargs.pop(name, None)
156
+ kwargs.pop("t_min", None)
157
+ kwargs.pop("t_max", None)
158
+ parent_kwargs = {
159
+ name: kwargs.pop(name)
160
+ for name in tuple(kwargs)
161
+ if name in _PARENT_GENERATION_KEYS
162
+ }
163
+ super().__init__(
164
+ sampler_config=None,
165
+ stability_threshold=None,
166
+ confidence_threshold=None,
167
+ **parent_kwargs,
168
+ )
169
+ self.t_min = self.denoise_temperature
170
+ self.t_max = self.denoise_temperature
171
+ self.validate()
172
+
173
+ def update(self, defaults_only=False, allow_custom_entries=False, **kwargs):
174
+ """Apply supported overrides, including temperature and failure budget."""
175
+
176
+ _reject_legacy_commit_fields({
177
+ name: kwargs.pop(name)
178
+ for name in _FORBIDDEN_COMMIT_FIELDS
179
+ if name in kwargs
180
+ })
181
+ if "turn_end_token_id" in kwargs:
182
+ self.turn_end_token_id = kwargs.pop("turn_end_token_id")
183
+ if "denoise_temperature" in kwargs:
184
+ self.denoise_temperature = float(kwargs.pop("denoise_temperature"))
185
+ if "commit_failure_budget" in kwargs:
186
+ self.commit_failure_budget = float(kwargs.pop("commit_failure_budget"))
187
+ if "jump_failure_budget" in kwargs:
188
+ self.jump_failure_budget = float(kwargs.pop("jump_failure_budget"))
189
+ if "max_ponder_steps" in kwargs:
190
+ self.max_ponder_steps = kwargs.pop("max_ponder_steps")
191
+ if "jump_on_no_progress_after" in kwargs:
192
+ self.jump_on_no_progress_after = kwargs.pop("jump_on_no_progress_after")
193
+ if "min_trajectory_progress" in kwargs:
194
+ self.min_trajectory_progress = float(kwargs.pop("min_trajectory_progress"))
195
+ if "repetition_penalty" in kwargs:
196
+ self.repetition_penalty = float(kwargs.pop("repetition_penalty"))
197
+ if "repetition_penalty_exclude_token_ids" in kwargs:
198
+ excluded_token_ids = kwargs.pop("repetition_penalty_exclude_token_ids")
199
+ self.repetition_penalty_exclude_token_ids = list(
200
+ dict.fromkeys(int(token_id) for token_id in excluded_token_ids or ())
201
+ )
202
+ for name in _IGNORED_GENERATION_KEYS:
203
+ kwargs.pop(name, None)
204
+ kwargs.pop("t_min", None)
205
+ kwargs.pop("t_max", None)
206
+ unused = super().update(
207
+ defaults_only=defaults_only,
208
+ allow_custom_entries=allow_custom_entries,
209
+ **kwargs,
210
+ )
211
+ self.sampler_config = None
212
+ self.stability_threshold = None
213
+ self.confidence_threshold = None
214
+ self.t_min = self.denoise_temperature
215
+ self.t_max = self.denoise_temperature
216
+ return unused
217
+
218
+ def validate(self, **kwargs):
219
+ if self.max_new_tokens is not None and (
220
+ not isinstance(self.max_new_tokens, int) or self.max_new_tokens <= 0
221
+ ):
222
+ raise ValueError(f"`max_new_tokens` must be a positive integer, but got {self.max_new_tokens}")
223
+ if self.max_length is not None and (
224
+ not isinstance(self.max_length, int) or self.max_length <= 0
225
+ ):
226
+ raise ValueError(f"`max_length` must be a positive integer, but got {self.max_length}")
227
+ if self.turn_end_token_id is not None and (
228
+ not isinstance(self.turn_end_token_id, int) or self.turn_end_token_id < 0
229
+ ):
230
+ raise ValueError("`turn_end_token_id` must be a non-negative integer.")
231
+ if not isinstance(self.max_ponder_steps, int) or self.max_ponder_steps <= 0:
232
+ raise ValueError("`max_ponder_steps` must be a positive integer.")
233
+ if not isinstance(self.jump_on_no_progress_after, int) or self.jump_on_no_progress_after <= 0:
234
+ raise ValueError("`jump_on_no_progress_after` must be a positive integer.")
235
+ if not isinstance(self.min_trajectory_progress, (int, float)) or self.min_trajectory_progress < 0:
236
+ raise ValueError("`min_trajectory_progress` must be a non-negative number.")
237
+ if not math.isfinite(self.denoise_temperature) or self.denoise_temperature <= 0:
238
+ raise ValueError("`denoise_temperature` must be positive.")
239
+ if not math.isfinite(self.commit_failure_budget) or self.commit_failure_budget <= 0:
240
+ raise ValueError("`commit_failure_budget` must be positive.")
241
+ if not math.isfinite(self.jump_failure_budget) or self.jump_failure_budget <= 0:
242
+ raise ValueError("`jump_failure_budget` must be positive.")
243
+ if not math.isfinite(self.repetition_penalty) or self.repetition_penalty <= 0:
244
+ raise ValueError("`repetition_penalty` must be a finite positive number.")
245
+ if any(
246
+ not isinstance(token_id, int) or isinstance(token_id, bool) or token_id < 0
247
+ for token_id in self.repetition_penalty_exclude_token_ids
248
+ ):
249
+ raise ValueError(
250
+ "`repetition_penalty_exclude_token_ids` must contain non-negative integers."
251
+ )
252
+
253
+ @classmethod
254
+ def from_model_config(cls, model_config):
255
+ """Build generation defaults from a model configuration."""
256
+
257
+ return cls(
258
+ turn_end_token_id=model_config.turn_end_token_id,
259
+ denoise_temperature=getattr(
260
+ model_config, "denoise_temperature", DENOISE_TEMPERATURE
261
+ ),
262
+ commit_failure_budget=getattr(
263
+ model_config, "commit_failure_budget", COMMIT_FAILURE_BUDGET
264
+ ),
265
+ jump_failure_budget=getattr(
266
+ model_config, "jump_failure_budget", JUMP_FAILURE_BUDGET
267
+ ),
268
+ max_ponder_steps=getattr(model_config, "max_ponder_steps", 64),
269
+ jump_on_no_progress_after=getattr(
270
+ model_config, "jump_on_no_progress_after", 12
271
+ ),
272
+ min_trajectory_progress=getattr(
273
+ model_config, "min_trajectory_progress", 0.005
274
+ ),
275
+ repetition_penalty=getattr(model_config, "repetition_penalty", 1.0),
276
+ eos_token_id=getattr(
277
+ model_config,
278
+ "eos_token_id",
279
+ model_config.text_config.eos_token_id,
280
+ ),
281
+ )
282
+
283
+ @staticmethod
284
+ def _get_default_generation_params() -> dict[str, object]:
285
+ """Return defaults with no inherited entropy/readiness commit controls."""
286
+
287
+ return {
288
+ "max_new_tokens": 256,
289
+ "max_denoising_steps": 48,
290
+ "t_min": DENOISE_TEMPERATURE,
291
+ "t_max": DENOISE_TEMPERATURE,
292
+ }
293
+
294
+
295
+ def deterministic_episode_iteration_bound(
296
+ response_lengths: torch.LongTensor,
297
+ *,
298
+ max_ponder_steps: int,
299
+ ) -> int:
300
+ """Return a safe watchdog bound without assuming canvas-sized jumps."""
301
+
302
+ if response_lengths.numel() == 0:
303
+ raise ValueError("`response_lengths` must be non-empty.")
304
+ if max_ponder_steps <= 0:
305
+ raise ValueError("`max_ponder_steps` must be positive.")
306
+ return max(1, int(response_lengths.max()) * max_ponder_steps)
307
+
308
+
309
+ @dataclass
310
+ class ModilifyMk2GenerationOutput(ModelOutput):
311
+ sequences: torch.LongTensor
312
+ generated_lengths: torch.LongTensor | None = None
313
+ tokens_per_forward: torch.FloatTensor | None = None
314
+ past_key_values: Cache | None = None
315
+ stop_reason: str | tuple[str, ...] | None = None
316
+ committed_tokens: int | torch.LongTensor | None = None
317
+ denoise_steps: int | torch.LongTensor | None = None
318
+ no_progress_steps: int | torch.LongTensor | None = None
319
+ jump_count: int | torch.LongTensor | None = None
320
+ forced_jump_bad_count: int | torch.LongTensor | None = None
321
+ heavy_forward_count: int | torch.LongTensor | None = None
322
+ latent_context_update_count: int | torch.LongTensor | None = None
323
+ average_commit_len: float | torch.FloatTensor | None = None
324
+ state_shift_count: int | torch.LongTensor | None = None
325
+ latent_memory_norm: float | torch.FloatTensor | None = None
326
+ state_retention_score: float | torch.FloatTensor | None = None
327
+ logits: None = None
328
+ scores: None = None
329
+ hidden_states: None = None
330
+
331
+
332
+ @dataclass
333
+ class ModilifyMk2RollingState:
334
+ """All real iterative state; no vocabulary-sized tensor is retained."""
335
+
336
+ canvas: torch.LongTensor
337
+ confidence: torch.FloatTensor
338
+ entropy: torch.FloatTensor
339
+ age: torch.IntTensor
340
+ latent_state: LatentDeliberationState
341
+ history: TrajectoryHistory
342
+ tape: TrajectoryTape
343
+
344
+
345
+ def qualified_single_turn_end_lengths(
346
+ token_ids: torch.LongTensor,
347
+ qualified_mask: torch.BoolTensor,
348
+ turn_end_token_id: int,
349
+ ) -> torch.LongTensor:
350
+ if token_ids.ndim != 2 or token_ids.shape != qualified_mask.shape:
351
+ raise ValueError("Token IDs and qualification mask must share shape [batch, canvas].")
352
+ qualified_prefix = qualified_mask.long().cumprod(dim=1).sum(dim=-1)
353
+ clipped = first_committed_token_lengths(
354
+ token_ids,
355
+ qualified_prefix,
356
+ turn_end_token_id,
357
+ )
358
+ found = clipped.lt(qualified_prefix)
359
+ boundary_is_turn = token_ids.gather(
360
+ 1,
361
+ clipped.clamp_min(1).sub(1)[:, None],
362
+ ).squeeze(1).eq(turn_end_token_id)
363
+ return torch.where(
364
+ found | boundary_is_turn,
365
+ clipped,
366
+ torch.zeros_like(clipped),
367
+ )
368
+
369
+
370
+ def _prefix_length(prefix_mask: torch.BoolTensor) -> int:
371
+ common = prefix_mask.all(dim=0)
372
+ rejected = (~common).nonzero(as_tuple=False)
373
+ return common.shape[0] if rejected.numel() == 0 else int(rejected[0, 0])
374
+
375
+
376
+ def _tensor_stats(value: torch.Tensor) -> dict[str, float]:
377
+ data = value.detach().float()
378
+ return {
379
+ "min": float(data.min()), "max": float(data.max()),
380
+ "mean": float(data.mean()), "norm": float(data.norm()),
381
+ }
382
+
383
+
384
+ def _memory_diversity_stats(memory_slots: torch.Tensor) -> dict[str, float]:
385
+ """Trace-only collapse diagnostics; the pairwise term is tiny (32 slots)."""
386
+
387
+ memory = memory_slots.detach().float()
388
+ centered = memory - memory.mean(dim=1, keepdim=True)
389
+ maximum_difference = (
390
+ (memory[:, 1:] - memory[:, :-1]).abs().amax()
391
+ if memory.shape[1] > 1 else memory.new_zeros(())
392
+ )
393
+ normalized = torch.nn.functional.normalize(memory, dim=-1, eps=1.0e-6)
394
+ pairwise = normalized @ normalized.transpose(-1, -2)
395
+ off_diagonal = ~torch.eye(memory.shape[1], device=memory.device, dtype=torch.bool)
396
+ return {
397
+ "memory_slot_std": float(centered.square().mean().sqrt()),
398
+ "memory_slot_max_difference": float(maximum_difference),
399
+ "memory_slot_pairwise_cosine_mean": float(pairwise[:, off_diagonal].mean()),
400
+ }
401
+
402
+
403
+ def build_denoise_trace_event(
404
+ *,
405
+ denoise_step: int,
406
+ prefix_length: int,
407
+ committed_before: int,
408
+ committed_after: int,
409
+ no_progress_steps: int,
410
+ policy_prefix_mask: torch.BoolTensor,
411
+ commit_length: int,
412
+ ponder_fallback: bool,
413
+ state: ModilifyMk2RollingState,
414
+ proposal: torch.LongTensor,
415
+ committed_token_ids: torch.LongTensor,
416
+ step_elapsed_seconds: float,
417
+ latent_residual_diagnostics: dict[str, torch.Tensor] | None = None,
418
+ ) -> dict[str, object]:
419
+ return {
420
+ "event": "denoise_step",
421
+ "denoise_step": denoise_step,
422
+ "prefix_length": prefix_length,
423
+ "committed_before": committed_before,
424
+ "committed_after": committed_after,
425
+ "no_progress_steps": no_progress_steps,
426
+ "policy_prefix_count": int(policy_prefix_mask.sum()),
427
+ "policy_prefix_length": _prefix_length(policy_prefix_mask),
428
+ "commit_length": commit_length,
429
+ "ponder_fallback": bool(ponder_fallback),
430
+ "confidence": _tensor_stats(state.confidence),
431
+ "entropy": _tensor_stats(state.entropy),
432
+ "age": _tensor_stats(state.age),
433
+ "token_changed": _tensor_stats(state.latent_state.token_changed),
434
+ "confidence_delta": _tensor_stats(state.latent_state.confidence_delta),
435
+ "entropy_delta": _tensor_stats(state.latent_state.entropy_delta),
436
+ "ponder_steps": int(state.latent_state.ponder_steps[0]),
437
+ "stagnation_steps": int(state.latent_state.stagnation_steps[0]),
438
+ "history_fill": float(state.history.valid[0].float().mean()),
439
+ "memory_slots": _tensor_stats(state.latent_state.memory_slots),
440
+ "memory_diversity": _memory_diversity_stats(state.latent_state.memory_slots),
441
+ "latent_residual": {
442
+ name: float(value.detach().float())
443
+ for name, value in (latent_residual_diagnostics or {}).items()
444
+ },
445
+ "proposal_token_ids": proposal[0].detach().cpu().tolist(),
446
+ "draft_token_ids": state.canvas[0].detach().cpu().tolist(),
447
+ "committed_token_ids": committed_token_ids[0].detach().cpu().tolist(),
448
+ "step_elapsed_seconds": step_elapsed_seconds,
449
+ }
450
+
451
+
452
+ class NoiseCanvasSampler:
453
+ """Uniform diffusion noise source with no commit-policy responsibilities."""
454
+
455
+ def __init__(self, *, canvas_length: int, vocab_size: int) -> None:
456
+ self.canvas_length = int(canvas_length)
457
+ self.vocab_size = int(vocab_size)
458
+ self.initial_entropy = math.log(self.vocab_size)
459
+
460
+ def initialize_canvas(
461
+ self,
462
+ batch_size: int,
463
+ device: torch.device,
464
+ generators: Sequence[torch.Generator] | None = None,
465
+ ) -> torch.LongTensor:
466
+ if generators is not None:
467
+ if len(generators) != batch_size:
468
+ raise ValueError("Canvas sampling requires one generator per batch row.")
469
+ return torch.cat(
470
+ [
471
+ torch.randint(
472
+ self.vocab_size,
473
+ (1, self.canvas_length),
474
+ device=device,
475
+ generator=generator,
476
+ )
477
+ for generator in generators
478
+ ],
479
+ dim=0,
480
+ )
481
+ return torch.randint(
482
+ self.vocab_size,
483
+ (batch_size, self.canvas_length),
484
+ device=device,
485
+ )
486
+
487
+
488
+ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
489
+ def init_continuous_batching(
490
+ self,
491
+ generation_config=None,
492
+ continuous_batching_config=None,
493
+ workload_hints=None,
494
+ ):
495
+ """Create the ModilifyMk2 continuous manager behind the official API surface."""
496
+
497
+ from .continuous_batching import (
498
+ ModilifyMk2ContinuousBatchingManager,
499
+ continuous_config_fingerprint,
500
+ )
501
+
502
+ cached = getattr(self, "_cached_continuous_batching_manager", None)
503
+ if isinstance(cached, ModilifyMk2ContinuousBatchingManager) and not cached.destroyed:
504
+ requested_generation = generation_config or getattr(
505
+ self, "generation_config", None
506
+ )
507
+ requested_fingerprint = continuous_config_fingerprint(
508
+ requested_generation,
509
+ continuous_batching_config,
510
+ )
511
+ if cached.config_fingerprint == requested_fingerprint:
512
+ cached._prepare_for_next_session()
513
+ return cached
514
+ cached.destroy()
515
+ delattr(self, "_cached_continuous_batching_manager")
516
+ return ModilifyMk2ContinuousBatchingManager(
517
+ model=self,
518
+ generation_config=generation_config or getattr(self, "generation_config", None),
519
+ continuous_batching_config=continuous_batching_config,
520
+ workload_hints=workload_hints,
521
+ )
522
+
523
+ def destroy_cached_continuous_batching_manager(self) -> None:
524
+ manager = getattr(self, "_cached_continuous_batching_manager", None)
525
+ if manager is not None:
526
+ manager.destroy()
527
+ delattr(self, "_cached_continuous_batching_manager")
528
+
529
+ @contextmanager
530
+ @torch.no_grad()
531
+ def continuous_batching_context_manager(
532
+ self,
533
+ generation_config=None,
534
+ block: bool = True,
535
+ timeout: float | None = None,
536
+ continuous_batching_config=None,
537
+ persistent_manager: bool = False,
538
+ warmup: bool = True,
539
+ workload_hints=None,
540
+ ):
541
+ manager = self.init_continuous_batching(
542
+ generation_config=generation_config,
543
+ continuous_batching_config=continuous_batching_config,
544
+ workload_hints=workload_hints,
545
+ )
546
+ if persistent_manager:
547
+ self._cached_continuous_batching_manager = manager
548
+ if warmup and not manager.warmed_up:
549
+ manager.warmup()
550
+ manager.start()
551
+ try:
552
+ yield manager
553
+ finally:
554
+ manager.stop(
555
+ block=block,
556
+ timeout=timeout,
557
+ keep_for_next_session=persistent_manager,
558
+ )
559
+ if not persistent_manager:
560
+ manager.destroy()
561
+
562
+ @torch.no_grad()
563
+ def generate_batch(
564
+ self,
565
+ inputs: list[list[int]],
566
+ generation_config=None,
567
+ continuous_batching_config=None,
568
+ record_timestamps: bool = False,
569
+ progress_bar: bool = True,
570
+ persistent_manager: bool = False,
571
+ warmup: bool = True,
572
+ **kwargs: Any,
573
+ ) -> dict[str, object]:
574
+ """Official-compatible convenience wrapper over the continuous manager."""
575
+
576
+ del progress_bar
577
+ if any(
578
+ not isinstance(input_ids, list)
579
+ or not input_ids
580
+ or any(
581
+ not isinstance(token_id, int) or isinstance(token_id, bool)
582
+ for token_id in input_ids
583
+ )
584
+ for input_ids in inputs
585
+ ):
586
+ raise ValueError("Every `inputs` row must be a non-empty list of integer token IDs.")
587
+ seeds = kwargs.pop("seeds", None)
588
+ if seeds is not None and len(seeds) != len(inputs):
589
+ raise ValueError("`seeds` must contain one seed per request.")
590
+ if seeds is not None:
591
+ seeds = [int(seed) for seed in seeds]
592
+ if not inputs:
593
+ return {}
594
+ manager = self.init_continuous_batching(
595
+ generation_config=generation_config,
596
+ continuous_batching_config=continuous_batching_config,
597
+ )
598
+ if persistent_manager:
599
+ self._cached_continuous_batching_manager = manager
600
+ if warmup and not manager.warmed_up:
601
+ manager.warmup()
602
+ completed = False
603
+ try:
604
+ request_ids = []
605
+ queue_limit = int(
606
+ manager.continuous_batching_config.max_queue_size or 0
607
+ )
608
+ preload_count = (
609
+ len(inputs) if queue_limit == 0 else min(len(inputs), queue_limit)
610
+ )
611
+ for index, input_ids in enumerate(inputs):
612
+ if index == preload_count and not manager.is_running():
613
+ manager.start()
614
+ request_kwargs = dict(kwargs)
615
+ if seeds is not None:
616
+ request_kwargs["seed"] = int(seeds[index])
617
+ request_ids.append(
618
+ manager.add_request(
619
+ input_ids=input_ids,
620
+ record_timestamps=record_timestamps,
621
+ **request_kwargs,
622
+ )
623
+ )
624
+ manager.close_input()
625
+ # Unlike the open-ended manager API, generate_batch has its whole
626
+ # initial workload. Queue it before starting so the first cohort is
627
+ # filled deterministically up to scheduler/queue capacity.
628
+ if not manager.is_running():
629
+ manager.start()
630
+ final_outputs = {}
631
+ for output in manager:
632
+ if output.status in {
633
+ getattr(output.status.__class__, "FINISHED", output.status),
634
+ getattr(output.status.__class__, "FAILED", output.status),
635
+ }:
636
+ final_outputs[output.request_id] = output
637
+ result = {
638
+ request_id: final_outputs[request_id]
639
+ for request_id in request_ids
640
+ if request_id is not None
641
+ }
642
+ completed = True
643
+ finally:
644
+ manager.stop(
645
+ block=True,
646
+ keep_for_next_session=persistent_manager,
647
+ hard_stop=not completed,
648
+ )
649
+ if not persistent_manager:
650
+ manager.destroy()
651
+ return result
652
+
653
+ def _prepare_sampler(
654
+ self, generation_config: ModilifyMk2GenerationConfig, canvas_length: int | None = None
655
+ ) -> NoiseCanvasSampler:
656
+ del generation_config
657
+ return NoiseCanvasSampler(
658
+ canvas_length=canvas_length or self.config.canvas_length,
659
+ vocab_size=self.config.text_config.vocab_size,
660
+ )
661
+
662
+ @staticmethod
663
+ def _shift_state(
664
+ state: ModilifyMk2RollingState,
665
+ commit_length: int,
666
+ sampler: NoiseCanvasSampler,
667
+ **_: object,
668
+ ) -> ModilifyMk2RollingState:
669
+ if commit_length == 0:
670
+ return state
671
+ canvas_length = state.canvas.shape[1]
672
+ if not 0 < commit_length <= canvas_length:
673
+ raise ValueError(f"`commit_length` must be in [1, {canvas_length}].")
674
+ tail = sampler.initialize_canvas(state.canvas.shape[0], state.canvas.device)[:, :commit_length]
675
+ canvas = torch.cat((state.canvas[:, commit_length:], tail), dim=1)
676
+
677
+ def shift(value: torch.Tensor | None, fill_value: float | int = 0) -> torch.Tensor | None:
678
+ if value is None:
679
+ return None
680
+ tail_state = torch.full(
681
+ (value.shape[0], commit_length, *value.shape[2:]),
682
+ fill_value, device=value.device, dtype=value.dtype,
683
+ )
684
+ return torch.cat((value[:, commit_length:], tail_state), dim=1)
685
+
686
+ if hasattr(sampler, "initial_entropy"):
687
+ unknown_entropy = float(sampler.initial_entropy)
688
+ elif hasattr(sampler, "vocab_size"):
689
+ unknown_entropy = math.log(sampler.vocab_size)
690
+ else:
691
+ unknown_entropy = float(state.entropy.max())
692
+
693
+ return ModilifyMk2RollingState(
694
+ canvas=canvas,
695
+ confidence=shift(state.confidence),
696
+ entropy=shift(state.entropy, unknown_entropy),
697
+ age=shift(state.age),
698
+ latent_state=state.latent_state.shift(
699
+ commit_length, entropy_fill_value=unknown_entropy
700
+ ),
701
+ history=state.history.shift(commit_length, entropy_fill_value=unknown_entropy),
702
+ tape=state.tape,
703
+ )
704
+
705
+ @staticmethod
706
+ def _merge_state_rows(
707
+ previous: ModilifyMk2RollingState,
708
+ updated: ModilifyMk2RollingState,
709
+ update_mask: torch.BoolTensor,
710
+ ) -> ModilifyMk2RollingState:
711
+ """Keep inactive batch rows bit-identical while active rows advance."""
712
+
713
+ def choose(old: torch.Tensor, new: torch.Tensor) -> torch.Tensor:
714
+ mask = update_mask.view(update_mask.shape[0], *([1] * (old.ndim - 1)))
715
+ return torch.where(mask, new, old)
716
+
717
+ old_latent = previous.latent_state
718
+ new_latent = updated.latent_state
719
+ latent = LatentDeliberationState(
720
+ memory_slots=choose(old_latent.memory_slots, new_latent.memory_slots),
721
+ confidence=choose(old_latent.confidence, new_latent.confidence),
722
+ entropy=choose(old_latent.entropy, new_latent.entropy),
723
+ age=choose(old_latent.age, new_latent.age),
724
+ token_changed=choose(old_latent.token_changed, new_latent.token_changed),
725
+ confidence_delta=choose(old_latent.confidence_delta, new_latent.confidence_delta),
726
+ entropy_delta=choose(old_latent.entropy_delta, new_latent.entropy_delta),
727
+ ponder_steps=choose(old_latent.ponder_steps, new_latent.ponder_steps),
728
+ stagnation_steps=choose(old_latent.stagnation_steps, new_latent.stagnation_steps),
729
+ )
730
+ return ModilifyMk2RollingState(
731
+ canvas=choose(previous.canvas, updated.canvas),
732
+ confidence=choose(previous.confidence, updated.confidence),
733
+ entropy=choose(previous.entropy, updated.entropy),
734
+ age=choose(previous.age, updated.age),
735
+ latent_state=latent,
736
+ history=choose_trajectory_history(
737
+ previous.history, updated.history, update_mask
738
+ ),
739
+ tape=choose_trajectory_tape(previous.tape, updated.tape, update_mask),
740
+ )
741
+
742
+ def _write_committed_memory(
743
+ self,
744
+ *,
745
+ previous_history: TrajectoryHistory,
746
+ next_state: ModilifyMk2RollingState,
747
+ working_state: torch.Tensor,
748
+ history_projected: torch.Tensor,
749
+ heavy_hidden: torch.Tensor,
750
+ commit_lengths: torch.Tensor,
751
+ prefix_lengths: torch.Tensor,
752
+ commit_reason: torch.Tensor,
753
+ **kwargs: Any,
754
+ ) -> ModilifyMk2RollingState:
755
+ if not bool(commit_lengths.gt(0).any()):
756
+ return next_state
757
+ memory, _diagnostics = self.latent_deliberation.commit_write(
758
+ memory=next_state.latent_state.memory_slots,
759
+ working_state=working_state,
760
+ history=previous_history,
761
+ history_projected=history_projected,
762
+ heavy_hidden=heavy_hidden,
763
+ commit_lengths=commit_lengths,
764
+ prefix_lengths=prefix_lengths,
765
+ commit_reason=commit_reason,
766
+ )
767
+ return replace(
768
+ next_state,
769
+ latent_state=replace(next_state.latent_state, memory_slots=memory),
770
+ )
771
+
772
+ @staticmethod
773
+ def _shift_state_rows(
774
+ state: ModilifyMk2RollingState,
775
+ commit_lengths: torch.LongTensor,
776
+ sampler: NoiseCanvasSampler,
777
+ generators: Sequence[torch.Generator] | None = None,
778
+ ) -> ModilifyMk2RollingState:
779
+ """Shift every rolling row by its own committed prefix length."""
780
+
781
+ batch_size, canvas_length = state.canvas.shape
782
+ if commit_lengths.shape != (batch_size,):
783
+ raise ValueError("Commit lengths must have shape [batch].")
784
+ if not bool(commit_lengths.gt(0).any()):
785
+ if generators is not None:
786
+ if len(generators) != batch_size:
787
+ raise ValueError("State shifting requires one generator per batch row.")
788
+ # Seeded generation deliberately advances every active request
789
+ # once per denoise step, independent of the other active rows.
790
+ for generator in generators:
791
+ try:
792
+ sampler.initialize_canvas(
793
+ 1, state.canvas.device, generators=[generator]
794
+ )
795
+ except TypeError:
796
+ sampler.initialize_canvas(1, state.canvas.device)
797
+ return state
798
+ positions = torch.arange(canvas_length, device=state.canvas.device)[None, :]
799
+ source = positions + commit_lengths[:, None]
800
+ retained = source.lt(canvas_length)
801
+
802
+ def shift(value: torch.Tensor, fill_value: float | int = 0) -> torch.Tensor:
803
+ index = source.clamp_max(canvas_length - 1)
804
+ index = index.view(
805
+ batch_size, canvas_length, *([1] * (value.ndim - 2))
806
+ ).expand_as(value)
807
+ gathered = value.gather(1, index)
808
+ mask = retained.view(
809
+ batch_size, canvas_length, *([1] * (value.ndim - 2))
810
+ )
811
+ fill = torch.as_tensor(fill_value, device=value.device, dtype=value.dtype)
812
+ return torch.where(mask, gathered, fill)
813
+
814
+ if generators is None:
815
+ tail = sampler.initialize_canvas(batch_size, state.canvas.device)
816
+ else:
817
+ if len(generators) != batch_size:
818
+ raise ValueError("State shifting requires one generator per batch row.")
819
+ tail = torch.zeros_like(state.canvas)
820
+ for row, (commit_length, generator) in enumerate(
821
+ zip(commit_lengths.detach().cpu().tolist(), generators, strict=True)
822
+ ):
823
+ try:
824
+ sampled = sampler.initialize_canvas(
825
+ 1,
826
+ state.canvas.device,
827
+ generators=[generator],
828
+ )
829
+ except TypeError:
830
+ # Preserve compatibility with deterministic test/custom
831
+ # samplers written before per-request RNG was introduced.
832
+ sampled = sampler.initialize_canvas(1, state.canvas.device)
833
+ tail[row] = sampled[0]
834
+ canvas = torch.cat((state.canvas, tail), dim=1).gather(1, source)
835
+ unknown_entropy = float(sampler.initial_entropy)
836
+ latent = state.latent_state
837
+ committed = commit_lengths.gt(0)
838
+ shifted_latent = LatentDeliberationState(
839
+ memory_slots=latent.memory_slots.clone(),
840
+ confidence=shift(latent.confidence),
841
+ entropy=shift(latent.entropy, unknown_entropy),
842
+ age=shift(latent.age),
843
+ token_changed=shift(latent.token_changed),
844
+ confidence_delta=shift(latent.confidence_delta),
845
+ entropy_delta=shift(latent.entropy_delta),
846
+ ponder_steps=torch.where(
847
+ committed, torch.zeros_like(latent.ponder_steps), latent.ponder_steps
848
+ ),
849
+ stagnation_steps=torch.where(
850
+ committed, torch.zeros_like(latent.stagnation_steps), latent.stagnation_steps
851
+ ),
852
+ )
853
+ return ModilifyMk2RollingState(
854
+ canvas=canvas,
855
+ confidence=shift(state.confidence),
856
+ entropy=shift(state.entropy, unknown_entropy),
857
+ age=shift(state.age),
858
+ latent_state=shifted_latent,
859
+ history=state.history.shift(commit_lengths, entropy_fill_value=unknown_entropy),
860
+ tape=state.tape,
861
+ )
862
+
863
+ @torch.inference_mode()
864
+ def generate(
865
+ self,
866
+ input_ids: torch.LongTensor | None = None,
867
+ past_key_values: Cache | None = None,
868
+ streamer: BaseStreamer | None = None,
869
+ generation_config: ModilifyMk2GenerationConfig | None = None,
870
+ logits_processor: LogitsProcessorList | None = None,
871
+ denoise_trace_callback: Callable[[dict[str, object]], None] | None = None,
872
+ **kwargs,
873
+ ) -> ModilifyMk2GenerationOutput:
874
+ request_seeds = kwargs.pop("seeds", None)
875
+ scalar_seed = kwargs.pop("seed", None)
876
+ if request_seeds is not None and scalar_seed is not None:
877
+ raise ValueError("Pass either `seed` or `seeds`, not both.")
878
+ generation_config, model_kwargs = self._prepare_generation_config(generation_config, **kwargs)
879
+ if input_ids is None or input_ids.ndim != 2 or input_ids.shape[0] < 1:
880
+ raise ValueError("ModilifyMk2 generation requires `input_ids` with shape [batch, sequence].")
881
+ if logits_processor:
882
+ raise ValueError(
883
+ "ModilifyMk2 uses its built-in exact sampler and does not accept "
884
+ "custom logits processors."
885
+ )
886
+ batch_size, input_width = input_ids.shape
887
+ if scalar_seed is not None:
888
+ request_seeds = [int(scalar_seed) + row for row in range(batch_size)]
889
+ elif batch_size > 1 and request_seeds is None:
890
+ # Static B>1 still uses independent row generators, but its default
891
+ # path must advance the caller's global device RNG just like normal
892
+ # generation instead of repeating torch.initial_seed() forever.
893
+ seed_parts = torch.randint(
894
+ 0,
895
+ (1 << 31) - 1,
896
+ (batch_size, 2),
897
+ device=input_ids.device,
898
+ dtype=torch.int64,
899
+ ).detach().cpu().tolist()
900
+ request_seeds = [
901
+ (int(high) << 31) | int(low) for high, low in seed_parts
902
+ ]
903
+ if request_seeds is not None:
904
+ if len(request_seeds) != batch_size:
905
+ raise ValueError("`seeds` must contain one seed per batch row.")
906
+ sampling_generators = []
907
+ for seed_value in request_seeds:
908
+ generator = torch.Generator(device=input_ids.device)
909
+ generator.manual_seed(int(seed_value) & ((1 << 63) - 1))
910
+ sampling_generators.append(generator)
911
+ else:
912
+ sampling_generators = None
913
+ if batch_size > 1 and streamer is not None:
914
+ raise ValueError("ModilifyMk2 streamers currently support batch size 1 only.")
915
+ if batch_size > 1 and denoise_trace_callback is not None:
916
+ raise ValueError("ModilifyMk2 denoise tracing currently supports batch size 1 only.")
917
+ if batch_size > 1 and past_key_values is not None:
918
+ raise ValueError("Batched ModilifyMk2 generation requires a fresh KV cache.")
919
+ if batch_size > 1:
920
+ from .continuous_batching import generate_static_batch_with_logical_cache
921
+
922
+ attention_mask = model_kwargs.pop("attention_mask", None)
923
+ canonical_mask = (
924
+ torch.ones_like(input_ids, dtype=torch.bool)
925
+ if attention_mask is None
926
+ else attention_mask.to(device=input_ids.device, dtype=torch.bool)
927
+ )
928
+ if canonical_mask.shape != input_ids.shape:
929
+ raise ValueError(
930
+ "`attention_mask` must have the same shape as `input_ids`."
931
+ )
932
+ row_limits = []
933
+ for prompt_length in canonical_mask.long().sum(dim=-1).tolist():
934
+ _, resolved_max_new_tokens = self._prepare_generated_length(
935
+ generation_config, int(prompt_length)
936
+ )
937
+ if resolved_max_new_tokens <= 0:
938
+ raise ValueError(
939
+ "The requested maximum length leaves no room to generate "
940
+ "for every batch row."
941
+ )
942
+ row_limits.append(int(resolved_max_new_tokens))
943
+ provided_positions = model_kwargs.pop("position_ids", None)
944
+ if provided_positions is not None:
945
+ canonical_positions = (
946
+ canonical_mask.long().cumsum(dim=-1).sub(1).clamp_min(0)
947
+ ).to(provided_positions)
948
+ if not torch.equal(provided_positions, canonical_positions):
949
+ raise ValueError(
950
+ "Continuous static batches require canonical row-local `position_ids`."
951
+ )
952
+ if model_kwargs:
953
+ unsupported = ", ".join(sorted(model_kwargs))
954
+ raise ValueError(
955
+ f"Unsupported batched ModilifyMk2 generation arguments: {unsupported}"
956
+ )
957
+ return generate_static_batch_with_logical_cache(
958
+ self,
959
+ input_ids,
960
+ canonical_mask,
961
+ generation_config,
962
+ seeds=request_seeds,
963
+ max_new_tokens=row_limits,
964
+ )
965
+ device = input_ids.device
966
+ dtype = self.model.decoder.embed_tokens.weight.dtype
967
+ canvas_length = self.config.canvas_length
968
+ cached_length = past_key_values.get_seq_length() if past_key_values is not None else 0
969
+ repetition_penalty = float(generation_config.repetition_penalty)
970
+ repetition_enabled = repetition_penalty != 1.0
971
+ if repetition_enabled and cached_length:
972
+ raise ValueError(
973
+ "Repetition penalty requires a fresh KV cache so the complete prompt "
974
+ "token history is available."
975
+ )
976
+ _, max_new_tokens = self._prepare_generated_length(
977
+ generation_config, cached_length + input_width
978
+ )
979
+ max_iterations = deterministic_episode_iteration_bound(
980
+ torch.tensor([max_new_tokens]),
981
+ max_ponder_steps=generation_config.max_ponder_steps,
982
+ )
983
+ if past_key_values is None:
984
+ past_key_values = self._prepare_cache_for_generation(
985
+ generation_config,
986
+ batch_size=batch_size,
987
+ # Ragged batches append dense, masked cache blocks. In the
988
+ # worst case only one row advances in each block.
989
+ max_length=input_width + batch_size * max_new_tokens,
990
+ )
991
+ expected_mask_width = cached_length + input_width
992
+ cache_attention_mask = model_kwargs.pop(
993
+ "attention_mask",
994
+ torch.ones(
995
+ batch_size, expected_mask_width, dtype=torch.bool, device=device
996
+ ),
997
+ ).bool()
998
+ if cache_attention_mask.shape != (batch_size, expected_mask_width):
999
+ raise ValueError(
1000
+ "`attention_mask` must have shape [batch, cached_length + sequence]."
1001
+ )
1002
+ provided_position_ids = model_kwargs.pop("position_ids", None)
1003
+ if provided_position_ids is not None:
1004
+ if provided_position_ids.shape != input_ids.shape:
1005
+ raise ValueError("`position_ids` must have the same shape as `input_ids`.")
1006
+ prompt_positions = provided_position_ids.to(device=device, dtype=torch.int32)
1007
+ elif cached_length:
1008
+ prompt_positions = torch.arange(
1009
+ cached_length,
1010
+ cached_length + input_width,
1011
+ device=device,
1012
+ dtype=torch.int32,
1013
+ ).unsqueeze(0)
1014
+ else:
1015
+ input_mask = cache_attention_mask[:, -input_width:]
1016
+ prompt_positions = input_mask.long().cumsum(dim=-1).sub(1).clamp_min(0).to(torch.int32)
1017
+ logical_lengths = cache_attention_mask.long().sum(dim=-1)
1018
+ if input_width:
1019
+ encoder_keys = ("pixel_values", "mm_token_type_ids", "image_position_ids")
1020
+ encoder_kwargs = {
1021
+ key: model_kwargs.pop(key)
1022
+ for key in encoder_keys
1023
+ if key in model_kwargs
1024
+ }
1025
+ past_key_values = self.model.encoder(
1026
+ input_ids=input_ids,
1027
+ attention_mask=cache_attention_mask,
1028
+ past_key_values=past_key_values,
1029
+ position_ids=prompt_positions,
1030
+ **encoder_kwargs,
1031
+ ).past_key_values
1032
+
1033
+ sampler = self._prepare_sampler(generation_config, canvas_length)
1034
+ latent = LatentDeliberationState.empty(
1035
+ batch_size=batch_size, canvas_length=canvas_length,
1036
+ latent_dim=self.config.latent_dim, memory_slots=self.config.latent_memory_slots,
1037
+ device=device, dtype=dtype,
1038
+ )
1039
+ history = TrajectoryHistory.empty(
1040
+ batch_size=batch_size,
1041
+ canvas_length=canvas_length,
1042
+ hidden_size=self.config.text_config.hidden_size,
1043
+ history_length=self.config.latent_history_length,
1044
+ device=device,
1045
+ dtype=dtype,
1046
+ )
1047
+ try:
1048
+ initial_canvas = sampler.initialize_canvas(
1049
+ batch_size, device, generators=sampling_generators
1050
+ )
1051
+ except TypeError:
1052
+ initial_canvas = sampler.initialize_canvas(batch_size, device)
1053
+ state = ModilifyMk2RollingState(
1054
+ canvas=initial_canvas,
1055
+ confidence=torch.zeros(
1056
+ batch_size, canvas_length, device=device, dtype=torch.float32
1057
+ ),
1058
+ entropy=torch.full(
1059
+ (batch_size, canvas_length), math.log(self.config.text_config.vocab_size),
1060
+ device=device, dtype=torch.float32,
1061
+ ),
1062
+ age=torch.zeros(
1063
+ batch_size, canvas_length, device=device, dtype=torch.int32
1064
+ ),
1065
+ latent_state=latent,
1066
+ history=history,
1067
+ tape=empty_trajectory_tape(
1068
+ batch_size=batch_size,
1069
+ config=self.config,
1070
+ device=device,
1071
+ dtype=dtype,
1072
+ ),
1073
+ )
1074
+ turn_end = (
1075
+ self.config.turn_end_token_id
1076
+ if generation_config.turn_end_token_id is None
1077
+ else generation_config.turn_end_token_id
1078
+ )
1079
+ configured_eos = generation_config.eos_token_id
1080
+ if configured_eos is None:
1081
+ configured_eos = self.config.eos_token_id
1082
+ if isinstance(configured_eos, int):
1083
+ configured_eos = [configured_eos]
1084
+ stop_token_ids = tuple(
1085
+ dict.fromkeys((int(turn_end), *(int(value) for value in configured_eos or ())))
1086
+ )
1087
+ pad_token_id = generation_config.pad_token_id
1088
+ if pad_token_id is None:
1089
+ pad_token_id = getattr(self.config, "pad_token_id", None)
1090
+ if isinstance(pad_token_id, (list, tuple)):
1091
+ pad_token_id = pad_token_id[0]
1092
+ pad_token_id = int(0 if pad_token_id is None else pad_token_id)
1093
+ excluded_repetition_token_ids = _flatten_token_ids(
1094
+ generation_config.repetition_penalty_exclude_token_ids,
1095
+ generation_config.pad_token_id,
1096
+ generation_config.bos_token_id,
1097
+ generation_config.eos_token_id,
1098
+ generation_config.turn_end_token_id,
1099
+ getattr(self.config, "image_token_id", None),
1100
+ )
1101
+ repetition_history = None
1102
+ if repetition_enabled:
1103
+ repetition_history = torch.zeros(
1104
+ (batch_size, self.config.text_config.vocab_size),
1105
+ dtype=torch.bool,
1106
+ device=device,
1107
+ )
1108
+ _add_repetition_history(
1109
+ repetition_history,
1110
+ input_ids,
1111
+ cache_attention_mask[:, -input_width:],
1112
+ excluded_repetition_token_ids,
1113
+ )
1114
+ generated = torch.full(
1115
+ (batch_size, max_new_tokens),
1116
+ pad_token_id,
1117
+ dtype=input_ids.dtype,
1118
+ device=device,
1119
+ )
1120
+ committed = torch.zeros(batch_size, dtype=torch.long, device=device)
1121
+ denoise_steps = torch.zeros_like(committed)
1122
+ jumps = torch.zeros_like(committed)
1123
+ forced_jump_tokens = torch.zeros_like(committed)
1124
+ shifts = torch.zeros_like(committed)
1125
+ retention_scores = torch.zeros(batch_size, dtype=torch.float32, device=device)
1126
+ stop_codes = torch.zeros_like(committed)
1127
+ active_rows = torch.ones(batch_size, dtype=torch.bool, device=device)
1128
+ canvas_positions = torch.arange(canvas_length, device=device)[None, :]
1129
+ if streamer is not None:
1130
+ streamer.put(input_ids.cpu())
1131
+
1132
+ while bool(active_rows.any()):
1133
+ started = time.perf_counter()
1134
+ prefix_length = past_key_values.get_seq_length()
1135
+ decoder_positions = (
1136
+ logical_lengths[:, None]
1137
+ + torch.arange(canvas_length, device=device)[None, :]
1138
+ ).to(torch.int32)
1139
+ denoise_steps += active_rows.long()
1140
+ decoder_attention_mask = torch.cat(
1141
+ (
1142
+ cache_attention_mask,
1143
+ torch.ones(
1144
+ batch_size,
1145
+ canvas_length,
1146
+ dtype=torch.bool,
1147
+ device=device,
1148
+ ),
1149
+ ),
1150
+ dim=-1,
1151
+ )
1152
+ output = self(
1153
+ input_ids=None, past_key_values=past_key_values,
1154
+ decoder_input_ids=state.canvas,
1155
+ previous_confidence=state.confidence, previous_entropy=state.entropy,
1156
+ token_age=state.age, latent_state=state.latent_state,
1157
+ history=state.history,
1158
+ tape=state.tape,
1159
+ decoder_position_ids=decoder_positions, decoder_read_cache=True,
1160
+ decoder_attention_mask=decoder_attention_mask,
1161
+ compact_vocab=True,
1162
+ denoise_temperature=generation_config.denoise_temperature,
1163
+ repetition_token_mask=repetition_history,
1164
+ repetition_penalty=repetition_penalty,
1165
+ sampling_generators=sampling_generators,
1166
+ **model_kwargs,
1167
+ )
1168
+ if (
1169
+ output.proposal is None
1170
+ or output.proposal_confidence is None
1171
+ or output.token_entropy is None
1172
+ or output.greedy_proposal is None
1173
+ or output.greedy_confidence is None
1174
+ ):
1175
+ raise RuntimeError(
1176
+ "Compact vocabulary forward did not return proposal statistics."
1177
+ )
1178
+ proposal = output.proposal
1179
+ proposal_confidence = output.proposal_confidence
1180
+ token_entropy = output.token_entropy
1181
+ greedy_proposal = output.greedy_proposal
1182
+ greedy_confidence = output.greedy_confidence
1183
+ next_canvas = proposal.clone()
1184
+ next_confidence = proposal_confidence.float()
1185
+ next_latent = replace(
1186
+ output.next_latent_state,
1187
+ confidence=next_confidence.detach().float(),
1188
+ entropy=token_entropy.detach().float(),
1189
+ age=state.age + 1,
1190
+ token_changed=next_canvas.ne(state.canvas).detach().float(),
1191
+ confidence_delta=next_confidence.detach().float() - state.confidence,
1192
+ entropy_delta=token_entropy.detach().float() - state.entropy,
1193
+ )
1194
+ remaining = torch.tensor(
1195
+ max_new_tokens, device=device, dtype=torch.long
1196
+ ).sub(committed)
1197
+ remaining_canvas = remaining[:, None].gt(canvas_positions)
1198
+ tape_probes, tape_valid = self.latent_deliberation.encode_tape_frame(
1199
+ output.heavy_hidden_state, remaining_canvas
1200
+ )
1201
+ next_state = ModilifyMk2RollingState(
1202
+ canvas=next_canvas, confidence=next_confidence,
1203
+ entropy=token_entropy, age=state.age + 1,
1204
+ latent_state=next_latent,
1205
+ history=state.history.append(
1206
+ output.heavy_hidden_state,
1207
+ next_confidence,
1208
+ token_entropy,
1209
+ next_canvas.ne(state.canvas).detach().float(),
1210
+ live_mask=remaining_canvas,
1211
+ ),
1212
+ tape=state.tape.append(tape_probes, tape_valid),
1213
+ )
1214
+ next_state = self._merge_state_rows(state, next_state, active_rows)
1215
+ normal_failure_rate = fused_commit_failure_rate(
1216
+ proposal_confidence, token_entropy,
1217
+ vocab_size=self.config.text_config.vocab_size,
1218
+ )
1219
+ jump_failure_rate = fused_commit_failure_rate(
1220
+ greedy_confidence, token_entropy,
1221
+ vocab_size=self.config.text_config.vocab_size,
1222
+ )
1223
+ previous_failure_rate = fused_commit_failure_rate(
1224
+ state.confidence, state.entropy,
1225
+ vocab_size=self.config.text_config.vocab_size,
1226
+ )
1227
+ policy_decision = select_commit_lengths(
1228
+ sampled_token_ids=proposal,
1229
+ normal_failure_rate=normal_failure_rate,
1230
+ previous_failure_rate=previous_failure_rate,
1231
+ greedy_token_ids=greedy_proposal,
1232
+ jump_failure_rate=jump_failure_rate,
1233
+ ponder_steps=state.latent_state.ponder_steps,
1234
+ stagnation_steps=state.latent_state.stagnation_steps,
1235
+ active_rows=active_rows,
1236
+ remaining_lengths=remaining,
1237
+ failure_budget=generation_config.commit_failure_budget,
1238
+ jump_failure_budget=generation_config.jump_failure_budget,
1239
+ stop_token_id=stop_token_ids,
1240
+ max_ponder_steps=generation_config.max_ponder_steps,
1241
+ stagnation_threshold=generation_config.jump_on_no_progress_after,
1242
+ min_progress=generation_config.min_trajectory_progress,
1243
+ )
1244
+ normal_commit = policy_decision.normal_lengths
1245
+ policy_prefix_mask = canvas_positions.lt(normal_commit[:, None])
1246
+ next_ponder = policy_decision.ponder_steps
1247
+ next_stagnation = policy_decision.stagnation_steps
1248
+ commit_lengths = policy_decision.commit_lengths
1249
+ jump_rows = policy_decision.jump_rows
1250
+ jumps += jump_rows.long()
1251
+ forced_jump_tokens += torch.where(
1252
+ jump_rows, commit_lengths, torch.zeros_like(commit_lengths)
1253
+ )
1254
+ commit_positions = canvas_positions.lt(commit_lengths[:, None])
1255
+ if bool(jump_rows.any()):
1256
+ next_state = replace(
1257
+ next_state,
1258
+ canvas=torch.where(
1259
+ commit_positions & jump_rows[:, None],
1260
+ policy_decision.commit_token_ids,
1261
+ next_state.canvas,
1262
+ ),
1263
+ )
1264
+ next_state = replace(
1265
+ next_state,
1266
+ latent_state=replace(
1267
+ next_state.latent_state,
1268
+ ponder_steps=next_ponder,
1269
+ stagnation_steps=next_stagnation,
1270
+ ),
1271
+ )
1272
+ commit_token_ids = policy_decision.commit_token_ids
1273
+ before = committed.clone()
1274
+ write_rows = torch.arange(batch_size, device=device)[:, None].expand_as(
1275
+ commit_token_ids
1276
+ )
1277
+ write_positions = before[:, None] + canvas_positions
1278
+ generated[
1279
+ write_rows[commit_positions], write_positions[commit_positions]
1280
+ ] = commit_token_ids[commit_positions]
1281
+ if repetition_history is not None:
1282
+ _add_repetition_history(
1283
+ repetition_history,
1284
+ commit_token_ids,
1285
+ commit_positions,
1286
+ excluded_repetition_token_ids,
1287
+ )
1288
+
1289
+ commit_width = int(commit_lengths.max())
1290
+ if commit_width:
1291
+ block_mask = torch.arange(commit_width, device=device)[None, :].lt(
1292
+ commit_lengths[:, None]
1293
+ )
1294
+ committed_block = torch.where(
1295
+ block_mask,
1296
+ commit_token_ids[:, :commit_width],
1297
+ torch.full(
1298
+ (batch_size, commit_width),
1299
+ pad_token_id,
1300
+ device=device,
1301
+ dtype=input_ids.dtype,
1302
+ ),
1303
+ )
1304
+ block_positions = (
1305
+ logical_lengths[:, None]
1306
+ + torch.arange(commit_width, device=device)[None, :]
1307
+ ).to(torch.int32)
1308
+ block_positions = torch.where(
1309
+ block_mask, block_positions, torch.zeros_like(block_positions)
1310
+ )
1311
+ cache_attention_mask = torch.cat(
1312
+ (cache_attention_mask, block_mask), dim=-1
1313
+ )
1314
+ past_key_values = self.model.encoder(
1315
+ input_ids=committed_block,
1316
+ attention_mask=cache_attention_mask,
1317
+ past_key_values=past_key_values,
1318
+ position_ids=block_positions,
1319
+ ).past_key_values
1320
+ if streamer is not None:
1321
+ streamer.put(committed_block.cpu())
1322
+ committed += commit_lengths
1323
+ logical_lengths += commit_lengths
1324
+ committed_rows = commit_lengths.gt(0)
1325
+ shifts += committed_rows.long()
1326
+ if (
1327
+ output.history_projected is None
1328
+ or output.working_state is None
1329
+ ):
1330
+ raise RuntimeError("Forward did not return working trajectory features.")
1331
+ next_state = self._write_committed_memory(
1332
+ previous_history=state.history,
1333
+ next_state=next_state,
1334
+ working_state=output.working_state,
1335
+ history_projected=output.history_projected,
1336
+ heavy_hidden=output.heavy_hidden_state,
1337
+ commit_lengths=commit_lengths,
1338
+ prefix_lengths=logical_lengths,
1339
+ commit_reason=infer_commit_reason(
1340
+ commit_lengths,
1341
+ jump_rows=jump_rows,
1342
+ commit_token_ids=commit_token_ids,
1343
+ terminal_token_ids=stop_token_ids,
1344
+ ),
1345
+ )
1346
+ shifted = self._shift_state_rows(
1347
+ next_state,
1348
+ commit_lengths,
1349
+ sampler,
1350
+ generators=sampling_generators,
1351
+ )
1352
+ retention_scores += committed_rows.float()
1353
+ state = shifted
1354
+
1355
+ turn_hits = (
1356
+ commit_token_ids.eq(turn_end) & commit_positions
1357
+ ).any(dim=-1)
1358
+ eos_hits = torch.zeros_like(turn_hits)
1359
+ for token_id in stop_token_ids:
1360
+ if token_id != turn_end:
1361
+ eos_hits |= (
1362
+ commit_token_ids.eq(token_id) & commit_positions
1363
+ ).any(dim=-1)
1364
+ stop_codes = torch.where(
1365
+ stop_codes.eq(0) & turn_hits,
1366
+ torch.ones_like(stop_codes),
1367
+ stop_codes,
1368
+ )
1369
+ stop_codes = torch.where(
1370
+ stop_codes.eq(0) & eos_hits,
1371
+ torch.full_like(stop_codes, 2),
1372
+ stop_codes,
1373
+ )
1374
+ stop_codes = torch.where(
1375
+ stop_codes.eq(0) & committed.ge(max_new_tokens),
1376
+ torch.full_like(stop_codes, 3),
1377
+ stop_codes,
1378
+ )
1379
+ if generation_config.max_denoising_steps is not None:
1380
+ stop_codes = torch.where(
1381
+ stop_codes.eq(0)
1382
+ & denoise_steps.ge(generation_config.max_denoising_steps),
1383
+ torch.full_like(stop_codes, 4),
1384
+ stop_codes,
1385
+ )
1386
+ stop_codes = torch.where(
1387
+ stop_codes.eq(0) & denoise_steps.ge(max_iterations),
1388
+ torch.full_like(stop_codes, 5),
1389
+ stop_codes,
1390
+ )
1391
+ if denoise_trace_callback is not None:
1392
+ denoise_trace_callback(build_denoise_trace_event(
1393
+ denoise_step=int(denoise_steps[0]),
1394
+ prefix_length=prefix_length, committed_before=int(before[0]),
1395
+ committed_after=int(committed[0]),
1396
+ no_progress_steps=int(next_stagnation[0]),
1397
+ policy_prefix_mask=policy_prefix_mask,
1398
+ commit_length=int(commit_lengths[0]),
1399
+ ponder_fallback=bool(jump_rows[0]), state=next_state, proposal=proposal,
1400
+ committed_token_ids=commit_token_ids[:, :commit_width],
1401
+ step_elapsed_seconds=time.perf_counter() - started,
1402
+ latent_residual_diagnostics=output.latent_residual_diagnostics,
1403
+ ))
1404
+ active_rows = stop_codes.eq(0)
1405
+
1406
+ output_width = int(committed.max())
1407
+ sequences = torch.cat((input_ids, generated[:, :output_width]), dim=-1)
1408
+ if streamer is not None:
1409
+ streamer.end()
1410
+ reason_names = {
1411
+ 1: "turn_end",
1412
+ 2: "eos",
1413
+ 3: "max_new_tokens",
1414
+ 4: "max_denoising_steps",
1415
+ 5: "episode_watchdog",
1416
+ }
1417
+ stop_reasons = tuple(
1418
+ reason_names.get(code, "unknown") for code in stop_codes.detach().cpu().tolist()
1419
+ )
1420
+ tokens_per_forward = committed.float() / denoise_steps.clamp_min(1).float()
1421
+ average_commit_len = committed.float() / shifts.clamp_min(1).float()
1422
+ latent_memory_norm = state.latent_state.memory_slots.float().norm(dim=-1).mean(dim=-1)
1423
+ state_retention_score = retention_scores / shifts.clamp_min(1).float()
1424
+
1425
+ def scalar_or_tensor(value: torch.Tensor, *, floating: bool = False):
1426
+ if batch_size > 1:
1427
+ return value
1428
+ item = value[0].item()
1429
+ return float(item) if floating else int(item)
1430
+
1431
+ return ModilifyMk2GenerationOutput(
1432
+ sequences=sequences,
1433
+ generated_lengths=committed.clone(),
1434
+ tokens_per_forward=tokens_per_forward,
1435
+ past_key_values=past_key_values,
1436
+ stop_reason=stop_reasons[0] if batch_size == 1 else stop_reasons,
1437
+ committed_tokens=scalar_or_tensor(committed),
1438
+ denoise_steps=scalar_or_tensor(denoise_steps),
1439
+ no_progress_steps=scalar_or_tensor(state.latent_state.stagnation_steps),
1440
+ jump_count=scalar_or_tensor(jumps),
1441
+ forced_jump_bad_count=scalar_or_tensor(forced_jump_tokens),
1442
+ heavy_forward_count=scalar_or_tensor(denoise_steps),
1443
+ latent_context_update_count=scalar_or_tensor(denoise_steps),
1444
+ average_commit_len=scalar_or_tensor(average_commit_len, floating=True),
1445
+ state_shift_count=scalar_or_tensor(shifts),
1446
+ latent_memory_norm=scalar_or_tensor(latent_memory_norm, floating=True),
1447
+ state_retention_score=scalar_or_tensor(state_retention_score, floating=True),
1448
+ )
1449
+
1450
+
1451
+ __all__ = [
1452
+ "ModilifyMk2GenerationConfig", "ModilifyMk2GenerationMixin", "ModilifyMk2GenerationOutput",
1453
+ "ModilifyMk2RollingState", "NoiseCanvasSampler",
1454
+ "build_denoise_trace_event", "qualified_single_turn_end_lengths",
1455
+ ]
latent_deliberation.py ADDED
@@ -0,0 +1,2134 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2026 Modilify
2
+ # SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.0
3
+ """Dual-timescale Transformer memory for Modilify Mk2 inference.
4
+
5
+ Working trajectory state is recomputed every denoise step from the packed
6
+ personal history. Persistent slots are commit-invariant and mutate only in
7
+ TransformerCommitWriter.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ from collections.abc import Sequence
13
+ from dataclasses import dataclass
14
+ import math
15
+
16
+ import torch
17
+ from torch import nn
18
+ from torch.nn import functional as F
19
+
20
+
21
+ _AGE_MAX = 4096
22
+ _PONDER_MAX = 1024
23
+ _STAGNATION_MAX = 1024
24
+ _METADATA_HIDDEN = 64
25
+ _FILM_RANK = 64
26
+ _FOURIER_WAVES = 4
27
+ _RETIREMENT_FRAMES = 4
28
+ _KEY_META_DIM = 13
29
+ _QUERY_META_DIM = 8 + 6
30
+ _ROW_META_DIM = 2
31
+ _HISTORY_VIEWS = 4
32
+ _EXPERIENCE_ROLES = 3
33
+ _EXPERIENCE_CANVAS_STRIPE = 8
34
+ _GATE_BIAS = -3.0
35
+
36
+ COMMIT_REASON_NONE = 0
37
+ COMMIT_REASON_NORMAL = 1
38
+ COMMIT_REASON_FORCED_JUMP = 2
39
+ COMMIT_REASON_TERMINAL = 3
40
+ COMMIT_REASON_FALLBACK = 4
41
+ COMMIT_REASON_TRAINING_RANDOM = 5
42
+ COMMIT_REASON_COUNT = 6
43
+
44
+
45
+ def _fp32_scaled_dot_product_attention(
46
+ query: torch.Tensor,
47
+ key: torch.Tensor,
48
+ value: torch.Tensor,
49
+ *,
50
+ attn_mask: torch.Tensor | None = None,
51
+ ) -> torch.Tensor:
52
+ """Run latent-memory attention reductions in FP32, then restore dtype.
53
+
54
+ These attention maps are small compared with the frozen decoder, while
55
+ their outputs feed recurrent trajectory and persistent-memory paths. A
56
+ BF16 reduction error therefore compounds across denoise/commit steps and
57
+ is much more expensive than the modest FP32 workspace.
58
+ """
59
+
60
+ output_dtype = query.dtype
61
+ stable_mask = attn_mask
62
+ if stable_mask is not None and stable_mask.is_floating_point():
63
+ stable_mask = stable_mask.float()
64
+ query_fp32 = query.float()
65
+ key_fp32 = key.float()
66
+ use_batched_mm = query.ndim == 4 and key.ndim == 4 and value.ndim == 4
67
+ if use_batched_mm:
68
+ query_length = int(query.shape[-2])
69
+ key_length = int(key.shape[-2])
70
+ scores = torch.bmm(
71
+ query_fp32.reshape(-1, query_length, query.shape[-1]),
72
+ key_fp32.reshape(-1, key_length, key.shape[-1]).transpose(1, 2),
73
+ ).reshape(*query.shape[:-2], query_length, key_length)
74
+ else:
75
+ scores = torch.matmul(query_fp32, key_fp32.transpose(-2, -1))
76
+ scores = scores / math.sqrt(max(query.shape[-1], 1))
77
+ if stable_mask is not None:
78
+ if stable_mask.dtype == torch.bool:
79
+ scores = scores.masked_fill(~stable_mask, _sdpa_mask_value(torch.float32))
80
+ else:
81
+ scores = scores + stable_mask
82
+ probabilities = torch.softmax(scores, dim=-1)
83
+ # All-masked rows are uniform under a finite mask, but keep a NaN
84
+ # barrier for any remaining -inf path that MPS softmax cannot invert.
85
+ probabilities = torch.nan_to_num(probabilities, nan=0.0)
86
+ if use_batched_mm:
87
+ output = torch.bmm(
88
+ probabilities.reshape(-1, query.shape[-2], key.shape[-2]),
89
+ value.float().reshape(-1, key.shape[-2], value.shape[-1]),
90
+ ).reshape(*query.shape[:-2], query.shape[-2], value.shape[-1])
91
+ else:
92
+ output = torch.matmul(probabilities, value.float())
93
+ return output.to(dtype=output_dtype)
94
+
95
+
96
+ @dataclass
97
+ class TrajectoryHistory:
98
+ """Detached ring of recent canvas heavies. Canvas is dimension 2."""
99
+
100
+ hidden: torch.Tensor
101
+ confidence: torch.Tensor
102
+ entropy: torch.Tensor
103
+ token_changed: torch.Tensor
104
+ valid: torch.Tensor
105
+
106
+ @classmethod
107
+ def empty(
108
+ cls,
109
+ *,
110
+ batch_size: int,
111
+ canvas_length: int,
112
+ hidden_size: int,
113
+ history_length: int,
114
+ device: torch.device,
115
+ dtype: torch.dtype,
116
+ ) -> "TrajectoryHistory":
117
+ return cls(
118
+ hidden=torch.zeros(
119
+ batch_size, history_length, canvas_length, hidden_size,
120
+ device=device, dtype=dtype,
121
+ ),
122
+ confidence=torch.zeros(
123
+ batch_size, history_length, canvas_length,
124
+ device=device, dtype=torch.float32,
125
+ ),
126
+ entropy=torch.zeros(
127
+ batch_size, history_length, canvas_length,
128
+ device=device, dtype=torch.float32,
129
+ ),
130
+ token_changed=torch.zeros(
131
+ batch_size, history_length, canvas_length,
132
+ device=device, dtype=torch.float32,
133
+ ),
134
+ valid=torch.zeros(
135
+ batch_size, history_length, canvas_length,
136
+ device=device, dtype=torch.bool,
137
+ ),
138
+ )
139
+
140
+ def detach(self) -> "TrajectoryHistory":
141
+ return TrajectoryHistory(
142
+ hidden=self.hidden.detach(),
143
+ confidence=self.confidence.detach(),
144
+ entropy=self.entropy.detach(),
145
+ token_changed=self.token_changed.detach(),
146
+ valid=self.valid.detach(),
147
+ )
148
+
149
+ def append(
150
+ self,
151
+ hidden: torch.Tensor,
152
+ confidence: torch.Tensor,
153
+ entropy: torch.Tensor,
154
+ token_changed: torch.Tensor,
155
+ live_mask: torch.Tensor | None = None,
156
+ ) -> "TrajectoryHistory":
157
+ # Truncated-BPTT observation. Replay uses temporal context / committed
158
+ # memory, not this ring; keeping frames live would retain every heavy
159
+ # decoder graph across the chunk.
160
+ frame = hidden.detach()
161
+ if live_mask is None:
162
+ newest_valid = torch.ones(
163
+ self.valid.shape[0],
164
+ self.valid.shape[2],
165
+ device=self.valid.device,
166
+ dtype=torch.bool,
167
+ )
168
+ else:
169
+ newest_valid = live_mask.to(device=self.valid.device, dtype=torch.bool)
170
+ if newest_valid.shape != self.valid[:, 0].shape:
171
+ raise ValueError("`live_mask` must have shape [batch, canvas].")
172
+ return TrajectoryHistory(
173
+ hidden=torch.cat((self.hidden[:, 1:], frame.unsqueeze(1)), dim=1),
174
+ confidence=torch.cat(
175
+ (self.confidence[:, 1:], confidence.detach().float().unsqueeze(1)),
176
+ dim=1,
177
+ ),
178
+ entropy=torch.cat(
179
+ (self.entropy[:, 1:], entropy.detach().float().unsqueeze(1)),
180
+ dim=1,
181
+ ),
182
+ token_changed=torch.cat(
183
+ (self.token_changed[:, 1:], token_changed.detach().float().unsqueeze(1)),
184
+ dim=1,
185
+ ),
186
+ valid=torch.cat((self.valid[:, 1:], newest_valid.unsqueeze(1)), dim=1),
187
+ )
188
+
189
+ def shift(
190
+ self,
191
+ commit_lengths: torch.Tensor | int,
192
+ *,
193
+ entropy_fill_value: float = 0.0,
194
+ ) -> "TrajectoryHistory":
195
+ batch, history_length, canvas_length, _hidden = self.hidden.shape
196
+ if isinstance(commit_lengths, int):
197
+ lengths = torch.full(
198
+ (batch,), commit_lengths, device=self.hidden.device, dtype=torch.long
199
+ )
200
+ else:
201
+ lengths = commit_lengths.to(device=self.hidden.device, dtype=torch.long)
202
+ if bool((lengths <= 0).all()):
203
+ return self.detach()
204
+ positions = torch.arange(canvas_length, device=self.hidden.device)[None, :]
205
+ source = positions + lengths[:, None]
206
+ retained = source.lt(canvas_length)
207
+
208
+ def shifted(tensor: torch.Tensor, fill_value: float | int = 0) -> torch.Tensor:
209
+ index = source.clamp_max(canvas_length - 1)
210
+ extra = tensor.ndim - 3
211
+ view = index.view(batch, 1, canvas_length, *([1] * extra)).expand_as(tensor)
212
+ gathered = tensor.gather(2, view)
213
+ fill = torch.as_tensor(fill_value, device=tensor.device, dtype=tensor.dtype)
214
+ mask = retained.view(batch, 1, canvas_length, *([1] * extra))
215
+ return torch.where(mask, gathered, fill)
216
+
217
+ return TrajectoryHistory(
218
+ hidden=shifted(self.hidden),
219
+ confidence=shifted(self.confidence),
220
+ entropy=shifted(self.entropy, entropy_fill_value),
221
+ token_changed=shifted(self.token_changed),
222
+ valid=shifted(self.valid, False),
223
+ )
224
+
225
+
226
+ @dataclass
227
+ class TrajectoryTape:
228
+ """Row-level denoise snapshots. Time axis does not follow canvas shift."""
229
+
230
+ probes: torch.Tensor
231
+ valid: torch.Tensor
232
+
233
+ @classmethod
234
+ def empty(
235
+ cls,
236
+ *,
237
+ batch_size: int,
238
+ tape_length: int,
239
+ num_probes: int,
240
+ probe_dim: int,
241
+ device: torch.device,
242
+ dtype: torch.dtype,
243
+ ) -> "TrajectoryTape":
244
+ return cls(
245
+ probes=torch.zeros(
246
+ batch_size, tape_length, num_probes, probe_dim,
247
+ device=device, dtype=dtype,
248
+ ),
249
+ valid=torch.zeros(
250
+ batch_size, tape_length, device=device, dtype=torch.bool,
251
+ ),
252
+ )
253
+
254
+ def detach(self) -> "TrajectoryTape":
255
+ return TrajectoryTape(probes=self.probes.detach(), valid=self.valid.detach())
256
+
257
+ def append(self, probes: torch.Tensor, valid: torch.Tensor) -> "TrajectoryTape":
258
+ # Canvas snapshots are detached before pooling. Keep the pool graph so
259
+ # the next denoise's tape read can train the compressor; TBPTT still
260
+ # cuts at `detach()`.
261
+ flag = valid.to(device=self.valid.device, dtype=torch.bool)
262
+ if flag.ndim == 0:
263
+ flag = flag.expand(self.valid.shape[0])
264
+ if flag.shape != self.valid[:, 0].shape:
265
+ raise ValueError("Tape frame validity must have shape [batch].")
266
+ if probes.shape[0] != self.probes.shape[0] or probes.shape[-2:] != self.probes.shape[-2:]:
267
+ raise ValueError("Tape probes do not match the ring.")
268
+ return TrajectoryTape(
269
+ probes=torch.cat((self.probes[:, 1:], probes.unsqueeze(1)), dim=1),
270
+ valid=torch.cat((self.valid[:, 1:], flag.unsqueeze(1)), dim=1),
271
+ )
272
+
273
+
274
+ @dataclass
275
+ class LatentDeliberationState:
276
+ """Persistent slots plus per-canvas trajectory clocks. No token latents."""
277
+
278
+ memory_slots: torch.Tensor
279
+ confidence: torch.Tensor
280
+ entropy: torch.Tensor
281
+ age: torch.Tensor
282
+ token_changed: torch.Tensor
283
+ confidence_delta: torch.Tensor
284
+ entropy_delta: torch.Tensor
285
+ ponder_steps: torch.Tensor
286
+ stagnation_steps: torch.Tensor
287
+
288
+ @classmethod
289
+ def empty(
290
+ cls,
291
+ *,
292
+ batch_size: int,
293
+ canvas_length: int,
294
+ latent_dim: int,
295
+ memory_slots: int,
296
+ device: torch.device,
297
+ dtype: torch.dtype,
298
+ ) -> "LatentDeliberationState":
299
+ return cls(
300
+ memory_slots=torch.zeros(
301
+ batch_size, memory_slots, latent_dim, device=device, dtype=dtype
302
+ ),
303
+ confidence=torch.zeros(
304
+ batch_size, canvas_length, device=device, dtype=torch.float32
305
+ ),
306
+ entropy=torch.zeros(
307
+ batch_size, canvas_length, device=device, dtype=torch.float32
308
+ ),
309
+ age=torch.zeros(
310
+ batch_size, canvas_length, device=device, dtype=torch.int32
311
+ ),
312
+ token_changed=torch.zeros(
313
+ batch_size, canvas_length, device=device, dtype=torch.float32
314
+ ),
315
+ confidence_delta=torch.zeros(
316
+ batch_size, canvas_length, device=device, dtype=torch.float32
317
+ ),
318
+ entropy_delta=torch.zeros(
319
+ batch_size, canvas_length, device=device, dtype=torch.float32
320
+ ),
321
+ ponder_steps=torch.zeros(batch_size, device=device, dtype=torch.int32),
322
+ stagnation_steps=torch.zeros(batch_size, device=device, dtype=torch.int32),
323
+ )
324
+
325
+ def detach(self) -> "LatentDeliberationState":
326
+ return LatentDeliberationState(
327
+ memory_slots=self.memory_slots.detach(),
328
+ confidence=self.confidence.detach(),
329
+ entropy=self.entropy.detach(),
330
+ age=self.age.detach(),
331
+ token_changed=self.token_changed.detach(),
332
+ confidence_delta=self.confidence_delta.detach(),
333
+ entropy_delta=self.entropy_delta.detach(),
334
+ ponder_steps=self.ponder_steps.detach(),
335
+ stagnation_steps=self.stagnation_steps.detach(),
336
+ )
337
+
338
+ def shift(
339
+ self, committed: int, *, entropy_fill_value: float = 0.0
340
+ ) -> "LatentDeliberationState":
341
+ canvas_length = self.confidence.shape[1]
342
+ if not 0 <= committed <= canvas_length:
343
+ raise ValueError("`committed` must be in [0, canvas_length].")
344
+ if committed == 0:
345
+ return self.detach()
346
+
347
+ def shifted(tensor: torch.Tensor, fill_value: float | int = 0) -> torch.Tensor:
348
+ result = torch.full_like(tensor, fill_value)
349
+ if committed < canvas_length:
350
+ result[:, : canvas_length - committed] = tensor[:, committed:]
351
+ return result
352
+
353
+ return LatentDeliberationState(
354
+ memory_slots=self.memory_slots.clone(),
355
+ confidence=shifted(self.confidence),
356
+ entropy=shifted(self.entropy, entropy_fill_value),
357
+ age=shifted(self.age),
358
+ token_changed=shifted(self.token_changed),
359
+ confidence_delta=shifted(self.confidence_delta),
360
+ entropy_delta=shifted(self.entropy_delta),
361
+ ponder_steps=torch.zeros_like(self.ponder_steps),
362
+ stagnation_steps=torch.zeros_like(self.stagnation_steps),
363
+ )
364
+
365
+
366
+ @dataclass
367
+ class LatentProcessorOutput:
368
+ context: torch.Tensor
369
+ working_state: torch.Tensor
370
+ state: LatentDeliberationState
371
+ history_projected: torch.Tensor
372
+
373
+
374
+ def advance_trajectory_clocks(
375
+ ponder_steps: torch.Tensor,
376
+ stagnation_steps: torch.Tensor,
377
+ *,
378
+ commit_lengths: torch.LongTensor,
379
+ active_rows: torch.BoolTensor,
380
+ progress_scores: torch.Tensor | None = None,
381
+ min_progress: float = 0.0,
382
+ ) -> tuple[torch.IntTensor, torch.IntTensor]:
383
+ """Advance useful-ponder and stagnation clocks for each row."""
384
+
385
+ if min_progress < 0:
386
+ raise ValueError("`min_progress` must be non-negative.")
387
+ if not (
388
+ ponder_steps.shape == stagnation_steps.shape == commit_lengths.shape
389
+ == active_rows.shape
390
+ ):
391
+ raise ValueError("Trajectory clock inputs must share shape [batch].")
392
+ committed = commit_lengths.gt(0)
393
+ waiting = active_rows & ~committed
394
+ next_ponder = torch.where(
395
+ committed, torch.zeros_like(ponder_steps), ponder_steps + waiting.to(torch.int32)
396
+ )
397
+ next_stagnation = torch.where(
398
+ committed,
399
+ torch.zeros_like(stagnation_steps),
400
+ stagnation_steps + waiting.to(torch.int32),
401
+ )
402
+ return next_ponder.to(torch.int32), next_stagnation.to(torch.int32)
403
+
404
+
405
+ def should_force_trajectory_jump(
406
+ stagnation_steps: torch.Tensor,
407
+ *,
408
+ progress_scores: torch.Tensor | None = None,
409
+ min_progress: float = 0.0,
410
+ stagnation_threshold: int,
411
+ ponder_steps: torch.Tensor | None = None,
412
+ max_ponder_steps: int | None = None,
413
+ ) -> torch.BoolTensor:
414
+ """Determine whether to force a trajectory JUMP for each row.
415
+
416
+ JUMP is triggered when stagnation_steps reaches or exceeds stagnation_threshold
417
+ (default 12) and the current single-step progress is not strictly greater than
418
+ min_progress (default 0.005).
419
+ """
420
+ if stagnation_threshold <= 0:
421
+ raise ValueError("`stagnation_threshold` must be positive.")
422
+ if min_progress < 0:
423
+ raise ValueError("`min_progress` must be non-negative.")
424
+
425
+ if progress_scores is None:
426
+ stagnation_jump = stagnation_steps.ge(stagnation_threshold)
427
+ else:
428
+ if progress_scores.shape != stagnation_steps.shape:
429
+ raise ValueError("`progress_scores` must share shape with `stagnation_steps`.")
430
+ stagnation_jump = stagnation_steps.ge(stagnation_threshold) & progress_scores.le(
431
+ float(min_progress)
432
+ )
433
+
434
+ if ponder_steps is not None and max_ponder_steps is not None and max_ponder_steps > 0:
435
+ ponder_jump = ponder_steps.ge(max_ponder_steps)
436
+ return (ponder_jump | stagnation_jump).to(torch.bool)
437
+
438
+ return stagnation_jump.to(torch.bool)
439
+
440
+
441
+ def _logit01(values: torch.Tensor) -> torch.Tensor:
442
+ clipped = values.clamp(1.0e-6, 1.0 - 1.0e-6)
443
+ return torch.log(clipped) - torch.log1p(-clipped)
444
+
445
+
446
+ def _renorm_confidence(confidence: torch.Tensor) -> torch.Tensor:
447
+ return _logit01(confidence).clamp(-8.0, 8.0) / 8.0
448
+
449
+
450
+ def _canvas_fourier(canvas_length: int, device: torch.device, dtype: torch.dtype) -> torch.Tensor:
451
+ positions = torch.arange(canvas_length, device=device, dtype=dtype) / max(canvas_length, 1)
452
+ features = []
453
+ for wave in range(_FOURIER_WAVES):
454
+ angle = (2.0 ** wave) * math.pi * positions
455
+ features.append(torch.sin(angle))
456
+ features.append(torch.cos(angle))
457
+ return torch.stack(features, dim=-1)
458
+
459
+
460
+ def _safe_cosine(left: torch.Tensor, right: torch.Tensor) -> torch.Tensor:
461
+ left_n = F.normalize(left.float(), dim=-1, eps=1.0e-6)
462
+ right_n = F.normalize(right.float(), dim=-1, eps=1.0e-6)
463
+ return (left_n * right_n).sum(dim=-1).clamp(-1.0, 1.0)
464
+
465
+
466
+ def _sdpa_mask_value(dtype: torch.dtype) -> float:
467
+ """Additive SDPA mask that stays finite on MPS fp16/bf16."""
468
+
469
+ if dtype in (torch.float16, torch.bfloat16):
470
+ return -1.0e4
471
+ return -1.0e9
472
+
473
+
474
+ def _apply_rotary(payload: torch.Tensor, positions: torch.Tensor) -> torch.Tensor:
475
+ dim = payload.shape[-1]
476
+ half = dim // 2
477
+ if half == 0:
478
+ return payload
479
+ device = payload.device
480
+ inv = torch.arange(half, device=device, dtype=torch.float32)
481
+ inv = 10000.0 ** (-inv / max(half, 1))
482
+ angle = positions.to(dtype=torch.float32).unsqueeze(-1) * inv
483
+ cos = angle.cos().to(dtype=payload.dtype)
484
+ sin = angle.sin().to(dtype=payload.dtype)
485
+ while cos.ndim < payload.ndim:
486
+ cos = cos.unsqueeze(1)
487
+ sin = sin.unsqueeze(1)
488
+ left, right = payload[..., :half], payload[..., half: half * 2]
489
+ rotated = torch.cat((left * cos - right * sin, left * sin + right * cos), dim=-1)
490
+ if dim > half * 2:
491
+ rotated = torch.cat((rotated, payload[..., half * 2 :]), dim=-1)
492
+ return rotated
493
+
494
+
495
+ class _RMSNorm(nn.Module):
496
+ def __init__(self, dim: int, eps: float = 1.0e-6) -> None:
497
+ super().__init__()
498
+ self.weight = nn.Parameter(torch.ones(dim))
499
+ self.eps = eps
500
+
501
+ def forward(self, hidden: torch.Tensor) -> torch.Tensor:
502
+ rms = hidden.float().square().mean(dim=-1, keepdim=True).add(self.eps).rsqrt()
503
+ return (hidden.float() * rms).to(dtype=hidden.dtype) * self.weight.to(dtype=hidden.dtype)
504
+
505
+
506
+ class _SwiGLU(nn.Module):
507
+ def __init__(self, dim: int, hidden: int) -> None:
508
+ super().__init__()
509
+ self.gate = nn.Linear(dim, hidden, bias=False)
510
+ self.up = nn.Linear(dim, hidden, bias=False)
511
+ self.down = nn.Linear(hidden, dim, bias=False)
512
+
513
+ def forward(self, hidden: torch.Tensor) -> torch.Tensor:
514
+ return self.down(F.silu(self.gate(hidden)) * self.up(hidden))
515
+
516
+
517
+ class SharedHistoryProjector(nn.Module):
518
+ """2816 → rank projection, once per denoise (and once more on the commit tail)."""
519
+
520
+ def __init__(self, hidden_size: int, rank: int) -> None:
521
+ super().__init__()
522
+ self.norm = _RMSNorm(hidden_size)
523
+ self.proj = nn.Linear(hidden_size, rank, bias=False)
524
+
525
+ def forward(self, hidden: torch.Tensor) -> torch.Tensor:
526
+ return self.proj(self.norm(hidden))
527
+
528
+
529
+ class CanvasProbePool(nn.Module):
530
+ """Learned queries compress one canvas snapshot into P rank-space probes."""
531
+
532
+ def __init__(self, rank: int, num_probes: int, num_heads: int) -> None:
533
+ super().__init__()
534
+ if rank % num_heads:
535
+ raise ValueError("Tape rank must be divisible by heads.")
536
+ if num_probes <= 0:
537
+ raise ValueError("`num_probes` must be positive.")
538
+ self.rank = rank
539
+ self.num_probes = num_probes
540
+ self.num_heads = num_heads
541
+ self.head_dim = rank // num_heads
542
+ self.queries = nn.Parameter(torch.empty(num_probes, rank))
543
+ self.k_proj = nn.Linear(rank, rank, bias=False)
544
+ self.v_proj = nn.Linear(rank, rank, bias=False)
545
+ self.reset_parameters()
546
+
547
+ @torch.no_grad()
548
+ def reset_parameters(self) -> None:
549
+ nn.init.normal_(self.queries, mean=0.0, std=0.02)
550
+
551
+ def forward(
552
+ self,
553
+ projected: torch.Tensor,
554
+ live_mask: torch.Tensor | None = None,
555
+ ) -> torch.Tensor:
556
+ batch, canvas, _rank = projected.shape
557
+ heads = self.num_heads
558
+ head_dim = self.head_dim
559
+ query = self.queries.to(dtype=projected.dtype).view(1, self.num_probes, heads, head_dim)
560
+ query = query.expand(batch, -1, -1, -1).permute(0, 2, 1, 3)
561
+ keys = self.k_proj(projected).view(batch, canvas, heads, head_dim).transpose(1, 2)
562
+ values = self.v_proj(projected).view(batch, canvas, heads, head_dim).transpose(1, 2)
563
+ if live_mask is None:
564
+ allowed = torch.ones(batch, canvas, device=projected.device, dtype=torch.bool)
565
+ else:
566
+ allowed = live_mask.to(device=projected.device, dtype=torch.bool)
567
+ if allowed.shape != (batch, canvas):
568
+ raise ValueError("`live_mask` must have shape [batch, canvas].")
569
+ has_live = allowed.any(dim=-1)
570
+ safe = allowed.clone()
571
+ safe[:, 0] = safe[:, 0] | ~has_live
572
+ # This small P×canvas pool is a poor place to trade stability for BF16:
573
+ # on MPS the first B16 tape frame can contain NaNs even with finite Q/K/V
574
+ # and at least one unmasked key per row. Perform only the attention
575
+ # reduction in FP32; projections and stored probes retain model dtype.
576
+ additive = torch.zeros(
577
+ batch, 1, 1, canvas, device=projected.device, dtype=torch.float32
578
+ )
579
+ additive = additive.masked_fill(
580
+ ~safe.view(batch, 1, 1, canvas), _sdpa_mask_value(torch.float32)
581
+ )
582
+ context = _fp32_scaled_dot_product_attention(
583
+ query, keys, values, attn_mask=additive
584
+ )
585
+ probes = context.transpose(1, 2).reshape(batch, self.num_probes, self.rank)
586
+ return probes * has_live.to(dtype=probes.dtype).view(batch, 1, 1)
587
+
588
+
589
+ class SharedPersistentKV(nn.Module):
590
+ """Frozen-M key/value projection shared across working-processor blocks."""
591
+
592
+ def __init__(self, dim: int, num_heads: int, kv_rank: int) -> None:
593
+ super().__init__()
594
+ if kv_rank % num_heads:
595
+ raise ValueError("`kv_rank` must be divisible by `num_heads`.")
596
+ self.num_heads = num_heads
597
+ self.kv_rank = kv_rank
598
+ self.head_dim = kv_rank // num_heads
599
+ self.address_norm = _RMSNorm(dim)
600
+ self.value_norm = _RMSNorm(dim)
601
+ self.k_proj = nn.Linear(dim, kv_rank, bias=False)
602
+ self.v_proj = nn.Linear(dim, kv_rank, bias=False)
603
+
604
+ def forward(
605
+ self, memory: torch.Tensor, slot_identity: torch.Tensor
606
+ ) -> tuple[torch.Tensor, torch.Tensor]:
607
+ batch, slots, _dim = memory.shape
608
+ keys = self.k_proj(self.address_norm(memory + slot_identity))
609
+ values = self.v_proj(self.value_norm(memory))
610
+ keys = keys.view(batch, slots, self.num_heads, self.head_dim).transpose(1, 2)
611
+ values = values.view(batch, slots, self.num_heads, self.head_dim).transpose(1, 2)
612
+ return keys, values
613
+
614
+
615
+ class _HistoryAttention(nn.Module):
616
+ """Per-position attention over shared rank-space history keys."""
617
+
618
+ def __init__(self, dim: int, num_heads: int, kv_rank: int) -> None:
619
+ super().__init__()
620
+ if kv_rank % num_heads:
621
+ raise ValueError("`kv_rank` must be divisible by `num_heads`.")
622
+ self.num_heads = num_heads
623
+ self.kv_rank = kv_rank
624
+ self.head_dim = kv_rank // num_heads
625
+ self.q_proj = nn.Linear(dim, kv_rank, bias=False)
626
+ self.o_proj = nn.Linear(kv_rank, dim, bias=False)
627
+ self.q_norm = _RMSNorm(dim)
628
+
629
+ def forward(
630
+ self,
631
+ query: torch.Tensor,
632
+ keys: torch.Tensor,
633
+ values: torch.Tensor,
634
+ *,
635
+ attn_bias: torch.Tensor,
636
+ key_mask: torch.Tensor,
637
+ ) -> torch.Tensor:
638
+ batch, canvas, _dim = query.shape
639
+ heads = self.num_heads
640
+ head_dim = self.head_dim
641
+ query = self.q_proj(self.q_norm(query))
642
+ query = query.view(batch, canvas, heads, head_dim).transpose(1, 2)
643
+ if keys.ndim == 5:
644
+ expected = (batch, heads, canvas, keys.shape[-2], head_dim)
645
+ if keys.shape != expected or values.shape != expected:
646
+ raise ValueError("Preformatted history K/V dimensions do not match.")
647
+ slots = int(keys.shape[-2])
648
+ else:
649
+ slots = int(keys.shape[2])
650
+ keys = keys.view(batch, canvas, slots, heads, head_dim).permute(0, 3, 1, 2, 4)
651
+ values = values.view(batch, canvas, slots, heads, head_dim).permute(0, 3, 1, 2, 4)
652
+ query = query.reshape(batch * heads * canvas, 1, head_dim)
653
+ keys = keys.reshape(batch * heads * canvas, slots, head_dim)
654
+ values = values.reshape(batch * heads * canvas, slots, head_dim)
655
+ has_hist = key_mask.any(dim=-1)
656
+ safe_mask = key_mask.clone()
657
+ safe_mask[..., 0] = safe_mask[..., 0] | ~has_hist
658
+ mask = safe_mask.reshape(batch, 1, canvas, slots)
659
+ mask = mask.expand(-1, heads, -1, -1).reshape(batch * heads * canvas, 1, slots)
660
+ bias = attn_bias.reshape(batch * heads * canvas, 1, slots)
661
+ additive = bias.masked_fill(~mask, _sdpa_mask_value(query.dtype))
662
+ context = _fp32_scaled_dot_product_attention(
663
+ query, keys, values, attn_mask=additive
664
+ )
665
+ context = context.view(batch, heads, canvas, head_dim).transpose(1, 2).reshape(
666
+ batch, canvas, self.kv_rank
667
+ )
668
+ output = self.o_proj(context)
669
+ keep = has_hist.unsqueeze(-1)
670
+ return torch.where(keep, output, output.new_zeros(output.shape))
671
+
672
+
673
+ class _QueryOutputAttention(nn.Module):
674
+ """Q/O attention against precomputed K/V."""
675
+
676
+ def __init__(self, dim: int, num_heads: int, kv_rank: int) -> None:
677
+ super().__init__()
678
+ if kv_rank % num_heads:
679
+ raise ValueError("`kv_rank` must be divisible by `num_heads`.")
680
+ self.num_heads = num_heads
681
+ self.kv_rank = kv_rank
682
+ self.head_dim = kv_rank // num_heads
683
+ self.q_proj = nn.Linear(dim, kv_rank, bias=False)
684
+ self.o_proj = nn.Linear(kv_rank, dim, bias=False)
685
+ self.q_norm = _RMSNorm(dim)
686
+
687
+ def forward(
688
+ self,
689
+ query: torch.Tensor,
690
+ keys: torch.Tensor,
691
+ values: torch.Tensor,
692
+ attn_mask: torch.Tensor | None = None,
693
+ ) -> torch.Tensor:
694
+ batch, queries, _dim = query.shape
695
+ heads = self.num_heads
696
+ head_dim = self.head_dim
697
+ query = self.q_proj(self.q_norm(query)).view(batch, queries, heads, head_dim).transpose(1, 2)
698
+ mask = attn_mask
699
+ keep_rows = None
700
+ if mask is not None and mask.ndim == 2:
701
+ if mask.shape[0] == batch:
702
+ if mask.dtype == torch.bool:
703
+ keep_rows = mask.any(dim=-1)
704
+ safe = mask.clone()
705
+ safe[:, 0] = safe[:, 0] | ~keep_rows
706
+ mask = safe
707
+ mask = mask.view(batch, 1, 1, mask.shape[-1])
708
+ else:
709
+ mask = mask.view(1, 1, queries, keys.shape[-2])
710
+ elif mask is not None and mask.ndim == 3:
711
+ mask = mask.unsqueeze(1)
712
+ if mask is not None and mask.dtype == torch.bool:
713
+ additive = torch.zeros(
714
+ mask.shape, device=query.device, dtype=query.dtype
715
+ )
716
+ mask = additive.masked_fill(~mask, _sdpa_mask_value(query.dtype))
717
+ context = _fp32_scaled_dot_product_attention(
718
+ query, keys, values, attn_mask=mask
719
+ )
720
+ context = context.transpose(1, 2).reshape(batch, queries, self.kv_rank)
721
+ output = self.o_proj(context)
722
+ if keep_rows is not None:
723
+ output = output * keep_rows.to(dtype=output.dtype).view(batch, 1, 1)
724
+ return output
725
+
726
+
727
+ class _RankAttention(nn.Module):
728
+ """Sequence attention in a rank-``kv_rank`` subspace, then map back to ``dim``."""
729
+
730
+ def __init__(self, dim: int, num_heads: int, kv_rank: int) -> None:
731
+ super().__init__()
732
+ if kv_rank % num_heads:
733
+ raise ValueError("`kv_rank` must be divisible by `num_heads`.")
734
+ self.num_heads = num_heads
735
+ self.kv_rank = kv_rank
736
+ self.head_dim = kv_rank // num_heads
737
+ self.q_proj = nn.Linear(dim, kv_rank, bias=False)
738
+ self.k_proj = nn.Linear(dim, kv_rank, bias=False)
739
+ self.v_proj = nn.Linear(dim, kv_rank, bias=False)
740
+ self.o_proj = nn.Linear(kv_rank, dim, bias=False)
741
+ self.q_norm = _RMSNorm(dim)
742
+ self.k_norm = _RMSNorm(dim)
743
+
744
+ def forward(
745
+ self,
746
+ query: torch.Tensor,
747
+ keys: torch.Tensor,
748
+ values: torch.Tensor,
749
+ attn_mask: torch.Tensor | None = None,
750
+ ) -> torch.Tensor:
751
+ batch, queries, _dim = query.shape
752
+ key_len = keys.shape[1]
753
+ heads = self.num_heads
754
+ head_dim = self.head_dim
755
+ query = self.q_proj(self.q_norm(query)).view(batch, queries, heads, head_dim).transpose(1, 2)
756
+ keys = self.k_proj(self.k_norm(keys)).view(batch, key_len, heads, head_dim).transpose(1, 2)
757
+ values = self.v_proj(values).view(batch, key_len, heads, head_dim).transpose(1, 2)
758
+ mask = attn_mask
759
+ if mask is not None and mask.ndim == 2:
760
+ mask = mask.view(1, 1, queries, key_len)
761
+ elif mask is not None and mask.ndim == 3:
762
+ mask = mask.unsqueeze(1)
763
+ context = _fp32_scaled_dot_product_attention(
764
+ query, keys, values, attn_mask=mask
765
+ )
766
+ context = context.transpose(1, 2).reshape(batch, queries, self.kv_rank)
767
+ return self.o_proj(context)
768
+
769
+
770
+ class _ProcessorBlock(nn.Module):
771
+ """History-read canvas state, local or global mixing, read-only slot CA, SwiGLU."""
772
+
773
+ def __init__(
774
+ self,
775
+ dim: int,
776
+ num_heads: int,
777
+ local_attention_window: int,
778
+ kv_rank: int,
779
+ ffn_dim: int,
780
+ *,
781
+ global_attention: bool,
782
+ ) -> None:
783
+ super().__init__()
784
+ self.history_attention = _HistoryAttention(dim, num_heads, kv_rank)
785
+ self.state_norm = _RMSNorm(dim)
786
+ self.local_attention = _RankAttention(dim, num_heads, kv_rank)
787
+ self.local_attention_window = local_attention_window
788
+ self.global_attention = global_attention
789
+ self.register_buffer("_local_attention_mask", torch.empty(0), persistent=False)
790
+ self.token_memory_attention = _QueryOutputAttention(dim, num_heads, kv_rank)
791
+ self.tape_attention = _QueryOutputAttention(dim, num_heads, kv_rank)
792
+ self.token_ff_norm = _RMSNorm(dim)
793
+ self.ff = _SwiGLU(dim, ffn_dim)
794
+
795
+ def _local_mask(self, canvas: int, device: torch.device, dtype: torch.dtype) -> torch.Tensor:
796
+ if (
797
+ self._local_attention_mask.shape != (canvas, canvas)
798
+ or self._local_attention_mask.device != device
799
+ or self._local_attention_mask.dtype != dtype
800
+ ):
801
+ positions = torch.arange(canvas, device=device)
802
+ allowed = (positions[:, None] - positions[None, :]).abs() < self.local_attention_window
803
+ mask = torch.zeros(canvas, canvas, device=device, dtype=dtype)
804
+ self._local_attention_mask = mask.masked_fill(~allowed, _sdpa_mask_value(dtype))
805
+ return self._local_attention_mask
806
+
807
+ def forward(
808
+ self,
809
+ canvas_state: torch.Tensor,
810
+ history_keys: torch.Tensor,
811
+ history_values: torch.Tensor,
812
+ attn_bias: torch.Tensor,
813
+ key_mask: torch.Tensor,
814
+ memory_keys: torch.Tensor,
815
+ memory_values: torch.Tensor,
816
+ history_query: torch.Tensor,
817
+ tape_keys: torch.Tensor | None = None,
818
+ tape_values: torch.Tensor | None = None,
819
+ tape_mask: torch.Tensor | None = None,
820
+ ) -> torch.Tensor:
821
+ canvas_state = canvas_state + self.history_attention(
822
+ history_query, history_keys, history_values,
823
+ attn_bias=attn_bias, key_mask=key_mask,
824
+ )
825
+ if tape_keys is not None and tape_values is not None:
826
+ canvas_state = canvas_state + self.tape_attention(
827
+ self.state_norm(canvas_state), tape_keys, tape_values, attn_mask=tape_mask
828
+ )
829
+ normalized = self.state_norm(canvas_state)
830
+ attn_mask = None if self.global_attention else self._local_mask(
831
+ canvas_state.shape[1], canvas_state.device, canvas_state.dtype
832
+ )
833
+ canvas_state = canvas_state + self.local_attention(
834
+ normalized, normalized, normalized, attn_mask=attn_mask
835
+ )
836
+ normalized = self.state_norm(canvas_state)
837
+ canvas_state = canvas_state + self.token_memory_attention(
838
+ normalized, memory_keys, memory_values
839
+ )
840
+ return canvas_state + self.ff(self.token_ff_norm(canvas_state))
841
+
842
+
843
+ class DecoderMemoryBus(nn.Module):
844
+ """Sidecar readers on full-attention decoder layers.
845
+
846
+ ``alpha`` starts at 0 so the residual is zero. Scale with ``tanh(alpha)``
847
+ rather than a hard ``where(alpha == 0)`` so the gate stays differentiable.
848
+ ``o_proj`` is *not* zeroed: that pair was a dead-gradient product.
849
+ """
850
+
851
+ def __init__(
852
+ self,
853
+ hidden_size: int,
854
+ num_heads: int,
855
+ num_readers: int,
856
+ memory_dim: int,
857
+ kv_rank: int,
858
+ *,
859
+ relative_bias: bool = False,
860
+ address_with_identity: bool = False,
861
+ max_relative_span: int = 256,
862
+ ) -> None:
863
+ super().__init__()
864
+ if num_readers > 0:
865
+ if kv_rank % num_heads:
866
+ raise ValueError("Memory bus rank must be divisible by heads.")
867
+ if hidden_size % num_heads:
868
+ raise ValueError("Memory bus hidden size must be divisible by heads.")
869
+ self.hidden_size = hidden_size
870
+ self.num_heads = num_heads
871
+ self.num_readers = num_readers
872
+ self.kv_rank = kv_rank
873
+ self.head_dim = kv_rank // num_heads if num_heads else kv_rank
874
+ self.relative_bias = relative_bias
875
+ self.address_with_identity = address_with_identity
876
+ self.memory_norm = _RMSNorm(memory_dim)
877
+ self.address_norm = _RMSNorm(memory_dim)
878
+ self.memory_to_hidden = (
879
+ nn.Identity()
880
+ if memory_dim == hidden_size
881
+ else nn.Linear(memory_dim, hidden_size, bias=False)
882
+ )
883
+ self.k_proj = nn.Linear(hidden_size, kv_rank, bias=False)
884
+ self.v_proj = nn.Linear(hidden_size, kv_rank, bias=False)
885
+ self.q_norm = _RMSNorm(hidden_size)
886
+ self.q_proj = nn.ModuleList(
887
+ [nn.Linear(hidden_size, kv_rank, bias=False) for _ in range(num_readers)]
888
+ )
889
+ self.o_proj = nn.ModuleList(
890
+ [nn.Linear(kv_rank, hidden_size, bias=False) for _ in range(num_readers)]
891
+ )
892
+ self.alpha = nn.Parameter(torch.zeros(max(num_readers, 1), max(num_heads, 1)))
893
+ span = max(2 * max_relative_span - 1, 1)
894
+ self.rel_bias = nn.Parameter(torch.zeros(max(num_heads, 1), span))
895
+ self.max_relative_span = max_relative_span
896
+ self.enabled = False
897
+ self.freeze()
898
+
899
+ def freeze(self) -> None:
900
+ self.enabled = False
901
+ for parameter in self.parameters():
902
+ parameter.requires_grad_(False)
903
+
904
+ def unfreeze(self) -> None:
905
+ if self.num_readers <= 0:
906
+ return
907
+ self.enabled = True
908
+ for parameter in self.parameters():
909
+ parameter.requires_grad_(True)
910
+
911
+ def prepare_kv(
912
+ self,
913
+ memory: torch.Tensor,
914
+ slot_identity: torch.Tensor | None = None,
915
+ ) -> tuple[torch.Tensor, torch.Tensor] | None:
916
+ if self.num_readers <= 0:
917
+ return None
918
+ if self.training and not self.enabled:
919
+ return None
920
+ if self.address_with_identity:
921
+ if slot_identity is None:
922
+ raise ValueError("Persistent bus requires slot identity on keys.")
923
+ mapped_keys = self.memory_to_hidden(self.address_norm(memory + slot_identity))
924
+ mapped_values = self.memory_to_hidden(self.memory_norm(memory))
925
+ else:
926
+ mapped_keys = mapped_values = self.memory_to_hidden(self.memory_norm(memory))
927
+ batch, slots, _dim = mapped_keys.shape
928
+ heads = self.num_heads
929
+ head_dim = self.head_dim
930
+ keys = self.k_proj(mapped_keys).view(batch, slots, heads, head_dim).transpose(1, 2)
931
+ values = self.v_proj(mapped_values).view(batch, slots, heads, head_dim).transpose(1, 2)
932
+ return keys, values
933
+
934
+ def _relative_mask(self, queries: int, keys: int, device: torch.device, dtype: torch.dtype) -> torch.Tensor | None:
935
+ if not self.relative_bias:
936
+ return None
937
+ q = torch.arange(queries, device=device)
938
+ k = torch.arange(keys, device=device)
939
+ rel = (q[:, None] - k[None, :] + (keys - 1)).clamp(0, self.rel_bias.shape[1] - 1)
940
+ return self.rel_bias[:, rel].to(dtype=dtype)
941
+
942
+ def read(
943
+ self,
944
+ hidden: torch.Tensor,
945
+ reader_index: int,
946
+ keys: torch.Tensor,
947
+ values: torch.Tensor,
948
+ ) -> torch.Tensor:
949
+ batch, canvas, _dim = hidden.shape
950
+ heads = self.num_heads
951
+ head_dim = self.head_dim
952
+ query = self.q_proj[reader_index](self.q_norm(hidden))
953
+ query = query.view(batch, canvas, heads, head_dim).transpose(1, 2)
954
+ bias = self._relative_mask(canvas, keys.shape[2], hidden.device, query.dtype)
955
+ if bias is not None:
956
+ bias = bias.unsqueeze(0)
957
+ context = _fp32_scaled_dot_product_attention(
958
+ query, keys, values, attn_mask=bias
959
+ )
960
+ scale = torch.tanh(self.alpha[reader_index]).to(dtype=hidden.dtype).view(1, heads, 1, 1)
961
+ context = context * scale
962
+ context = context.transpose(1, 2).reshape(batch, canvas, self.kv_rank)
963
+ return hidden + self.o_proj[reader_index](context)
964
+
965
+ @torch.no_grad()
966
+ def reset_identity_parameters(self) -> None:
967
+ self.alpha.zero_()
968
+ self.rel_bias.zero_()
969
+
970
+
971
+ class _SequenceBlock(nn.Module):
972
+ def __init__(self, dim: int, num_heads: int, ffn_dim: int) -> None:
973
+ super().__init__()
974
+ self.attn = _RankAttention(dim, num_heads, dim)
975
+ self.norm = _RMSNorm(dim)
976
+ self.ff_norm = _RMSNorm(dim)
977
+ self.ff = _SwiGLU(dim, ffn_dim)
978
+
979
+ def forward(self, hidden: torch.Tensor, attn_mask: torch.Tensor | None) -> torch.Tensor:
980
+ normalized = self.norm(hidden)
981
+ hidden = hidden + self.attn(normalized, normalized, normalized, attn_mask=attn_mask)
982
+ return hidden + self.ff(self.ff_norm(hidden))
983
+
984
+
985
+ class ExperienceRoleEncoder(nn.Module):
986
+ """Three role queries over a committed token's packed rank-space trajectory."""
987
+
988
+ def __init__(self, rank: int, num_heads: int, hidden_size: int) -> None:
989
+ super().__init__()
990
+ self.rank = rank
991
+ self.num_heads = num_heads
992
+ self.head_dim = rank // num_heads
993
+ self.role_queries = nn.Parameter(torch.empty(_EXPERIENCE_ROLES, rank))
994
+ self.working_proj = nn.Linear(hidden_size, rank, bias=False)
995
+ self.q_proj = nn.Linear(rank, rank, bias=False)
996
+ self.k_proj = nn.Linear(rank, rank, bias=False)
997
+ self.v_proj = nn.Linear(rank, rank, bias=False)
998
+ self.o_proj = nn.Linear(rank, rank, bias=False)
999
+ self.reason_embed = nn.Embedding(COMMIT_REASON_COUNT, rank)
1000
+ self.reset_parameters()
1001
+
1002
+ @torch.no_grad()
1003
+ def reset_parameters(self) -> None:
1004
+ nn.init.normal_(self.role_queries, mean=0.0, std=0.02)
1005
+
1006
+ def forward(
1007
+ self,
1008
+ *,
1009
+ working_state: torch.Tensor,
1010
+ history_keys: torch.Tensor,
1011
+ history_values: torch.Tensor,
1012
+ history_mask: torch.Tensor,
1013
+ z_final: torch.Tensor,
1014
+ commit_reason: torch.Tensor,
1015
+ ) -> torch.Tensor:
1016
+ batch, canvas, slots, rank = history_keys.shape
1017
+ working = self.working_proj(working_state)
1018
+ extra = torch.stack((z_final, working), dim=2)
1019
+ keys = torch.cat((history_keys, extra), dim=2)
1020
+ values = torch.cat((history_values, extra), dim=2)
1021
+ extra_mask = torch.ones(batch, canvas, 2, device=history_mask.device, dtype=torch.bool)
1022
+ mask = torch.cat((history_mask, extra_mask), dim=2)
1023
+ roles = self.role_queries.to(dtype=keys.dtype).view(1, 1, _EXPERIENCE_ROLES, rank)
1024
+ roles = roles.expand(batch, canvas, -1, -1)
1025
+ reason = self.reason_embed(commit_reason.clamp(0, COMMIT_REASON_COUNT - 1))
1026
+ roles = roles + reason.to(dtype=roles.dtype).view(batch, 1, 1, rank)
1027
+ heads = self.num_heads
1028
+ head_dim = self.head_dim
1029
+ query = self.q_proj(roles).view(batch, canvas, _EXPERIENCE_ROLES, heads, head_dim)
1030
+ query = query.permute(0, 3, 1, 2, 4).reshape(
1031
+ batch * heads * canvas, _EXPERIENCE_ROLES, head_dim
1032
+ )
1033
+ key = self.k_proj(keys).view(batch, canvas, slots + 2, heads, head_dim)
1034
+ key = key.permute(0, 3, 1, 2, 4).reshape(batch * heads * canvas, slots + 2, head_dim)
1035
+ value = self.v_proj(values).view(batch, canvas, slots + 2, heads, head_dim)
1036
+ value = value.permute(0, 3, 1, 2, 4).reshape(batch * heads * canvas, slots + 2, head_dim)
1037
+ attn_mask = mask.view(batch, 1, canvas, 1, slots + 2)
1038
+ attn_mask = attn_mask.expand(-1, heads, -1, _EXPERIENCE_ROLES, -1)
1039
+ attn_mask = attn_mask.reshape(batch * heads * canvas, _EXPERIENCE_ROLES, slots + 2)
1040
+ additive = torch.zeros(
1041
+ query.shape[0],
1042
+ query.shape[1],
1043
+ key.shape[1],
1044
+ device=query.device,
1045
+ dtype=query.dtype,
1046
+ )
1047
+ additive = additive.masked_fill(~attn_mask, _sdpa_mask_value(query.dtype))
1048
+ context = _fp32_scaled_dot_product_attention(
1049
+ query, key, value, attn_mask=additive
1050
+ )
1051
+ context = context.view(batch, heads, canvas, _EXPERIENCE_ROLES, head_dim)
1052
+ context = context.permute(0, 2, 3, 1, 4).reshape(batch, canvas, _EXPERIENCE_ROLES, rank)
1053
+ return self.o_proj(context)
1054
+
1055
+
1056
+ class CommitSequenceTransformer(nn.Module):
1057
+ """Bidirectional phrase-level mixer over 3L experience role tokens."""
1058
+
1059
+ def __init__(self, dim: int, num_heads: int, num_layers: int, ffn_dim: int) -> None:
1060
+ super().__init__()
1061
+ self.role_embed = nn.Embedding(_EXPERIENCE_ROLES, dim)
1062
+ self.blocks = nn.ModuleList(
1063
+ [_SequenceBlock(dim, num_heads, ffn_dim) for _ in range(num_layers)]
1064
+ )
1065
+ self.norm = _RMSNorm(dim)
1066
+
1067
+ def forward(
1068
+ self,
1069
+ tokens: torch.Tensor,
1070
+ positions: torch.Tensor,
1071
+ valid: torch.Tensor,
1072
+ ) -> torch.Tensor:
1073
+ batch, length, dim = tokens.shape
1074
+ roles = torch.arange(_EXPERIENCE_ROLES, device=tokens.device).repeat(length // _EXPERIENCE_ROLES + 1)
1075
+ roles = roles[:length]
1076
+ hidden = tokens + self.role_embed(roles).to(dtype=tokens.dtype)
1077
+ heads = self.blocks[0].attn.num_heads
1078
+ head_dim = dim // heads
1079
+ hidden = hidden.view(batch, length, heads, head_dim).transpose(1, 2)
1080
+ hidden = _apply_rotary(hidden, positions)
1081
+ hidden = hidden.transpose(1, 2).reshape(batch, length, dim)
1082
+ keep_rows = valid.any(dim=-1)
1083
+ safe = valid.clone()
1084
+ if length > 0:
1085
+ safe[:, 0] = safe[:, 0] | ~keep_rows
1086
+ keep = safe.unsqueeze(1) & safe.unsqueeze(2)
1087
+ attn_mask = torch.zeros(
1088
+ batch, length, length, device=tokens.device, dtype=tokens.dtype
1089
+ )
1090
+ attn_mask = attn_mask.masked_fill(
1091
+ ~keep, _sdpa_mask_value(tokens.dtype)
1092
+ )
1093
+ for block in self.blocks:
1094
+ hidden = block(hidden, attn_mask)
1095
+ hidden = self.norm(hidden)
1096
+ return torch.where(valid.unsqueeze(-1), hidden, hidden.new_zeros(hidden.shape))
1097
+
1098
+
1099
+ class TransformerCommitWriter(nn.Module):
1100
+ """Identity-init slot-gated persistent write."""
1101
+
1102
+ def __init__(
1103
+ self,
1104
+ dim: int,
1105
+ rank: int,
1106
+ num_heads: int,
1107
+ ffn_dim: int,
1108
+ experience_dim: int,
1109
+ ) -> None:
1110
+ super().__init__()
1111
+ self.cross = _RankAttention(dim, num_heads, rank)
1112
+ self.experience_up = (
1113
+ nn.Identity()
1114
+ if experience_dim == dim
1115
+ else nn.Linear(experience_dim, dim, bias=False)
1116
+ )
1117
+ self.self_attn = _RankAttention(dim, num_heads, rank)
1118
+ self.ff = _SwiGLU(dim, ffn_dim)
1119
+ self.ff_norm = _RMSNorm(dim)
1120
+ self.norm = _RMSNorm(dim)
1121
+ self.gate = nn.Linear(dim * 2, 1, bias=True)
1122
+ self.beta_write = nn.Parameter(torch.zeros(()))
1123
+ self.gamma_ca = nn.Parameter(torch.tensor(0.1))
1124
+ self.gamma_sa = nn.Parameter(torch.zeros(()))
1125
+ self.gamma_ffn = nn.Parameter(torch.tensor(0.1))
1126
+ self.reset_identity_parameters()
1127
+
1128
+ @torch.no_grad()
1129
+ def reset_identity_parameters(self) -> None:
1130
+ nn.init.zeros_(self.gate.weight)
1131
+ nn.init.constant_(self.gate.bias, _GATE_BIAS)
1132
+ self.beta_write.zero_()
1133
+ self.gamma_sa.zero_()
1134
+ self.gamma_ca.copy_(self.gamma_ca.new_tensor(0.1))
1135
+ self.gamma_ffn.copy_(self.gamma_ffn.new_tensor(0.1))
1136
+
1137
+ def forward(
1138
+ self,
1139
+ memory: torch.Tensor,
1140
+ experience: torch.Tensor,
1141
+ experience_mask: torch.Tensor,
1142
+ ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
1143
+ mapped = self.experience_up(experience)
1144
+ row_has_experience = experience_mask.any(dim=-1)
1145
+ safe_mask = experience_mask
1146
+ if experience_mask.shape[-1] > 0:
1147
+ safe_mask = experience_mask.clone()
1148
+ safe_mask[:, 0] = safe_mask[:, 0] | ~row_has_experience
1149
+ keep = safe_mask.unsqueeze(1) & torch.ones(
1150
+ memory.shape[0], memory.shape[1], 1, device=memory.device, dtype=torch.bool
1151
+ )
1152
+ attn_mask = torch.zeros(
1153
+ memory.shape[0],
1154
+ memory.shape[1],
1155
+ mapped.shape[1],
1156
+ device=memory.device,
1157
+ dtype=memory.dtype,
1158
+ )
1159
+ attn_mask = attn_mask.masked_fill(~keep, _sdpa_mask_value(memory.dtype))
1160
+ delta_ca = self.cross(self.norm(memory), mapped, mapped, attn_mask=attn_mask)
1161
+ hidden = memory + self.gamma_ca.to(dtype=memory.dtype) * delta_ca
1162
+ delta_sa = self.self_attn(self.norm(hidden), hidden, hidden)
1163
+ hidden = hidden + self.gamma_sa.to(dtype=memory.dtype) * delta_sa
1164
+ delta_ff = self.ff(self.ff_norm(hidden))
1165
+ proposed = hidden + self.gamma_ffn.to(dtype=memory.dtype) * delta_ff
1166
+ delta = proposed - memory
1167
+ gate = torch.sigmoid(self.gate(torch.cat((memory, proposed), dim=-1)))
1168
+ scale = torch.tanh(self.beta_write).to(dtype=memory.dtype)
1169
+ written = memory + scale * gate * delta
1170
+ written = torch.where(row_has_experience.view(memory.shape[0], 1, 1), written, memory)
1171
+ return written, gate, delta
1172
+
1173
+
1174
+ class LatentDeliberationTransformer(nn.Module):
1175
+ """Working trajectory processor plus commit-only persistent writer."""
1176
+
1177
+ def __init__(
1178
+ self,
1179
+ *,
1180
+ hidden_size: int,
1181
+ vocab_size: int,
1182
+ latent_dim: int = 2816,
1183
+ ffn_dim: int = 7168,
1184
+ memory_slots: int = 256,
1185
+ num_layers: int = 4,
1186
+ num_heads: int = 16,
1187
+ local_attention_window: int = 128,
1188
+ dropout: float = 0.0,
1189
+ history_length: int = 16,
1190
+ tape_probes: int = 16,
1191
+ history_kv_rank: int = 1024,
1192
+ num_memory_readers: int = 0,
1193
+ num_working_readers: int | None = None,
1194
+ num_persistent_readers: int | None = None,
1195
+ working_last_block_global: bool = True,
1196
+ experience_roles: int = _EXPERIENCE_ROLES,
1197
+ commit_sequence_layers: int = 2,
1198
+ commit_sequence_dim: int | None = None,
1199
+ writer_ffn_dim: int | None = None,
1200
+ max_canvas_length: int = 256,
1201
+ **kwargs: Any,
1202
+ ) -> None:
1203
+ super().__init__()
1204
+ del dropout, experience_roles
1205
+ if latent_dim % num_heads:
1206
+ raise ValueError("`latent_dim` must be divisible by `num_heads`.")
1207
+ if local_attention_window <= 0:
1208
+ raise ValueError("`local_attention_window` must be positive.")
1209
+ if history_length <= 0:
1210
+ raise ValueError("`history_length` must be positive.")
1211
+ if ffn_dim <= 0:
1212
+ raise ValueError("`ffn_dim` must be positive.")
1213
+ if history_kv_rank % num_heads or history_kv_rank > latent_dim:
1214
+ raise ValueError("Invalid history K/V rank.")
1215
+ self.hidden_size = hidden_size
1216
+ self.vocab_size = vocab_size
1217
+ self.latent_dim = latent_dim
1218
+ self.memory_slots = memory_slots
1219
+ self.history_length = history_length
1220
+ self.tape_probes = int(tape_probes)
1221
+ if self.tape_probes <= 0:
1222
+ raise ValueError("`tape_probes` must be positive.")
1223
+ self.history_views = _HISTORY_VIEWS
1224
+ self.log_vocab = math.log(max(vocab_size, 2))
1225
+ packet_dim = int(commit_sequence_dim or history_kv_rank)
1226
+ if packet_dim % num_heads:
1227
+ raise ValueError("`commit_sequence_dim` must be divisible by `num_heads`.")
1228
+ self.packet_dim = packet_dim
1229
+ self.history_in = (
1230
+ nn.Identity()
1231
+ if hidden_size == latent_dim
1232
+ else nn.Linear(hidden_size, latent_dim, bias=False)
1233
+ )
1234
+ self.query_in = (
1235
+ nn.Identity()
1236
+ if hidden_size == latent_dim
1237
+ else nn.Linear(hidden_size, latent_dim, bias=False)
1238
+ )
1239
+ self.history_projector = SharedHistoryProjector(latent_dim, history_kv_rank)
1240
+ self.tape_pool = CanvasProbePool(history_kv_rank, self.tape_probes, num_heads)
1241
+ self.persistent_kv = SharedPersistentKV(latent_dim, num_heads, history_kv_rank)
1242
+ self.bias_in = nn.Linear(
1243
+ _KEY_META_DIM + _QUERY_META_DIM + _ROW_META_DIM, _METADATA_HIDDEN, bias=True
1244
+ )
1245
+ self.bias_out = nn.Linear(_METADATA_HIDDEN, num_heads, bias=True)
1246
+ nn.init.zeros_(self.bias_out.weight)
1247
+ nn.init.zeros_(self.bias_out.bias)
1248
+ self.film_in = nn.Linear(_KEY_META_DIM, _FILM_RANK, bias=True)
1249
+ self.film_out = nn.Linear(_FILM_RANK, 2 * history_kv_rank, bias=True)
1250
+ nn.init.zeros_(self.film_out.weight)
1251
+ nn.init.zeros_(self.film_out.bias)
1252
+ self.blocks = nn.ModuleList(
1253
+ [
1254
+ _ProcessorBlock(
1255
+ latent_dim,
1256
+ num_heads,
1257
+ local_attention_window,
1258
+ history_kv_rank,
1259
+ ffn_dim,
1260
+ global_attention=bool(
1261
+ working_last_block_global and index == num_layers - 1
1262
+ ),
1263
+ )
1264
+ for index in range(num_layers)
1265
+ ]
1266
+ )
1267
+ self.output_norm = _RMSNorm(latent_dim)
1268
+ self.output_to_hidden = (
1269
+ nn.Identity()
1270
+ if hidden_size == latent_dim
1271
+ else nn.Linear(latent_dim, hidden_size, bias=False)
1272
+ )
1273
+ self.memory_slot_identity = nn.Parameter(torch.empty(memory_slots, latent_dim))
1274
+ bus_heads = num_heads if hidden_size % num_heads == 0 else 1
1275
+ bus_rank = history_kv_rank if history_kv_rank % bus_heads == 0 else bus_heads
1276
+ working_readers = num_memory_readers if num_working_readers is None else num_working_readers
1277
+ persistent_readers = (
1278
+ num_memory_readers if num_persistent_readers is None else num_persistent_readers
1279
+ )
1280
+ self.working_memory_bus = DecoderMemoryBus(
1281
+ hidden_size=hidden_size,
1282
+ num_heads=bus_heads,
1283
+ num_readers=working_readers,
1284
+ memory_dim=hidden_size,
1285
+ kv_rank=bus_rank,
1286
+ relative_bias=True,
1287
+ address_with_identity=False,
1288
+ max_relative_span=max_canvas_length,
1289
+ )
1290
+ self.persistent_memory_bus = DecoderMemoryBus(
1291
+ hidden_size=hidden_size,
1292
+ num_heads=bus_heads,
1293
+ num_readers=persistent_readers,
1294
+ memory_dim=latent_dim,
1295
+ kv_rank=bus_rank,
1296
+ relative_bias=False,
1297
+ address_with_identity=True,
1298
+ max_relative_span=max_canvas_length,
1299
+ )
1300
+ self.experience_encoder = ExperienceRoleEncoder(
1301
+ packet_dim, num_heads, hidden_size
1302
+ )
1303
+ self.commit_sequence = CommitSequenceTransformer(
1304
+ packet_dim,
1305
+ num_heads,
1306
+ commit_sequence_layers,
1307
+ max(packet_dim * 2, packet_dim),
1308
+ )
1309
+ self.commit_writer = TransformerCommitWriter(
1310
+ latent_dim,
1311
+ history_kv_rank,
1312
+ num_heads,
1313
+ writer_ffn_dim or ffn_dim,
1314
+ packet_dim,
1315
+ )
1316
+ self.reset_identity_parameters()
1317
+
1318
+ @property
1319
+ def memory_bus(self) -> DecoderMemoryBus:
1320
+ return self.persistent_memory_bus
1321
+
1322
+ @torch.no_grad()
1323
+ def reset_memory_slot_identity(self) -> None:
1324
+ workspace = torch.empty_like(self.memory_slot_identity, dtype=torch.float32)
1325
+ if self.memory_slots <= self.latent_dim:
1326
+ nn.init.orthogonal_(workspace)
1327
+ else:
1328
+ nn.init.normal_(workspace, mean=0.0, std=1.0)
1329
+ workspace = F.normalize(workspace, dim=-1)
1330
+ self.memory_slot_identity.copy_(workspace.to(dtype=self.memory_slot_identity.dtype))
1331
+
1332
+ @torch.no_grad()
1333
+ def reset_identity_parameters(self) -> None:
1334
+ """Initialize direct parameters after generic PreTrainedModel init."""
1335
+
1336
+ self.reset_memory_slot_identity()
1337
+ # These direct Parameters are created on `meta` during low-memory
1338
+ # from_pretrained loading. Generic initialization covers Linear,
1339
+ # Embedding, and RMSNorm modules, but not standalone query tensors.
1340
+ self.tape_pool.reset_parameters()
1341
+ self.experience_encoder.reset_parameters()
1342
+ nn.init.zeros_(self.bias_out.weight)
1343
+ nn.init.zeros_(self.bias_out.bias)
1344
+ nn.init.zeros_(self.film_out.weight)
1345
+ nn.init.zeros_(self.film_out.bias)
1346
+ self.commit_writer.reset_identity_parameters()
1347
+ self.working_memory_bus.reset_identity_parameters()
1348
+ self.persistent_memory_bus.reset_identity_parameters()
1349
+
1350
+ def scaled_memory_slot_identity(
1351
+ self,
1352
+ *,
1353
+ batch_size: int,
1354
+ device: torch.device,
1355
+ dtype: torch.dtype,
1356
+ ) -> torch.Tensor:
1357
+ identity = F.normalize(self.memory_slot_identity.float(), dim=-1)
1358
+ identity = identity * math.sqrt(self.latent_dim)
1359
+ return identity.to(device=device, dtype=dtype).unsqueeze(0).expand(
1360
+ batch_size, -1, -1
1361
+ )
1362
+
1363
+ def project_context(self, canvas_state: torch.Tensor) -> torch.Tensor:
1364
+ return self.output_to_hidden(self.output_norm(canvas_state))
1365
+
1366
+ def encode_tape_frame(
1367
+ self,
1368
+ hidden: torch.Tensor,
1369
+ live_mask: torch.Tensor | None = None,
1370
+ ) -> tuple[torch.Tensor, torch.Tensor]:
1371
+ """Compress a detached canvas snapshot into tape probes."""
1372
+
1373
+ projected = self.history_projector(self.history_in(hidden.detach()))
1374
+ if not bool(torch.isfinite(projected.detach()).all()):
1375
+ raise FloatingPointError("Non-finite trajectory tape projection.")
1376
+ probes = self.tape_pool(projected, live_mask)
1377
+ # Materialize the MPS FP32 reduction before the probes enter the
1378
+ # recurrent ring. Besides fail-fast validation, this is a required
1379
+ # producer/consumer barrier for the next BF16 denoise on MPS.
1380
+ if not bool(torch.isfinite(probes.detach()).all()):
1381
+ raise FloatingPointError("Non-finite trajectory tape probes.")
1382
+ if live_mask is None:
1383
+ valid = torch.ones(hidden.shape[0], device=hidden.device, dtype=torch.bool)
1384
+ else:
1385
+ valid = live_mask.to(device=hidden.device, dtype=torch.bool).any(dim=-1)
1386
+ return probes, valid
1387
+
1388
+ def _tape_keys(
1389
+ self,
1390
+ tape: TrajectoryTape,
1391
+ dtype: torch.dtype,
1392
+ ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor] | tuple[None, None, None]:
1393
+ if not bool(tape.valid.any()):
1394
+ return None, None, None
1395
+ rank = self.tape_pool.rank
1396
+ heads = self.tape_pool.num_heads
1397
+ head_dim = self.tape_pool.head_dim
1398
+ batch, tape_length, probes, probe_dim = tape.probes.shape
1399
+ if probe_dim != rank:
1400
+ raise ValueError("Tape probe width does not match the projector rank.")
1401
+ keys = tape.probes.to(dtype=dtype).reshape(batch, tape_length * probes, heads, head_dim)
1402
+ keys = keys.transpose(1, 2)
1403
+ mask = tape.valid.unsqueeze(-1).expand(-1, -1, probes).reshape(batch, tape_length * probes)
1404
+ return keys, keys, mask
1405
+
1406
+ def _geometry(self, history: TrajectoryHistory) -> dict[str, torch.Tensor]:
1407
+ valid = history.valid
1408
+ batch, history_length, canvas, hidden_size = history.hidden.shape
1409
+ velocity = torch.zeros(
1410
+ batch, history_length, canvas,
1411
+ device=history.hidden.device, dtype=torch.float32,
1412
+ )
1413
+ acceleration = torch.zeros_like(velocity)
1414
+ reversal = torch.zeros_like(velocity)
1415
+ recurrence = torch.zeros_like(velocity)
1416
+ osc = torch.zeros_like(velocity)
1417
+ scale = math.sqrt(max(hidden_size, 1))
1418
+
1419
+ # Geometry is diagnostic conditioning over a detached history ring.
1420
+ # Processing one time edge at a time keeps only three FP32 frames and
1421
+ # two deltas live instead of materializing FP32 hidden/delta/accel for
1422
+ # the complete [B, H, C, D] ring. Each scalar uses the same FP32
1423
+ # subtraction, norm, and cosine operations as the dense formulation.
1424
+ if history_length > 1:
1425
+ previous_previous: torch.Tensor | None = None
1426
+ previous = history.hidden[:, 0].float()
1427
+ previous_delta: torch.Tensor | None = None
1428
+ for index in range(1, history_length):
1429
+ current = history.hidden[:, index].float()
1430
+ delta = current - previous
1431
+ velocity[:, index] = delta.norm(dim=-1) / scale
1432
+ if previous_previous is not None and previous_delta is not None:
1433
+ accel = current - 2.0 * previous + previous_previous
1434
+ acceleration[:, index] = accel.norm(dim=-1) / scale
1435
+ reversal[:, index] = -_safe_cosine(delta, previous_delta)
1436
+ recurrence[:, index] = _safe_cosine(current, previous_previous)
1437
+ osc[:, index] = (
1438
+ recurrence[:, index] - _safe_cosine(current, previous)
1439
+ )
1440
+ previous_previous = previous
1441
+ previous = current
1442
+ previous_delta = delta
1443
+ raw_mask = valid
1444
+ delta_mask = valid.clone()
1445
+ delta_mask[:, 0] = False
1446
+ if valid.shape[1] > 1:
1447
+ delta_mask[:, 1:] = valid[:, 1:] & valid[:, :-1]
1448
+ accel_mask = valid.clone()
1449
+ accel_mask[:, :2] = False
1450
+ if valid.shape[1] > 2:
1451
+ accel_mask[:, 2:] = valid[:, 2:] & valid[:, 1:-1] & valid[:, :-2]
1452
+ return {
1453
+ "velocity": velocity.masked_fill(~delta_mask, 0.0),
1454
+ "acceleration": acceleration.masked_fill(~accel_mask, 0.0),
1455
+ "reversal": reversal.masked_fill(~accel_mask, 0.0),
1456
+ "recurrence": recurrence.masked_fill(~accel_mask, 0.0),
1457
+ "oscillation": osc.masked_fill(~accel_mask, 0.0),
1458
+ "raw_mask": raw_mask,
1459
+ "delta_mask": delta_mask,
1460
+ "accel_mask": accel_mask,
1461
+ }
1462
+
1463
+ def _frame_metadata(
1464
+ self,
1465
+ history: TrajectoryHistory,
1466
+ state: LatentDeliberationState,
1467
+ dtype: torch.dtype,
1468
+ geometry: dict[str, torch.Tensor],
1469
+ ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
1470
+ batch, history_length, canvas, _dim = history.hidden.shape
1471
+ log_vocab = self.log_vocab
1472
+ confidence = _renorm_confidence(history.confidence.to(dtype=dtype))
1473
+ entropy = (history.entropy.to(dtype=dtype) / log_vocab).clamp(0.0, 1.0)
1474
+ if history_length == 1:
1475
+ recency = torch.ones(
1476
+ batch, history_length, canvas, device=history.hidden.device, dtype=dtype
1477
+ )
1478
+ else:
1479
+ recency = torch.linspace(
1480
+ 0.0, 1.0, history_length, device=history.hidden.device, dtype=dtype
1481
+ )
1482
+ recency = recency.view(1, history_length, 1).expand(batch, history_length, canvas)
1483
+ retirement = torch.zeros(
1484
+ batch, history_length, canvas, device=history.hidden.device, dtype=dtype
1485
+ )
1486
+ oldest = min(_RETIREMENT_FRAMES, history_length)
1487
+ if oldest:
1488
+ scale = torch.linspace(
1489
+ 1.0, 1.0 / oldest, oldest, device=history.hidden.device, dtype=dtype
1490
+ )
1491
+ retirement[:, :oldest] = scale.view(1, oldest, 1)
1492
+ retirement = retirement * history.valid.to(dtype=dtype)
1493
+ changed = history.token_changed.to(dtype=dtype)
1494
+ delta_c = torch.zeros_like(confidence)
1495
+ delta_e = torch.zeros_like(entropy)
1496
+ if history_length > 1:
1497
+ delta_c[:, 1:] = (confidence[:, 1:] - confidence[:, :-1]).clamp(-1.0, 1.0)
1498
+ delta_e[:, 1:] = ((history.entropy[:, 1:] - history.entropy[:, :-1]) / log_vocab).clamp(
1499
+ -1.0, 1.0
1500
+ ).to(dtype=dtype)
1501
+ age = (
1502
+ state.age.to(dtype=dtype).clamp_max(_AGE_MAX).log1p()
1503
+ / math.log1p(_AGE_MAX)
1504
+ )
1505
+ age_frames = torch.zeros_like(confidence)
1506
+ age_frames[:, -1] = age
1507
+ key_meta = torch.stack(
1508
+ (
1509
+ confidence, entropy, age_frames, changed, delta_c, delta_e,
1510
+ recency, retirement,
1511
+ geometry["velocity"].to(dtype=dtype),
1512
+ geometry["acceleration"].to(dtype=dtype),
1513
+ geometry["reversal"].to(dtype=dtype),
1514
+ geometry["recurrence"].to(dtype=dtype),
1515
+ geometry["oscillation"].to(dtype=dtype),
1516
+ ),
1517
+ dim=-1,
1518
+ )
1519
+ query_scalars = torch.stack(
1520
+ (
1521
+ _renorm_confidence(state.confidence.to(dtype=dtype)),
1522
+ (state.entropy.to(dtype=dtype) / log_vocab).clamp(0.0, 1.0),
1523
+ age,
1524
+ state.token_changed.to(dtype=dtype),
1525
+ state.confidence_delta.to(dtype=dtype).clamp(-1.0, 1.0),
1526
+ (state.entropy_delta.to(dtype=dtype) / log_vocab).clamp(-1.0, 1.0),
1527
+ ),
1528
+ dim=-1,
1529
+ )
1530
+ fourier = _canvas_fourier(canvas, history.hidden.device, dtype).unsqueeze(0).expand(
1531
+ batch, -1, -1
1532
+ )
1533
+ query_meta = torch.cat((fourier, query_scalars), dim=-1)
1534
+ ponder = (
1535
+ state.ponder_steps.to(dtype=dtype).clamp_max(_PONDER_MAX).log1p()
1536
+ / math.log1p(_PONDER_MAX)
1537
+ )
1538
+ stagnation = (
1539
+ state.stagnation_steps.to(dtype=dtype).clamp_max(_STAGNATION_MAX).log1p()
1540
+ / math.log1p(_STAGNATION_MAX)
1541
+ )
1542
+ row_meta = torch.stack((ponder, stagnation), dim=-1)
1543
+ return key_meta, query_meta, row_meta, history.valid
1544
+
1545
+ def _history_views(
1546
+ self,
1547
+ history: TrajectoryHistory,
1548
+ key_meta: torch.Tensor,
1549
+ geometry: dict[str, torch.Tensor],
1550
+ dtype: torch.dtype,
1551
+ ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
1552
+ mapped = self.history_in(history.hidden.to(dtype=dtype))
1553
+ projected = self.history_projector(mapped)
1554
+ return self._views_from_projected(
1555
+ projected, history, key_meta, geometry, dtype,
1556
+ preformat_heads=True,
1557
+ )
1558
+
1559
+ def _views_from_projected(
1560
+ self,
1561
+ projected: torch.Tensor,
1562
+ history: TrajectoryHistory,
1563
+ key_meta: torch.Tensor,
1564
+ geometry: dict[str, torch.Tensor],
1565
+ dtype: torch.dtype,
1566
+ *,
1567
+ preformat_heads: bool = False,
1568
+ ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
1569
+ valid = history.valid
1570
+ raw_mask = geometry["raw_mask"]
1571
+ delta_mask = geometry["delta_mask"]
1572
+ accel_mask = geometry["accel_mask"]
1573
+ latest_index = (
1574
+ valid.to(torch.int64) * (
1575
+ torch.arange(valid.shape[1], device=valid.device).view(1, -1, 1) + 1
1576
+ )
1577
+ ).amax(dim=1) - 1
1578
+ has_latest = latest_index.ge(0)
1579
+ latest_index = latest_index.clamp_min(0)
1580
+ gather = latest_index.view(projected.shape[0], 1, projected.shape[2], 1).expand(
1581
+ -1, 1, -1, projected.shape[-1]
1582
+ )
1583
+ z_latest = projected.gather(1, gather).squeeze(1)
1584
+ residual = projected - z_latest.unsqueeze(1)
1585
+ residual_mask = raw_mask & has_latest.unsqueeze(1)
1586
+ if projected.shape[1] > 1:
1587
+ delta = torch.cat(
1588
+ (torch.zeros_like(projected[:, :1]), projected[:, 1:] - projected[:, :-1]),
1589
+ dim=1,
1590
+ )
1591
+ else:
1592
+ delta = torch.zeros_like(projected)
1593
+ accel = torch.zeros_like(projected)
1594
+ if projected.shape[1] > 2:
1595
+ accel[:, 2:] = projected[:, 2:] - 2.0 * projected[:, 1:-1] + projected[:, :-2]
1596
+ views = torch.stack((projected, delta, accel, residual), dim=2)
1597
+ view_mask = torch.stack((raw_mask, delta_mask, accel_mask, residual_mask), dim=2)
1598
+ film = self.film_out(F.silu(self.film_in(key_meta.to(dtype=dtype))))
1599
+ scale, shift = film.chunk(2, dim=-1)
1600
+ numeric_mask = view_mask.unsqueeze(-1).to(dtype=views.dtype)
1601
+ # The unmasked `views` tensor is dead here. Mask it in-place and use
1602
+ # it as the key storage instead of allocating a second full copy.
1603
+ # Applying the same mask again after the FiLM shift preserves invalid
1604
+ # entries as exact zeros and leaves every valid entry unchanged.
1605
+ views.mul_(numeric_mask)
1606
+ values = views * (1.0 + scale.unsqueeze(2))
1607
+ values.add_(shift.unsqueeze(2)).mul_(numeric_mask)
1608
+ keys = views
1609
+ batch, history_length, views_n, canvas, dim = keys.shape
1610
+ if preformat_heads:
1611
+ heads = self.blocks[0].history_attention.num_heads
1612
+ head_dim = dim // heads
1613
+ keys = keys.view(
1614
+ batch, history_length, views_n, canvas, heads, head_dim
1615
+ ).permute(0, 4, 3, 1, 2, 5).reshape(
1616
+ batch, heads, canvas, history_length * views_n, head_dim
1617
+ )
1618
+ values = values.view(
1619
+ batch, history_length, views_n, canvas, heads, head_dim
1620
+ ).permute(0, 4, 3, 1, 2, 5).reshape(
1621
+ batch, heads, canvas, history_length * views_n, head_dim
1622
+ )
1623
+ else:
1624
+ keys = keys.permute(0, 3, 1, 2, 4).reshape(
1625
+ batch, canvas, history_length * views_n, dim
1626
+ )
1627
+ values = values.permute(0, 3, 1, 2, 4).reshape(
1628
+ batch, canvas, history_length * views_n, dim
1629
+ )
1630
+ key_mask = view_mask.permute(0, 3, 1, 2).reshape(batch, canvas, history_length * views_n)
1631
+ return keys, values, key_mask, projected
1632
+
1633
+ def _attention_bias(
1634
+ self,
1635
+ key_meta: torch.Tensor,
1636
+ query_meta: torch.Tensor,
1637
+ row_meta: torch.Tensor,
1638
+ view_mask: torch.Tensor,
1639
+ num_heads: int,
1640
+ dtype: torch.dtype,
1641
+ ) -> torch.Tensor:
1642
+ batch, history_length, canvas, _meta = key_meta.shape
1643
+ views = _HISTORY_VIEWS
1644
+ query = query_meta[:, None, :, :].expand(-1, history_length, -1, -1)
1645
+ row = row_meta[:, None, None, :].expand(-1, history_length, canvas, -1)
1646
+ packed = torch.cat((key_meta.to(dtype=dtype), query, row), dim=-1)
1647
+ bias = self.bias_out(F.silu(self.bias_in(packed)))
1648
+ bias = bias.permute(0, 3, 2, 1).unsqueeze(-1).expand(-1, -1, -1, -1, views)
1649
+ return bias.reshape(batch, num_heads, canvas, history_length * views).to(dtype=dtype)
1650
+
1651
+ def forward(
1652
+ self,
1653
+ *,
1654
+ token_embeddings: torch.Tensor,
1655
+ confidence: torch.Tensor,
1656
+ entropy: torch.Tensor,
1657
+ state: LatentDeliberationState,
1658
+ history: TrajectoryHistory,
1659
+ tape: TrajectoryTape | None = None,
1660
+ ) -> LatentProcessorOutput:
1661
+ if token_embeddings.ndim != 3:
1662
+ raise ValueError("`token_embeddings` must have shape [batch, canvas, hidden].")
1663
+ batch_size, canvas_length, hidden_size = token_embeddings.shape
1664
+ if hidden_size != self.hidden_size:
1665
+ raise ValueError("Unexpected hidden size for latent deliberation.")
1666
+ if state.memory_slots.shape != (batch_size, self.memory_slots, self.latent_dim):
1667
+ raise ValueError("State memory slots do not match this module.")
1668
+ if history.hidden.shape[:3] != (batch_size, self.history_length, canvas_length):
1669
+ raise ValueError("Trajectory history does not match the current canvas.")
1670
+ if state.age.dtype is not torch.int32:
1671
+ raise TypeError("Latent deliberation ages must use int32.")
1672
+ if tape is None:
1673
+ tape = TrajectoryTape.empty(
1674
+ batch_size=batch_size,
1675
+ tape_length=self.history_length,
1676
+ num_probes=self.tape_probes,
1677
+ probe_dim=self.tape_pool.rank,
1678
+ device=token_embeddings.device,
1679
+ dtype=token_embeddings.dtype,
1680
+ )
1681
+ if tape.probes.shape[:2] != (batch_size, self.history_length):
1682
+ raise ValueError("Trajectory tape does not match the current batch.")
1683
+
1684
+ dtype = token_embeddings.dtype
1685
+ query = self.query_in(token_embeddings)
1686
+ geometry = self._geometry(history)
1687
+ key_meta, query_meta, row_meta, _valid = self._frame_metadata(
1688
+ history, state, dtype, geometry
1689
+ )
1690
+ keys, values, key_mask, projected = self._history_views(
1691
+ history, key_meta, geometry, dtype
1692
+ )
1693
+ num_heads = self.blocks[0].history_attention.num_heads
1694
+ attn_bias = self._attention_bias(
1695
+ key_meta, query_meta, row_meta, key_mask, num_heads, dtype
1696
+ )
1697
+ tape_keys, tape_values, tape_mask = self._tape_keys(tape, dtype)
1698
+ # Working state accumulates from zero. Empty history and zero memory
1699
+ # therefore produce a zero self-conditioning residual.
1700
+ canvas_state = torch.zeros_like(query)
1701
+ slot_identity = self.scaled_memory_slot_identity(
1702
+ batch_size=batch_size, device=state.memory_slots.device, dtype=query.dtype
1703
+ )
1704
+ memory_keys, memory_values = self.persistent_kv(state.memory_slots, slot_identity)
1705
+ for block in self.blocks:
1706
+ canvas_state = block(
1707
+ canvas_state, keys, values, attn_bias, key_mask,
1708
+ memory_keys, memory_values, query,
1709
+ tape_keys, tape_values, tape_mask,
1710
+ )
1711
+ context = self.project_context(canvas_state)
1712
+ next_state = LatentDeliberationState(
1713
+ memory_slots=state.memory_slots,
1714
+ confidence=confidence.to(dtype=torch.float32),
1715
+ entropy=entropy.to(dtype=torch.float32),
1716
+ age=state.age,
1717
+ token_changed=state.token_changed,
1718
+ confidence_delta=state.confidence_delta,
1719
+ entropy_delta=state.entropy_delta,
1720
+ ponder_steps=state.ponder_steps,
1721
+ stagnation_steps=state.stagnation_steps,
1722
+ )
1723
+ return LatentProcessorOutput(
1724
+ context=context,
1725
+ working_state=context,
1726
+ state=next_state,
1727
+ history_projected=projected,
1728
+ )
1729
+
1730
+ def _slice_commit_canvas(
1731
+ self,
1732
+ *,
1733
+ working_state: torch.Tensor,
1734
+ history: TrajectoryHistory,
1735
+ history_projected: torch.Tensor,
1736
+ heavy_hidden: torch.Tensor,
1737
+ max_commit: int,
1738
+ ) -> tuple[torch.Tensor, TrajectoryHistory, torch.Tensor, torch.Tensor]:
1739
+ return (
1740
+ working_state[:, :max_commit],
1741
+ TrajectoryHistory(
1742
+ hidden=history.hidden[:, :, :max_commit],
1743
+ confidence=history.confidence[:, :, :max_commit],
1744
+ entropy=history.entropy[:, :, :max_commit],
1745
+ token_changed=history.token_changed[:, :, :max_commit],
1746
+ valid=history.valid[:, :, :max_commit],
1747
+ ),
1748
+ history_projected[:, :, :max_commit],
1749
+ heavy_hidden[:, :max_commit],
1750
+ )
1751
+
1752
+ def _experience_encoder_step(
1753
+ self,
1754
+ working_state: torch.Tensor,
1755
+ history_keys: torch.Tensor,
1756
+ history_values: torch.Tensor,
1757
+ history_mask: torch.Tensor,
1758
+ z_final: torch.Tensor,
1759
+ commit_reason: torch.Tensor,
1760
+ ) -> torch.Tensor:
1761
+ return self.experience_encoder(
1762
+ working_state=working_state,
1763
+ history_keys=history_keys,
1764
+ history_values=history_values,
1765
+ history_mask=history_mask,
1766
+ z_final=z_final,
1767
+ commit_reason=commit_reason,
1768
+ )
1769
+
1770
+ def _encode_experience_roles(
1771
+ self,
1772
+ *,
1773
+ working_state: torch.Tensor,
1774
+ history_keys: torch.Tensor,
1775
+ history_values: torch.Tensor,
1776
+ history_mask: torch.Tensor,
1777
+ z_final: torch.Tensor,
1778
+ commit_reason: torch.Tensor,
1779
+ ) -> torch.Tensor:
1780
+ canvas = int(working_state.shape[1])
1781
+ stripe = min(_EXPERIENCE_CANVAS_STRIPE, canvas)
1782
+ parts: list[torch.Tensor] = []
1783
+ for start in range(0, canvas, stripe):
1784
+ stop = min(start + stripe, canvas)
1785
+ parts.append(
1786
+ self._experience_encoder_step(
1787
+ working_state[:, start:stop],
1788
+ history_keys[:, start:stop],
1789
+ history_values[:, start:stop],
1790
+ history_mask[:, start:stop],
1791
+ z_final[:, start:stop],
1792
+ commit_reason,
1793
+ )
1794
+ )
1795
+ return parts[0] if len(parts) == 1 else torch.cat(parts, dim=1)
1796
+
1797
+ def _pack_experience(
1798
+ self,
1799
+ *,
1800
+ working_state: torch.Tensor,
1801
+ history: TrajectoryHistory,
1802
+ history_projected: torch.Tensor,
1803
+ heavy_hidden: torch.Tensor,
1804
+ commit_lengths: torch.Tensor,
1805
+ prefix_lengths: torch.Tensor,
1806
+ commit_reason: torch.Tensor,
1807
+ ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
1808
+ batch, canvas, _hidden = working_state.shape
1809
+ max_commit = min(int(commit_lengths.max().clamp_min(0)), canvas)
1810
+ if max_commit <= 0:
1811
+ empty = working_state.new_zeros(batch, 0, self.packet_dim)
1812
+ return empty, empty.new_zeros(batch, 0, dtype=torch.bool), empty.new_zeros(batch, 0)
1813
+ working_state, history, history_projected, heavy_hidden = self._slice_commit_canvas(
1814
+ working_state=working_state,
1815
+ history=history,
1816
+ history_projected=history_projected,
1817
+ heavy_hidden=heavy_hidden,
1818
+ max_commit=max_commit,
1819
+ )
1820
+ dtype = working_state.dtype
1821
+ geometry = self._geometry(history)
1822
+ key_meta, _query_meta, _row_meta, _valid = self._frame_metadata(
1823
+ history,
1824
+ LatentDeliberationState.empty(
1825
+ batch_size=batch,
1826
+ canvas_length=max_commit,
1827
+ latent_dim=self.latent_dim,
1828
+ memory_slots=self.memory_slots,
1829
+ device=working_state.device,
1830
+ dtype=dtype,
1831
+ ),
1832
+ dtype,
1833
+ geometry,
1834
+ )
1835
+ keys, values, key_mask, _projected = self._views_from_projected(
1836
+ history_projected.to(dtype=dtype), history, key_meta, geometry, dtype
1837
+ )
1838
+ # Heavy is a TBPTT observation here, same as history frames: CE already
1839
+ # backpropagated through this decoder stack. The writer stays in the
1840
+ # temporal graph via `working_state` and the committed memory output.
1841
+ z_final = self.history_projector(
1842
+ self.history_in(heavy_hidden.detach().to(dtype=dtype))
1843
+ )
1844
+ roles = self._encode_experience_roles(
1845
+ working_state=working_state,
1846
+ history_keys=keys,
1847
+ history_values=values,
1848
+ history_mask=key_mask,
1849
+ z_final=z_final,
1850
+ commit_reason=commit_reason,
1851
+ )
1852
+ tokens = roles.reshape(batch, max_commit * _EXPERIENCE_ROLES, -1)
1853
+ token_valid = torch.arange(max_commit, device=working_state.device)[None, :] < commit_lengths[:, None]
1854
+ valid = token_valid.unsqueeze(-1).expand(-1, -1, _EXPERIENCE_ROLES).reshape(
1855
+ batch, max_commit * _EXPERIENCE_ROLES
1856
+ )
1857
+ abs_pos = prefix_lengths[:, None] + torch.arange(
1858
+ max_commit, device=working_state.device
1859
+ )[None, :]
1860
+ pos = abs_pos.unsqueeze(-1).expand(-1, -1, _EXPERIENCE_ROLES).reshape(
1861
+ batch, max_commit * _EXPERIENCE_ROLES
1862
+ )
1863
+ mixed = self.commit_sequence(tokens, pos, valid)
1864
+ return mixed, valid, pos
1865
+
1866
+ def commit_write(
1867
+ self,
1868
+ *,
1869
+ memory: torch.Tensor,
1870
+ working_state: torch.Tensor,
1871
+ history: TrajectoryHistory,
1872
+ history_projected: torch.Tensor,
1873
+ heavy_hidden: torch.Tensor,
1874
+ commit_lengths: torch.Tensor,
1875
+ prefix_lengths: torch.Tensor | None = None,
1876
+ commit_reason: torch.Tensor | None = None,
1877
+ **kwargs: Any,
1878
+ ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
1879
+ batch = memory.shape[0]
1880
+ canvas = working_state.shape[1]
1881
+ lengths = commit_lengths.to(device=memory.device, dtype=torch.long)
1882
+ if prefix_lengths is None:
1883
+ prefixes = torch.zeros(batch, device=memory.device, dtype=torch.long)
1884
+ else:
1885
+ prefixes = prefix_lengths.to(device=memory.device, dtype=torch.long)
1886
+ if commit_reason is None:
1887
+ reasons = torch.full(
1888
+ (batch,), COMMIT_REASON_NORMAL, device=memory.device, dtype=torch.long
1889
+ )
1890
+ else:
1891
+ reasons = commit_reason.to(device=memory.device, dtype=torch.long)
1892
+ if bool((lengths <= 0).all()):
1893
+ zero = memory.new_zeros(batch, self.memory_slots, 1)
1894
+ return memory, {
1895
+ "gate_mean": zero.mean(),
1896
+ "gate_max": zero.amax(),
1897
+ "gate_gt_01": zero.new_zeros(()),
1898
+ "gate_gt_05": zero.new_zeros(()),
1899
+ "delta_norm_mean": memory.new_zeros(()),
1900
+ }
1901
+ experience, experience_mask, _pos = self._pack_experience(
1902
+ working_state=working_state,
1903
+ history=history,
1904
+ history_projected=history_projected,
1905
+ heavy_hidden=heavy_hidden,
1906
+ commit_lengths=lengths,
1907
+ prefix_lengths=prefixes,
1908
+ commit_reason=reasons,
1909
+ )
1910
+ max_commit = int(lengths.max())
1911
+ role_stop = max_commit * _EXPERIENCE_ROLES
1912
+ chunk_tokens = experience[:, :role_stop]
1913
+ chunk_mask = experience_mask[:, :role_stop]
1914
+ if chunk_tokens.shape[1] > 0:
1915
+ written, last_gate, last_delta = self.commit_writer(
1916
+ memory, chunk_tokens, chunk_mask
1917
+ )
1918
+ else:
1919
+ written = memory
1920
+ last_gate = memory.new_zeros(batch, self.memory_slots, 1)
1921
+ last_delta = memory.new_zeros(memory.shape)
1922
+ gate = last_gate.detach()
1923
+ delta_norm = last_delta.detach().float().norm(dim=-1)
1924
+ diagnostics = {
1925
+ "gate_mean": gate.mean(),
1926
+ "gate_max": gate.amax(),
1927
+ "gate_gt_01": gate.gt(0.1).float().sum(),
1928
+ "gate_gt_05": gate.gt(0.5).float().sum(),
1929
+ "delta_norm_mean": delta_norm.mean(),
1930
+ }
1931
+ unchanged = lengths.le(0).view(batch, 1, 1)
1932
+ written = torch.where(unchanged, memory, written)
1933
+ return written, diagnostics
1934
+
1935
+
1936
+ def slice_trajectory_history(
1937
+ history: TrajectoryHistory, rows: slice | torch.Tensor
1938
+ ) -> TrajectoryHistory:
1939
+ return TrajectoryHistory(
1940
+ hidden=history.hidden[rows],
1941
+ confidence=history.confidence[rows],
1942
+ entropy=history.entropy[rows],
1943
+ token_changed=history.token_changed[rows],
1944
+ valid=history.valid[rows],
1945
+ )
1946
+
1947
+
1948
+ def cat_trajectory_history(
1949
+ histories: Sequence[TrajectoryHistory],
1950
+ ) -> TrajectoryHistory:
1951
+ return TrajectoryHistory(
1952
+ hidden=torch.cat([history.hidden for history in histories], dim=0),
1953
+ confidence=torch.cat([history.confidence for history in histories], dim=0),
1954
+ entropy=torch.cat([history.entropy for history in histories], dim=0),
1955
+ token_changed=torch.cat([history.token_changed for history in histories], dim=0),
1956
+ valid=torch.cat([history.valid for history in histories], dim=0),
1957
+ )
1958
+
1959
+
1960
+ def choose_trajectory_history(
1961
+ previous: TrajectoryHistory,
1962
+ updated: TrajectoryHistory,
1963
+ update_mask: torch.Tensor,
1964
+ ) -> TrajectoryHistory:
1965
+ def choose(old: torch.Tensor, new: torch.Tensor) -> torch.Tensor:
1966
+ mask = update_mask.view(update_mask.shape[0], *([1] * (old.ndim - 1)))
1967
+ return torch.where(mask, new, old)
1968
+
1969
+ return TrajectoryHistory(
1970
+ hidden=choose(previous.hidden, updated.hidden),
1971
+ confidence=choose(previous.confidence, updated.confidence),
1972
+ entropy=choose(previous.entropy, updated.entropy),
1973
+ token_changed=choose(previous.token_changed, updated.token_changed),
1974
+ valid=choose(previous.valid, updated.valid),
1975
+ )
1976
+
1977
+
1978
+ def slice_trajectory_tape(
1979
+ tape: TrajectoryTape, rows: slice | torch.Tensor
1980
+ ) -> TrajectoryTape:
1981
+ return TrajectoryTape(probes=tape.probes[rows], valid=tape.valid[rows])
1982
+
1983
+
1984
+ def cat_trajectory_tape(tapes: Sequence[TrajectoryTape]) -> TrajectoryTape:
1985
+ return TrajectoryTape(
1986
+ probes=torch.cat([tape.probes for tape in tapes], dim=0),
1987
+ valid=torch.cat([tape.valid for tape in tapes], dim=0),
1988
+ )
1989
+
1990
+
1991
+ def choose_trajectory_tape(
1992
+ previous: TrajectoryTape,
1993
+ updated: TrajectoryTape,
1994
+ update_mask: torch.Tensor,
1995
+ ) -> TrajectoryTape:
1996
+ def choose(old: torch.Tensor, new: torch.Tensor) -> torch.Tensor:
1997
+ mask = update_mask.view(update_mask.shape[0], *([1] * (old.ndim - 1)))
1998
+ return torch.where(mask, new, old)
1999
+
2000
+ return TrajectoryTape(
2001
+ probes=choose(previous.probes, updated.probes),
2002
+ valid=choose(previous.valid, updated.valid),
2003
+ )
2004
+
2005
+
2006
+ def empty_trajectory_tape(
2007
+ *,
2008
+ batch_size: int,
2009
+ config: object,
2010
+ device: torch.device,
2011
+ dtype: torch.dtype,
2012
+ ) -> TrajectoryTape:
2013
+ rank = int(getattr(config, "latent_history_kv_rank"))
2014
+ return TrajectoryTape.empty(
2015
+ batch_size=batch_size,
2016
+ tape_length=int(getattr(config, "latent_history_length")),
2017
+ num_probes=int(getattr(config, "latent_tape_probes", 16)),
2018
+ probe_dim=rank,
2019
+ device=device,
2020
+ dtype=dtype,
2021
+ )
2022
+
2023
+
2024
+ def slice_latent_state(
2025
+ state: LatentDeliberationState, rows: slice | torch.Tensor
2026
+ ) -> LatentDeliberationState:
2027
+ return LatentDeliberationState(
2028
+ memory_slots=state.memory_slots[rows],
2029
+ confidence=state.confidence[rows],
2030
+ entropy=state.entropy[rows],
2031
+ age=state.age[rows],
2032
+ token_changed=state.token_changed[rows],
2033
+ confidence_delta=state.confidence_delta[rows],
2034
+ entropy_delta=state.entropy_delta[rows],
2035
+ ponder_steps=state.ponder_steps[rows],
2036
+ stagnation_steps=state.stagnation_steps[rows],
2037
+ )
2038
+
2039
+
2040
+ def cat_latent_states(states: Sequence[LatentDeliberationState]) -> LatentDeliberationState:
2041
+ def cat(name: str) -> torch.Tensor:
2042
+ return torch.cat([getattr(state, name) for state in states], dim=0)
2043
+
2044
+ return LatentDeliberationState(
2045
+ memory_slots=cat("memory_slots"),
2046
+ confidence=cat("confidence"),
2047
+ entropy=cat("entropy"),
2048
+ age=cat("age"),
2049
+ token_changed=cat("token_changed"),
2050
+ confidence_delta=cat("confidence_delta"),
2051
+ entropy_delta=cat("entropy_delta"),
2052
+ ponder_steps=cat("ponder_steps"),
2053
+ stagnation_steps=cat("stagnation_steps"),
2054
+ )
2055
+
2056
+
2057
+ def infer_commit_reason(
2058
+ commit_lengths: torch.Tensor,
2059
+ *,
2060
+ jump_rows: torch.Tensor | None = None,
2061
+ commit_token_ids: torch.Tensor | None = None,
2062
+ terminal_token_ids: Sequence[int] = (),
2063
+ training_random: bool = False,
2064
+ ) -> torch.Tensor:
2065
+ """Return per-row commit-reason codes. No hard skip; writer sees the label."""
2066
+
2067
+ reasons = torch.full(
2068
+ commit_lengths.shape,
2069
+ COMMIT_REASON_NONE,
2070
+ device=commit_lengths.device,
2071
+ dtype=torch.long,
2072
+ )
2073
+ committed = commit_lengths.gt(0)
2074
+ default = (
2075
+ COMMIT_REASON_TRAINING_RANDOM if training_random else COMMIT_REASON_NORMAL
2076
+ )
2077
+ reasons = torch.where(committed, torch.full_like(reasons, default), reasons)
2078
+ if jump_rows is not None:
2079
+ reasons = torch.where(
2080
+ committed & jump_rows.to(dtype=torch.bool),
2081
+ torch.full_like(reasons, COMMIT_REASON_FORCED_JUMP),
2082
+ reasons,
2083
+ )
2084
+ if commit_token_ids is not None and terminal_token_ids:
2085
+ positions = torch.arange(
2086
+ commit_token_ids.shape[1], device=commit_token_ids.device
2087
+ )[None, :]
2088
+ selected = positions.lt(commit_lengths[:, None])
2089
+ terminal = torch.zeros_like(committed)
2090
+ for token_id in terminal_token_ids:
2091
+ terminal |= (commit_token_ids.eq(int(token_id)) & selected).any(dim=-1)
2092
+ reasons = torch.where(
2093
+ committed & terminal,
2094
+ torch.full_like(reasons, COMMIT_REASON_TERMINAL),
2095
+ reasons,
2096
+ )
2097
+ return reasons
2098
+
2099
+
2100
+ def memory_bus_parameter_names(module: nn.Module) -> list[str]:
2101
+ return [
2102
+ name for name, _parameter in module.named_parameters()
2103
+ if "working_memory_bus." in name or "persistent_memory_bus." in name
2104
+ or "memory_bus." in name
2105
+ ]
2106
+
2107
+
2108
+ __all__ = [
2109
+ "COMMIT_REASON_FALLBACK",
2110
+ "COMMIT_REASON_FORCED_JUMP",
2111
+ "COMMIT_REASON_NONE",
2112
+ "COMMIT_REASON_NORMAL",
2113
+ "COMMIT_REASON_TERMINAL",
2114
+ "COMMIT_REASON_TRAINING_RANDOM",
2115
+ "DecoderMemoryBus",
2116
+ "LatentDeliberationState",
2117
+ "LatentDeliberationTransformer",
2118
+ "LatentProcessorOutput",
2119
+ "TrajectoryHistory",
2120
+ "TrajectoryTape",
2121
+ "advance_trajectory_clocks",
2122
+ "cat_latent_states",
2123
+ "cat_trajectory_history",
2124
+ "cat_trajectory_tape",
2125
+ "choose_trajectory_history",
2126
+ "choose_trajectory_tape",
2127
+ "empty_trajectory_tape",
2128
+ "infer_commit_reason",
2129
+ "memory_bus_parameter_names",
2130
+ "should_force_trajectory_jump",
2131
+ "slice_latent_state",
2132
+ "slice_trajectory_history",
2133
+ "slice_trajectory_tape",
2134
+ ]
model-00001-of-00011.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:437da1276f5e908bd5cd0d771377b31d0089188af394a9670b6e0e1d86b80a37
3
+ size 4718357052
model-00002-of-00011.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:7790953e9060f3e95d125ca89bdfe0fd888230b88afc6e35553a7e25c39104bf
3
+ size 4913414758
model-00003-of-00011.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:5a7a10f3ff4cab7e3ac7b4fe49cc8c5ef05703b7294399948fc046ef7cd5e4b6
3
+ size 4884578046
model-00004-of-00011.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d411b2543b3e3dae71430cb3a6339292f320ac2d7e116e6e5a529dbd0d094da1
3
+ size 4913414782
model-00005-of-00011.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:0a2acd7790f9def59c36ad66c58bd9a8c1b6889c5a024f1e57f7d1d8b34d5e6a
3
+ size 4884578022
model-00006-of-00011.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e5909f8cd062f1150d7adbe8e246a6de6d4eac9aeadec0204f5a6847aaeb820e
3
+ size 4884578046
model-00007-of-00011.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:08c0d652218e400128d4320b5825e38cb1c09abd79ae0d795c66a5ab774f99ae
3
+ size 4913414782