ydy9038074 commited on
Commit
b88f761
·
verified ·
1 Parent(s): 4316924

Publish Modilify Mk2 Preview step 1250 schema25

Browse files
Files changed (50) hide show
  1. .gitattributes +0 -1
  2. NOTICE.md +7 -6
  3. README.md +30 -248
  4. chat_template.jinja +4 -4
  5. commit_policy.py +40 -30
  6. config.json +38 -81
  7. configuration_modilify_mk2.py +73 -297
  8. continuous_batching.py +64 -154
  9. gdn2_memory.py +112 -0
  10. gdn2_trajectory.py +86 -0
  11. generation_config.json +1 -15
  12. generation_modilify_mk2.py +118 -461
  13. latent_deliberation.py +247 -1757
  14. lora.py +96 -0
  15. model-00001-of-00032.safetensors +3 -0
  16. model-00002-of-00032.safetensors +3 -0
  17. model-00003-of-00032.safetensors +3 -0
  18. model-00004-of-00032.safetensors +3 -0
  19. model-00005-of-00032.safetensors +3 -0
  20. model-00006-of-00032.safetensors +3 -0
  21. model-00007-of-00032.safetensors +3 -0
  22. model-00008-of-00032.safetensors +3 -0
  23. model-00009-of-00032.safetensors +3 -0
  24. model-00010-of-00032.safetensors +3 -0
  25. model-00011-of-00032.safetensors +3 -0
  26. model-00012-of-00032.safetensors +3 -0
  27. model-00013-of-00032.safetensors +3 -0
  28. model-00014-of-00032.safetensors +3 -0
  29. model-00015-of-00032.safetensors +3 -0
  30. model-00016-of-00032.safetensors +3 -0
  31. model-00017-of-00032.safetensors +3 -0
  32. model-00018-of-00032.safetensors +3 -0
  33. model-00019-of-00032.safetensors +3 -0
  34. model-00020-of-00032.safetensors +3 -0
  35. model-00021-of-00032.safetensors +3 -0
  36. model-00022-of-00032.safetensors +3 -0
  37. model-00023-of-00032.safetensors +3 -0
  38. model-00024-of-00032.safetensors +3 -0
  39. model-00025-of-00032.safetensors +3 -0
  40. model-00026-of-00032.safetensors +3 -0
  41. model-00027-of-00032.safetensors +3 -0
  42. model-00028-of-00032.safetensors +3 -0
  43. model-00029-of-00032.safetensors +3 -0
  44. model-00030-of-00032.safetensors +3 -0
  45. model-00031-of-00032.safetensors +3 -0
  46. model-00032-of-00032.safetensors +3 -0
  47. model.safetensors.index.json +0 -0
  48. modeling_modilify_mk2.py +157 -254
  49. mps_ops.py +99 -375
  50. vocab_ops.py +136 -438
.gitattributes CHANGED
@@ -1,3 +1,2 @@
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
 
1
  *.safetensors filter=lfs diff=lfs merge=lfs -text
2
  tokenizer.json filter=lfs diff=lfs merge=lfs -text
 
NOTICE.md CHANGED
@@ -3,12 +3,13 @@
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
 
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. This PyTorch distribution retains
7
+ the text encoder/decoder backbone, tokenizer, and chat template. It migrates
8
+ the step 1250 schema25 weights from the native MLX release, preserving unfused
9
+ LoRA adapters and dual-timescale GDN2 trajectory memory. The runtime follows
10
+ the Modilify-Mk reference implementation and its confidence-and-entropy prefix
11
+ commit policy. It contains no vision tower, vision projection, or image/video
12
+ inference path.
13
 
14
  The Modilify Open Model License 1.0 applies to Modilify's distribution and
15
  original contributions. It does not erase, narrow, or replace rights and notices
README.md CHANGED
@@ -3,278 +3,60 @@ 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 · 26B-A5B
18
 
19
- **Refine a whole canvas. Think in latent space. Carry memory forward.**
 
 
 
20
 
21
- Modilify Mk2 is a **26B-A5B multimodal block-diffusion model** that brings parallel token refinement, latent deliberation, and persistent trajectory memory into one generation loop. A rolling 256-token canvas gives the model room to revise upcoming text together. A dedicated latent Transformer turns the history of those revisions into context for the next step. As output advances, a learned memory writer carries information across canvas boundaries.
 
22
 
23
- The central idea is simple: **give each answer a workspace, a memory, and a variable compute budget.** Mk2 can refine its internal state over multiple passes before committing text, and release multiple tokens together when its confidence-and-entropy policy permits. Deliberation happens in continuous hidden states, without requiring every internal update to become a visible reasoning token.
24
 
25
- Built on DiffusionGemma, Mk2 combines sparse expert routing with a dual-timescale latent architecture. Text, images, and sampled video frames feed the same generation path.
26
-
27
- ## What makes Mk2 different
28
-
29
- | Architecture choice | What it enables |
30
- | --- | --- |
31
- | **Parallel block diffusion** | Revise a 256-token canvas jointly and commit a variable-length prefix, allowing multiple output tokens per denoising pass. |
32
- | **Trajectory-aware latent deliberation** | Condition the next revision on how hidden states have evolved, including their changes, acceleration, and residuals relative to the latest state. |
33
- | **Memory at two timescales** | Rebuild working state every pass while retaining persistent slots across rolling-window shifts within a generation. |
34
- | **Memory inside the decoder** | Feed working and persistent memory directly into full-attention decoder layers so both can influence token refinement. |
35
- | **Adaptive commitment** | Use confidence and entropy to decide how much text to release, with explicit budgets for continued refinement and forced progress. |
36
- | **Sparse multimodal foundation** | Select 8 of 128 experts per token and bring text, image, and video context into a shared decoder. |
37
-
38
- ## Inside the generation loop
39
-
40
- ### A canvas built for revision
41
-
42
- Mk2 maintains a rolling canvas of 256 candidate tokens. Each denoising pass updates the candidates using the encoded prompt, the current canvas, and latent context. Attention lets positions within the canvas inform one another before the output prefix is finalized. After a commit, the canvas shifts forward and opens space for new candidates.
43
-
44
- This creates two useful degrees of freedom: **how many times to refine** and **how many tokens to release**. A pass can commit a longer prefix when the policy permits, or spend additional computation refining an uncertain frontier.
45
-
46
- ### Deliberation that reads its own trajectory
47
-
48
- A 4-layer latent Transformer, 2,816 dimensions wide, builds working context before each decoder denoising pass. It reads the current canvas alongside three complementary sources of state:
49
-
50
- - **Per-token trajectory history:** 16 recent frames represented through four views—hidden state, first difference, second difference, and residual from the latest state—projected to rank 1,024. These give the processor access to the direction and stability of recent revisions. History follows surviving canvas tokens; new positions start empty.
51
- - **Denoising tape:** 16 pooled probes per frame summarize canvas activity in a time-indexed ring. The tape preserves recent step-level context as token positions move through the window.
52
- - **Persistent memory:** 256 slots of 2,816 dimensions carry learned summaries across commits within the generation.
53
-
54
- The resulting working context enters the decoder through its self-conditioning bridge and working-memory bus. Full-attention layers also read a separate persistent-memory bus. **The refinement history becomes an input to the next refinement.**
55
-
56
- ### Fast working state, lasting commit memory
57
-
58
- Working state is recomputed on every denoising pass. Persistent slots update only when tokens commit, using a Transformer writer with a separate gate for each slot. That writer draws on the committed region's trajectory, working state, and final decoder representation.
59
-
60
- This separates rapid revision from memory consolidation: the canvas can keep changing while persistent memory stays stable between commits. When the window advances, the memory slots remain available to later tokens.
61
-
62
- ### Compute that follows the commit frontier
63
-
64
- Mk2 combines proposal confidence with an excess-entropy penalty and selects the longest prefix whose cumulative failure score stays below the configured budget. A tighter budget requires stronger evidence before normal commitment; a looser budget admits longer prefixes for the same scores. Stagnation handling and a pondering watchdog bound continued refinement.
65
-
66
- Temperature and commit budget are exposed directly at inference time, making the generation policy adjustable per request. These scores govern commitment; they are not calibrated guarantees of factual correctness.
67
-
68
- ### One generation path for text, images, and video
69
-
70
- The Gemma 4 vision tower supplies visual features through the DiffusionGemma multimodal encoder. Text, image, and sampled video-frame inputs are encoded into the prefix KV cache that conditions the decoder. The rolling text canvas then uses the same latent deliberation and memory loop across all three input modalities.
71
-
72
- ## Model Summary
73
-
74
- | | |
75
- | --- | ---: |
76
- | Architecture | Mixture-of-Experts block diffusion + dual-timescale latent Transformer |
77
- | Model Size | 26B-A5B |
78
- | Vision Encoder | Gemma 4 Vision |
79
- | Layers | 30 |
80
- | Number of Experts | 128 |
81
- | Selected Experts per Token | 8 |
82
- | Vocabulary Size | 262,144 |
83
- | Configured Context Length | 262,144 tokens |
84
- | Sliding Window | 1024 |
85
- | Canvas Length | 256 |
86
- | Latent Width | 2,816 |
87
- | Latent Transformer | 4 layers, 16 attention heads |
88
- | Persistent Memory | 256 slots × 2,816 dimensions |
89
- | Memory Writer | 2-layer commit-sequence Transformer + per-slot gated writer |
90
- | Trajectory History | 16 frames, 4 views, rank 1,024 |
91
- | Denoise Tape | 16 probes per frame |
92
- | Modality | Text, Image, Video |
93
- | Preview checkpoint | Training step 900 |
94
- | Adaptation tokens | ~12.4 million |
95
-
96
- ## Preview release
97
-
98
- This first public preview includes **merged BF16 weights, inference code, processor, and tokenizer**. The text backbone and latent stack are exported from training step 900, after approximately 12.4 million adaptation tokens. The Gemma 4 vision tower is restored from DiffusionGemma.
99
-
100
- The release makes the architecture available for hands-on exploration and evaluation. Comprehensive benchmark results are not included; measured speed, reasoning quality, and multimodal reliability remain to be established for specific workloads.
101
-
102
- ## Getting Started
103
-
104
- Transformers 5.14.1 is the minimum supported version.
105
-
106
- ```shell
107
- pip install -U transformers torch accelerate
108
- ```
109
-
110
- ### Text generation
111
 
112
  ```python
113
  import torch
114
- from transformers import AutoModelForMultimodalLM, AutoProcessor
115
 
116
- model_id = "modilify/Modilify-Mk2-preview"
117
- processor = AutoProcessor.from_pretrained(model_id, trust_remote_code=True)
118
- model = AutoModelForMultimodalLM.from_pretrained(
119
- model_id,
120
  trust_remote_code=True,
121
  dtype=torch.bfloat16,
122
  device_map="auto",
123
- )
124
-
125
- messages = [{"role": "user", "content": "Explain why the sky is blue."}]
126
- inputs = processor.apply_chat_template(
127
- messages,
128
  tokenize=True,
129
  add_generation_prompt=True,
130
  enable_thinking=False,
131
  return_dict=True,
132
  return_tensors="pt",
133
  ).to(model.device)
134
-
135
- output = model.generate(
136
- **inputs,
137
- max_new_tokens=256,
138
- denoise_temperature=0.8,
139
- commit_failure_budget=0.2,
140
- )
141
- new_tokens = output.sequences[:, inputs["input_ids"].shape[1]:]
142
- print(processor.batch_decode(new_tokens, skip_special_tokens=False)[0])
143
- ```
144
-
145
- ### Image input
146
-
147
- ```python
148
- from PIL import Image
149
-
150
- image = Image.open("example.jpg").convert("RGB")
151
- messages = [{
152
- "role": "user",
153
- "content": [
154
- {"type": "image", "image": image},
155
- {"type": "text", "text": "Describe the image and identify uncertainty."},
156
- ],
157
- }]
158
- inputs = processor.apply_chat_template(
159
- messages,
160
- tokenize=True,
161
- add_generation_prompt=True,
162
- enable_thinking=True,
163
- return_dict=True,
164
- return_tensors="pt",
165
- ).to(model.device)
166
- output = model.generate(**inputs, max_new_tokens=256)
167
- ```
168
-
169
- ### Video-frame input
170
-
171
- 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.
172
-
173
- ```python
174
- import av
175
- from PIL import Image
176
-
177
- container = av.open("short_clip.mp4")
178
- decoded = [Image.fromarray(frame.to_rgb().to_ndarray()) for frame in container.decode(video=0)]
179
- stride = max(1, len(decoded) // 32)
180
- frames = decoded[::stride][:32]
181
-
182
- messages = [{
183
- "role": "user",
184
- "content": [
185
- {"type": "video", "video": frames},
186
- {"type": "text", "text": "Summarize the main visual events in order."},
187
- ],
188
- }]
189
- inputs = processor.apply_chat_template(
190
- messages,
191
- tokenize=True,
192
- add_generation_prompt=True,
193
- enable_thinking=True,
194
- return_dict=True,
195
- return_tensors="pt",
196
- ).to(model.device)
197
- output = model.generate(**inputs, max_new_tokens=256)
198
- ```
199
-
200
- ## Thinking mode
201
-
202
- Latent deliberation runs inside the generation loop regardless of the chat template's thinking flag. The flag controls the prompt's request for a textual thought channel:
203
-
204
- - `enable_thinking=True` inserts a system turn that contains `<|think|>` and still ends the prompt at `<|turn>model`.
205
- - `enable_thinking=False` does **not** inject an empty thought channel. The prompt ends at `<|turn>model`.
206
-
207
- The model may still open `<|channel>thought` on its own when thinking is disabled. The template flag does not guarantee suppression of generated thought-channel text. Applications should handle that channel explicitly before displaying an answer.
208
-
209
- ## Configurable inference
210
-
211
- 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.
212
-
213
- | Parameter | Default | Meaning |
214
- | --- | ---: | --- |
215
- | `denoise_temperature` | 0.8 | Sampling temperature for every canvas step |
216
- | `commit_failure_budget` | 0.2 | Cumulative prefix risk limit for normal commits |
217
- | `jump_failure_budget` | 2.0 | Cumulative risk limit for forced jumps |
218
- | `jump_on_no_progress_after` | 12 | Stagnation steps before a forced jump |
219
- | `max_ponder_steps` | 64 | Watchdog multiplier per requested token |
220
- | `min_trajectory_progress` | 0.005 | Minimum fused-risk improvement counted as progress |
221
- | `canvas_length` | 256 | Rolling diffusion canvas length |
222
- | `repetition_penalty` | 1.0 | Transformers-style repetition penalty |
223
- | `turn_end_token_id` | 106 | Gemma turn terminator |
224
-
225
- Call-site override:
226
-
227
- ```python
228
- output = model.generate(
229
- **inputs,
230
- max_new_tokens=256,
231
- denoise_temperature=0.4,
232
- commit_failure_budget=0.05,
233
- )
234
- ```
235
-
236
- Load-time override:
237
-
238
- ```python
239
- from transformers import AutoConfig
240
-
241
- config = AutoConfig.from_pretrained(model_id, trust_remote_code=True)
242
- config.denoise_temperature = 0.4
243
- config.commit_failure_budget = 0.05
244
- model = AutoModelForMultimodalLM.from_pretrained(
245
- model_id,
246
- config=config,
247
- trust_remote_code=True,
248
- dtype=torch.bfloat16,
249
- device_map="auto",
250
- )
251
- ```
252
-
253
- A tighter commit budget allows fewer tokens for the same confidence-and-entropy scores; a looser budget allows more. Temperature changes the sampling distribution and also affects those scores, so its effect on throughput depends on the prompt and generation trajectory. Measure latency and output quality together when tuning these controls.
254
-
255
- 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.
256
-
257
- ## Evaluation status, limitations, and risks
258
-
259
- This preview does not include a complete accuracy, robustness, calibration, fairness, or safety evaluation. Architectural features describe how Mk2 generates; they do not establish benchmark superiority or fitness for a particular deployment. The configured context limit is not a validated long-context quality result.
260
-
261
- 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.
262
-
263
- 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.
264
-
265
- ## License
266
-
267
- 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).
268
-
269
- ## Citation
270
-
271
- ```bibtex
272
- @software{modilify_mk2_preview_2026,
273
- title = {Modilify Mk2 Preview: 26B-A5B},
274
- author = {Modilify},
275
- year = {2026},
276
- note = {A multimodal dual-timescale latent-deliberation derivative of DiffusionGemma}
277
- }
278
  ```
279
 
280
- Also cite the upstream DiffusionGemma release as requested by Google DeepMind.
 
 
 
 
 
3
  license_name: modilify-open-model-license-1.0
4
  license_link: LICENSE
5
  library_name: transformers
6
+ pipeline_tag: text-generation
7
  tags:
8
  - diffusion
 
 
9
  - mixture-of-experts
10
  - trust-remote-code
11
+ - safetensors
12
  ---
13
 
14
  ![LOGO](assets/01-LOGO.jpg)
15
 
16
+ # Modilify Mk2 Preview
17
 
18
+ PyTorch text inference for the **step 1250, schema25** checkpoint migrated
19
+ from `Modilify-Mk2-preview-mlx`. The model uses a shared DiffusionGemma text
20
+ encoder/decoder, a rolling 256-token canvas, and dual-timescale GDN2 memory.
21
+ The release contains 32 Safetensors shards totaling **48.23 GiB**.
22
 
23
+ Dense and expert LoRA adapters remain unfused. Model loading preserves the
24
+ 26 FP32 GDN2/norm parameters alongside BF16 weights. No vision tower is included.
25
 
26
+ ## Inference
27
 
28
+ Use Transformers **5.14.1**, PyTorch, Accelerate, and Safetensors. Run from
29
+ this directory or replace `model_path` with its location. Inference needs
30
+ additional memory beyond the weights.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
31
 
32
  ```python
33
  import torch
34
+ from transformers import AutoModelForCausalLM, AutoTokenizer
35
 
36
+ model_path = "."
37
+ tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)
38
+ model = AutoModelForCausalLM.from_pretrained(
39
+ model_path,
40
  trust_remote_code=True,
41
  dtype=torch.bfloat16,
42
  device_map="auto",
43
+ attn_implementation="sdpa",
44
+ ).eval()
45
+ inputs = tokenizer.apply_chat_template(
46
+ [{"role": "user", "content": "Explain why the sky is blue."}],
 
47
  tokenize=True,
48
  add_generation_prompt=True,
49
  enable_thinking=False,
50
  return_dict=True,
51
  return_tensors="pt",
52
  ).to(model.device)
53
+ output = model.generate(**inputs, max_new_tokens=128, seed=42)
54
+ print(tokenizer.decode(output.sequences[0, inputs.input_ids.shape[1]:],
55
+ skip_special_tokens=True))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
56
  ```
57
 
58
+ For MPS, use `device_map={"": "mps"}`. Set `enable_thinking=True` in the chat
59
+ template to request thinking. Generation uses temperature 0.8, top-k 40,
60
+ min-p 0.05, target confidence 0.5, and failure budget 0.2. `max_denoising_steps`
61
+ can bound generation. Static batches and continuous batching share the
62
+ confidence-prefix commit policy; each request retains its own cache and RNG.
chat_template.jinja CHANGED
@@ -1,6 +1,6 @@
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
  #}
@@ -268,7 +268,7 @@
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 -%}
@@ -384,4 +384,4 @@
384
  {%- elif ns.prev_message_type == 'tool_response' and enable_thinking -%}
385
  {{- '<|channel>thought\n' -}}
386
  {%- endif -%}
387
- {%- endif -%}
 
1
  {#
2
+ Template: Google Gemma 4 Canonical Chat Template
3
+ Author: Google Gemma Engineering Team
4
  Published: 2026-07-09
5
  Context: Fixed tool-calling loops, turn closures, and thinking content-ordering.
6
  #}
 
268
 
269
  {%- set ns_tr_out = namespace(flag=false) -%}
270
  {%- if message.get('tool_responses') -%}
271
+ {#- Legacy: tool_responses embedded on the assistant message (Google/Gemma native) -#}
272
  {%- for tool_response in message.get('tool_responses') -%}
273
  {{- format_tool_response_block(tool_response['name'] | default('unknown', true), tool_response['response']) -}}
274
  {%- set ns_tr_out.flag = true -%}
 
384
  {%- elif ns.prev_message_type == 'tool_response' and enable_thinking -%}
385
  {{- '<|channel>thought\n' -}}
386
  {%- endif -%}
387
+ {%- endif -%}
commit_policy.py CHANGED
@@ -1,13 +1,11 @@
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 (
@@ -20,12 +18,32 @@ JUMP_FAILURE_BUDGET = 2.0
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
 
@@ -34,11 +52,9 @@ def fused_commit_confidence(
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)
@@ -48,12 +64,23 @@ def fused_commit_confidence(
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
 
@@ -71,7 +98,7 @@ def fused_commit_failure_rate(
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
@@ -184,7 +211,6 @@ def select_commit_lengths(
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.
@@ -242,8 +268,6 @@ def select_commit_lengths(
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,
@@ -256,9 +280,7 @@ def select_commit_lengths(
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,
@@ -291,15 +313,3 @@ def select_commit_lengths(
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
- ]
 
1
+ """Confidence-and-entropy prefix commitment."""
 
 
2
 
3
  from __future__ import annotations
4
 
5
  from collections.abc import Sequence
6
  from dataclasses import dataclass
 
7
 
8
+ import math
9
  import torch
10
 
11
  from .latent_deliberation import (
 
18
  FUSED_EPS = 1e-6
19
 
20
 
21
+ def commit_target_confidence_bias(
22
+ target_confidence: float | None,
23
+ failure_budget: float = 0.2,
24
+ budget_safety_ratio: float = 0.85,
25
+ ) -> float:
26
+ if target_confidence is None or target_confidence <= 0.0 or target_confidence >= 1.0:
27
+ return 0.0
28
+ target_failure = min(budget_safety_ratio * failure_budget, 1.0 - target_confidence)
29
+ target_failure = max(target_failure, 1.0e-4)
30
+ target_conf = 1.0 - target_failure
31
+ logit_c = math.log(target_conf / (1.0 - target_conf))
32
+ logit_p = math.log(target_confidence / (1.0 - target_confidence))
33
+ return max(logit_c - logit_p, 0.0)
34
+
35
+
36
  def fused_commit_confidence(
37
  proposal_confidence: torch.Tensor,
38
  token_entropy: torch.Tensor,
39
  *,
 
40
  eps: float = FUSED_EPS,
41
+ entropy_weight: float = 1.0,
42
+ confidence_power: float = 1.0,
43
+ top_k: int | None = None,
44
+ min_p: float | None = None,
45
+ target_confidence: float | None = None,
46
+ failure_budget: float = 0.2,
47
  ) -> torch.Tensor:
48
  """Fuse proposal confidence with token entropy into effective commit confidence.
49
 
 
52
  p = clamp(proposal_confidence, eps, 1 - eps)
53
  h2 = -p * log(p) - (1 - p) * log(1 - p) # binary entropy of p
54
  excess = max(token_entropy - h2, 0)
55
+ base_fused = sigmoid(logit(p) + bias - entropy_weight * excess)
56
+ fused = base_fused ** confidence_power
57
 
 
 
58
  """
59
 
60
  p = proposal_confidence.float().nan_to_num(0.5).clamp(min=eps, max=1.0 - eps)
 
64
  h2 = -p * torch.log(p) - (1.0 - p) * torch.log1p(-p)
65
 
66
  excess = (H - h2).clamp(min=0.0)
67
+ k_eff = None
68
+ if top_k is not None and top_k > 0:
69
+ k_eff = torch.full_like(p, float(top_k))
70
+ if min_p is not None and min_p > 0:
71
+ thresh = torch.clamp_min(float(min_p) * p, 1.0e-6)
72
+ k_min_p = torch.clamp_min((1.0 - p) / thresh, 1.0)
73
+ k_eff = k_min_p if k_eff is None else torch.minimum(k_eff, k_min_p)
74
+ if k_eff is not None:
75
+ max_excess = (1.0 - p) * torch.log(k_eff)
76
+ excess = torch.minimum(excess, max_excess)
77
 
78
  # logit(p) = log(p / (1-p)) = log(p) - log1p(-p)
79
  logit_p = torch.log(p) - torch.log1p(-p)
80
 
81
+ bias = commit_target_confidence_bias(target_confidence, failure_budget=failure_budget)
82
+ base_fused = torch.sigmoid(logit_p + bias - float(entropy_weight) * excess)
83
+ fused = base_fused.pow(float(confidence_power))
84
  return fused.clamp(min=eps, max=1.0 - eps)
85
 
86
 
 
98
 
99
  @dataclass(frozen=True)
100
  class CommitPolicyDecision:
101
+ """Proposal-to-commit transition."""
102
 
103
  normal_lengths: torch.LongTensor
104
  commit_lengths: torch.LongTensor
 
211
  stagnation_threshold: int,
212
  min_progress: float,
213
  max_ponder_steps: int | None = None,
 
214
  valid_mask: torch.BoolTensor | None = None,
215
  ) -> CommitPolicyDecision:
216
  """Use normal sampled commits and a fixed-budget greedy JUMP.
 
268
  stagnation_steps,
269
  commit_lengths=normal,
270
  active_rows=active_rows,
 
 
271
  )
272
  jump_rows = normal.eq(0) & active_rows & should_force_trajectory_jump(
273
  next_stagnation,
 
280
  jump_commit = bounded_prefix_failure_commit_lengths(
281
  greedy_token_ids,
282
  jump_failure_rate,
283
+ failure_budget=JUMP_FAILURE_BUDGET,
 
 
284
  remaining_lengths=remaining_lengths,
285
  stop_token_id=stop_token_id,
286
  valid_mask=valid_mask,
 
313
  ponder_steps=next_ponder,
314
  stagnation_steps=next_stagnation,
315
  )
 
 
 
 
 
 
 
 
 
 
 
 
config.json CHANGED
@@ -1,54 +1,27 @@
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
@@ -122,56 +95,40 @@
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
  }
 
1
  {
 
 
 
 
 
 
 
 
 
 
 
2
  "canvas_length": 256,
3
+ "commit_confidence_power": 1.0,
4
+ "commit_entropy_weight": 1.0,
5
  "commit_failure_budget": 0.2,
6
+ "commit_min_p": 0.05,
7
  "commit_sequence_dim": 1024,
8
+ "commit_target_confidence": 0.5,
9
+ "commit_top_k": 40,
10
+ "eos_token_id": 1,
 
 
 
 
 
 
 
11
  "initializer_range": 0.02,
 
 
12
  "kv_cache_bucket_size": 128,
13
  "latent_dim": 2816,
 
14
  "latent_ffn_dim": 7168,
15
  "latent_history_kv_rank": 1024,
 
 
16
  "latent_local_attention_window": 128,
 
17
  "latent_num_heads": 16,
18
  "latent_num_layers": 4,
19
+ "latent_tape_probes": 4,
20
  "latent_working_last_block_global": true,
21
+ "memory_architecture": "compact_gdn2_v2",
 
 
22
  "model_type": "modilify_mk2",
 
23
  "persistent_memory_bus": true,
24
+ "state_schema_version": 25,
 
 
25
  "terminal_token_ids": [
26
  106,
27
  50
 
95
  "sliding_window": 1024,
96
  "tie_word_embeddings": true,
97
  "top_k_experts": 8,
98
+ "use_bidirectional_attention": null,
99
  "vocab_size": 262144
100
  },
101
  "tie_word_embeddings": true,
102
  "transformers_version": "5.14.1",
103
  "turn_end_token_id": 106,
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
104
  "vocab_chunk_size": 32768,
105
  "working_memory_bus": true,
106
+ "architectures": [
107
+ "ModilifyMk2ForBlockDiffusion"
108
+ ],
109
+ "auto_map": {
110
+ "AutoConfig": "configuration_modilify_mk2.ModilifyMk2Config",
111
+ "AutoModel": "modeling_modilify_mk2.ModilifyMk2Model",
112
+ "AutoModelForCausalLM": "modeling_modilify_mk2.ModilifyMk2ForBlockDiffusion"
113
+ },
114
+ "dtype": "bfloat16",
115
+ "lora_config": {
116
+ "r": 16,
117
+ "alpha": 16,
118
+ "target_modules": [
119
+ "q_proj",
120
+ "k_proj",
121
+ "v_proj",
122
+ "o_proj",
123
+ "gate_proj",
124
+ "up_proj",
125
+ "down_proj",
126
+ "proj"
127
+ ],
128
+ "expert_r": 8,
129
+ "expert_alpha": 8
130
+ },
131
+ "precision_policy": "gdn2_small_fp32_v1",
132
+ "global_step": 1250,
133
+ "base_model": "google/diffusiongemma-26B-A4B-it"
134
  }
configuration_modilify_mk2.py CHANGED
@@ -1,323 +1,99 @@
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
- ]
 
1
+ """Text-only schema25 inference configuration."""
 
 
 
2
  from __future__ import annotations
 
 
3
  from collections.abc import Sequence
4
  from typing import Any
5
+ from transformers import PreTrainedConfig
6
+ from transformers.models.diffusion_gemma import DiffusionGemmaTextConfig
7
 
 
 
 
 
 
 
 
 
 
 
8
  DENOISE_TEMPERATURE = 0.8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9
 
10
  class ModilifyMk2TextConfig(DiffusionGemmaTextConfig):
 
 
11
  model_type = "modilify_mk2_text"
12
+ vocab_size: int = 262_144
13
+ hidden_size: int = 2816
14
+ intermediate_size: int = 2112
15
+ num_hidden_layers: int = 30
16
+ num_attention_heads: int = 16
17
+ num_key_value_heads: int = 8
18
+ head_dim: int = 256
19
+ max_position_embeddings: int = 262_144
20
+ sliding_window: int = 1024
21
+ use_bidirectional_attention: str | None = None
22
+ num_global_key_value_heads: int | None = 2
23
+ global_head_dim: int = 512
24
+ num_experts: int | None = 128
25
+ top_k_experts: int | None = 8
26
+ moe_intermediate_size: int | None = 704
27
+
28
+ class ModilifyMk2Config(PreTrainedConfig):
 
 
 
 
 
 
 
 
 
 
 
 
 
29
  model_type = "modilify_mk2"
30
+ sub_configs = {"text_config": ModilifyMk2TextConfig}
 
 
 
 
 
 
 
31
 
32
  def __init__(
33
+ self, text_config: ModilifyMk2TextConfig | dict[str, Any] | None = None, *,
 
 
 
 
 
 
 
 
 
 
 
34
  canvas_length: int = 256,
35
  initializer_range: float = 0.02,
36
+ tie_word_embeddings: bool = True,
37
+ state_schema_version: int = 25,
38
+ memory_architecture: str = "compact_gdn2_v2",
39
  latent_dim: int = 2816,
40
  latent_ffn_dim: int = 7168,
 
41
  latent_num_layers: int = 4,
42
  latent_num_heads: int = 16,
43
  latent_local_attention_window: int = 128,
44
+ latent_tape_probes: int = 4,
45
+ latent_history_kv_rank: int = 1024,
 
 
 
46
  latent_working_last_block_global: bool = True,
47
  working_memory_bus: bool = True,
48
  persistent_memory_bus: bool = True,
49
+ commit_sequence_dim: int = 1024,
 
 
 
 
50
  kv_cache_bucket_size: int = 128,
51
+ vocab_chunk_size: int = 32768,
52
  turn_end_token_id: int = 106,
53
+ terminal_token_ids: Sequence[int] = (106, 50),
54
+ eos_token_id: int = 1,
55
+ commit_failure_budget: float = 0.2,
56
+ commit_top_k: int | None = 40,
57
+ commit_min_p: float | None = 0.05,
58
+ commit_target_confidence: float | None = 0.5,
59
+ commit_entropy_weight: float = 1.0,
60
+ commit_confidence_power: float = 1.0,
61
  **kwargs: Any,
62
  ) -> None:
63
+ if state_schema_version != 25 or memory_architecture != "compact_gdn2_v2":
64
+ raise ValueError("This release requires schema25 compact_gdn2_v2 weights.")
65
+ if text_config is None:
66
+ text_config = ModilifyMk2TextConfig()
 
 
 
 
 
67
  elif isinstance(text_config, dict):
68
+ payload = dict(text_config)
69
+ payload.pop("model_type", None)
70
+ text_config = ModilifyMk2TextConfig(**payload)
71
+ text_config.use_bidirectional_attention = None
72
+ self.text_config = text_config
73
+ self.canvas_length = canvas_length
74
+ self.initializer_range = initializer_range
75
+ self.state_schema_version = state_schema_version
76
+ self.memory_architecture = memory_architecture
77
+ self.latent_dim = latent_dim
78
+ self.latent_ffn_dim = latent_ffn_dim
79
+ self.latent_num_layers = latent_num_layers
80
+ self.latent_num_heads = latent_num_heads
81
+ self.latent_local_attention_window = latent_local_attention_window
82
+ self.latent_tape_probes = latent_tape_probes
83
+ self.latent_history_kv_rank = latent_history_kv_rank
84
+ self.latent_working_last_block_global = latent_working_last_block_global
85
+ self.working_memory_bus = working_memory_bus
86
+ self.persistent_memory_bus = persistent_memory_bus
87
+ self.commit_sequence_dim = commit_sequence_dim
88
+ self.kv_cache_bucket_size = kv_cache_bucket_size
89
+ self.vocab_chunk_size = vocab_chunk_size
90
+ self.turn_end_token_id = turn_end_token_id
91
+ self.terminal_token_ids = terminal_token_ids
92
+ self.commit_failure_budget = commit_failure_budget
93
+ self.commit_top_k = commit_top_k
94
+ self.commit_min_p = commit_min_p
95
+ self.commit_target_confidence = commit_target_confidence
96
+ self.commit_entropy_weight = commit_entropy_weight
97
+ self.commit_confidence_power = commit_confidence_power
98
+ super().__init__(tie_word_embeddings=tie_word_embeddings,
99
+ eos_token_id=eos_token_id, **kwargs)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
continuous_batching.py CHANGED
@@ -1,6 +1,4 @@
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
@@ -47,21 +45,9 @@ from .generation_modilify_mk2 import (
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(
@@ -100,16 +86,13 @@ def continuous_config_fingerprint(
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
@@ -121,8 +104,6 @@ class ModilifyMk2ContinuousGenerationOutput(GenerationOutput):
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."""
@@ -133,7 +114,6 @@ class ModilifyMk2ContinuousGenerationOutput(GenerationOutput):
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
@@ -142,7 +122,6 @@ class ModilifyMk2RequestState:
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
@@ -178,10 +157,9 @@ def _slice_rolling_state(state: ModilifyMk2RollingState, row: int) -> ModilifyMk
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
 
@@ -190,10 +168,9 @@ def _pack_rolling_states(states: Sequence[ModilifyMk2RollingState]) -> ModilifyM
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
 
@@ -320,9 +297,6 @@ class ModilifyMk2ContinuousBatchingManager:
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)
@@ -348,6 +322,12 @@ class ModilifyMk2ContinuousBatchingManager:
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
@@ -707,20 +687,12 @@ class ModilifyMk2ContinuousBatchingManager:
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.")
@@ -788,7 +760,6 @@ class ModilifyMk2ContinuousBatchingManager:
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
@@ -955,8 +926,6 @@ class ModilifyMk2ContinuousBatchingManager:
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,
@@ -974,16 +943,6 @@ class ModilifyMk2ContinuousBatchingManager:
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(
@@ -1052,21 +1011,15 @@ class ModilifyMk2ContinuousBatchingManager:
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,
@@ -1077,22 +1030,9 @@ class ModilifyMk2ContinuousBatchingManager:
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)
@@ -1203,8 +1143,7 @@ class ModilifyMk2ContinuousBatchingManager:
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,
@@ -1222,6 +1161,14 @@ class ModilifyMk2ContinuousBatchingManager:
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):
@@ -1261,15 +1208,12 @@ class ModilifyMk2ContinuousBatchingManager:
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,
@@ -1295,46 +1239,56 @@ class ModilifyMk2ContinuousBatchingManager:
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,
@@ -1372,17 +1326,13 @@ class ModilifyMk2ContinuousBatchingManager:
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(
@@ -1399,6 +1349,8 @@ class ModilifyMk2ContinuousBatchingManager:
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)
@@ -1462,31 +1414,6 @@ class ModilifyMk2ContinuousBatchingManager:
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)
@@ -1691,7 +1618,6 @@ def generate_static_batch_with_logical_cache(
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(
@@ -1705,22 +1631,6 @@ def generate_static_batch_with_logical_cache(
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
- ]
 
1
+ """Continuous batching for ModilifyMk2 behind the Transformers public API shape.
 
 
2
 
3
  The upstream continuous runner is autoregressive: it persists every query in a
4
  paged cache and emits exactly one token per request and step. ModilifyMk2 instead
 
45
  NoiseCanvasSampler,
46
  _add_repetition_history,
47
  _flatten_token_ids,
 
 
 
 
 
 
 
 
 
 
 
 
 
 
48
  )
49
+ from .latent_deliberation import LatentDeliberationState, cat_latent_states, infer_commit_reason, slice_latent_state
50
+ from .generation_modilify_mk2 import deterministic_episode_iteration_bound
51
 
52
 
53
  _TERMINAL_REASONS = frozenset(
 
86
 
87
  @dataclass
88
  class ModilifyMk2ContinuousGenerationOutput(GenerationOutput):
89
+ """Request output for the continuous generation runner."""
 
90
  stop_reason: str | None = None
91
  committed_tokens: int = 0
92
  denoise_steps: int = 0
93
  no_progress_steps: int = 0
94
  jump_count: int = 0
95
  forced_jump_bad_count: int = 0
 
 
96
  average_commit_len: float = 0.0
97
  tokens_per_forward: float = 0.0
98
  seed: int | None = None
 
104
  is_stream_update: bool = False
105
  delta_tokens: list[int] = field(default_factory=list)
106
  state_shift_count: int = 0
 
 
107
 
108
  def is_finished(self) -> bool:
109
  """Treat failed/cancelled requests as terminal for every consumer API."""
 
114
  @dataclass
115
  class ModilifyMk2RequestState:
116
  """All mutable state required to suspend and re-batch one request."""
 
117
  request_id: str
118
  prompt_ids: list[int]
119
  max_new_tokens: int
 
122
  record_timestamps: bool
123
  seed: int
124
  max_denoising_steps: int | None
 
125
  created_time: float = field(default_factory=time.perf_counter)
126
  status: RequestStatus = RequestStatus.PENDING
127
  started_time: float = -1.0
 
157
  canvas=_clone_tensor_row(state.canvas, row),
158
  confidence=_clone_tensor_row(state.confidence, row),
159
  entropy=_clone_tensor_row(state.entropy, row),
 
160
  latent_state=slice_latent_state(state.latent_state, selected),
161
+
162
+ head=_clone_tensor_row(state.head, row),
163
  )
164
 
165
 
 
168
  canvas=torch.cat([state.canvas for state in states], dim=0),
169
  confidence=torch.cat([state.confidence for state in states], dim=0),
170
  entropy=torch.cat([state.entropy for state in states], dim=0),
 
171
  latent_state=cat_latent_states([state.latent_state for state in states]),
172
+
173
+ head=torch.cat([state.head for state in states], dim=0),
174
  )
175
 
176
 
 
297
  workload_hints: Any = None,
298
  ) -> None:
299
  del workload_hints
 
 
 
300
  self.model = model
301
  self.generation_config = copy.deepcopy(
302
  generation_config or getattr(model, "generation_config", None)
 
322
  self.sampler: NoiseCanvasSampler = model._prepare_sampler(
323
  self.generation_config, model.config.canvas_length
324
  )
325
+ pad_token_id = self.generation_config.pad_token_id
326
+ if pad_token_id is None:
327
+ pad_token_id = getattr(model.config, "pad_token_id", None)
328
+ if isinstance(pad_token_id, (list, tuple)):
329
+ pad_token_id = pad_token_id[0]
330
+ self.pad_token_id = int(0 if pad_token_id is None else pad_token_id)
331
  self.run_id = uuid.uuid4().hex
332
  self.warmed_up = False
333
  self.destroyed = False
 
687
  ):
688
  raise ValueError("`input_ids` must be a non-empty list of integer token IDs.")
689
  seed = request_kwargs.pop("seed", None)
 
690
  max_denoising_steps = request_kwargs.pop(
691
  "max_denoising_steps", self.generation_config.max_denoising_steps
692
  )
693
  if request_kwargs:
694
  unsupported = ", ".join(sorted(request_kwargs))
695
  raise ValueError(f"Unsupported per-request generation options: {unsupported}")
 
 
 
 
 
 
 
696
  limit = self.generation_config.max_new_tokens if max_new_tokens is None else max_new_tokens
697
  if not isinstance(limit, int) or limit <= 0:
698
  raise ValueError("`max_new_tokens` must be a positive integer.")
 
760
  record_timestamps=bool(record_timestamps),
761
  seed=resolved_seed & ((1 << 63) - 1),
762
  max_denoising_steps=max_denoising_steps,
 
763
  )
764
  state.reserved_blocks = math.ceil(
765
  (len(state.prompt_ids) + state.max_new_tokens) / self.block_size
 
926
  ),
927
  jump_count=state.jumps,
928
  forced_jump_bad_count=state.forced_jump_tokens,
 
 
929
  average_commit_len=len(state.generated_tokens) / shifts,
930
  tokens_per_forward=len(state.generated_tokens) / steps,
931
  seed=state.seed,
 
943
  state.last_delta_tokens if delta_tokens is None else delta_tokens
944
  ),
945
  state_shift_count=state.shifts,
 
 
 
 
 
 
 
 
 
 
946
  )
947
 
948
  def _finish(
 
1011
  generator = torch.Generator(device=self.device)
1012
  generator.manual_seed(state.seed)
1013
  state.generator = generator
1014
+ canvas = self.sampler.initialize_canvas(
1015
+ 1, self.device, generators=[generator]
1016
+ )
1017
+ canvas[:, state.max_new_tokens:] = self.pad_token_id
 
 
 
1018
  canvas_length = int(self.model.config.canvas_length)
1019
  latent = LatentDeliberationState.empty(
1020
  batch_size=1,
1021
  canvas_length=canvas_length,
 
 
1022
  device=self.device,
 
1023
  )
1024
  state.rolling_state = ModilifyMk2RollingState(
1025
  canvas=canvas,
 
1030
  device=self.device,
1031
  dtype=torch.float32,
1032
  ),
 
1033
  latent_state=latent,
1034
+
1035
+ head=torch.zeros(1, device=self.device, dtype=torch.long),
 
 
 
 
 
 
 
 
 
 
 
 
1036
  )
1037
  state.cache = self.cache_pool.prefill(state.prompt_ids)
1038
  state.logical_length = len(state.prompt_ids)
 
1143
  stagnation_steps=rolling.latent_state.stagnation_steps[row : row + 1],
1144
  active_rows=torch.ones(1, device=self.device, dtype=torch.bool),
1145
  remaining_lengths=torch.tensor([remaining], device=self.device),
1146
+ failure_budget=self.model.config.commit_failure_budget,
 
1147
  stop_token_id=state.eos_token_ids,
1148
  max_ponder_steps=self.generation_config.max_ponder_steps,
1149
  stagnation_threshold=self.generation_config.jump_on_no_progress_after,
 
1161
 
1162
  @torch.inference_mode()
1163
  def _run_batch_step(self, states: Sequence[ModilifyMk2RequestState]) -> list[str]:
1164
+ if len(states) > 1:
1165
+ # The BF16 trunk changes reductions with batch shape. Routed
1166
+ # experts amplify those differences into different commits, so
1167
+ # run each active row with its independent-request numerics.
1168
+ finished = []
1169
+ for state in states:
1170
+ finished.extend(self._run_batch_step([state]))
1171
+ return finished
1172
  started = time.perf_counter()
1173
  rolling_states = [state.rolling_state for state in states]
1174
  if any(state is None for state in rolling_states):
 
1208
  decoder_input_ids=rolling.canvas,
1209
  previous_confidence=rolling.confidence,
1210
  previous_entropy=rolling.entropy,
 
1211
  latent_state=rolling.latent_state,
1212
+
 
1213
  decoder_position_ids=decoder_positions,
1214
  decoder_read_cache=True,
1215
  decoder_attention_mask=decoder_mask,
1216
  compact_vocab=True,
 
1217
  repetition_token_mask=repetition_history,
1218
  repetition_penalty=self.repetition_penalty,
1219
  sampling_generators=generators,
 
1239
  output.next_latent_state,
1240
  confidence=next_confidence.detach().float(),
1241
  entropy=token_entropy.detach().float(),
 
 
 
 
1242
  )
1243
+ remaining_lengths = torch.tensor(
1244
+ [state.max_new_tokens - len(state.generated_tokens) for state in states],
1245
+ device=self.device,
1246
+ dtype=torch.long,
1247
+ )
1248
+ live_mask = torch.arange(canvas_length, device=self.device)[None, :].lt(
1249
+ remaining_lengths[:, None]
1250
  )
1251
+ next_latent = self.model.latent_deliberation.observe_state(
1252
+ next_latent, output.heavy_hidden_state, output.working_state,
1253
+ live_mask, rolling.head,
1254
  )
1255
  next_state = ModilifyMk2RollingState(
1256
  canvas=next_canvas,
1257
  confidence=next_confidence,
1258
  entropy=token_entropy,
 
1259
  latent_state=next_latent,
1260
+
1261
+ head=rolling.head,
 
 
 
 
 
 
1262
  )
1263
  normal_failure_rate = fused_commit_failure_rate(
1264
  proposal_confidence,
1265
  token_entropy,
1266
+ entropy_weight=self.model.config.commit_entropy_weight,
1267
+ confidence_power=self.model.config.commit_confidence_power,
1268
+ top_k=getattr(self.model.config, "commit_top_k", None),
1269
+ min_p=getattr(self.model.config, "commit_min_p", None),
1270
+ target_confidence=getattr(self.model.config, "commit_target_confidence", None),
1271
+ failure_budget=float(self.model.config.commit_failure_budget),
1272
  )
1273
  jump_failure_rate = fused_commit_failure_rate(
1274
  greedy_confidence,
1275
  token_entropy,
1276
+ entropy_weight=self.model.config.commit_entropy_weight,
1277
+ confidence_power=self.model.config.commit_confidence_power,
1278
+ top_k=getattr(self.model.config, "commit_top_k", None),
1279
+ min_p=getattr(self.model.config, "commit_min_p", None),
1280
+ target_confidence=getattr(self.model.config, "commit_target_confidence", None),
1281
+ failure_budget=float(self.model.config.commit_failure_budget),
1282
  )
1283
  previous_failure_rate = fused_commit_failure_rate(
1284
  rolling.confidence,
1285
  rolling.entropy,
1286
+ entropy_weight=self.model.config.commit_entropy_weight,
1287
+ confidence_power=self.model.config.commit_confidence_power,
1288
+ top_k=getattr(self.model.config, "commit_top_k", None),
1289
+ min_p=getattr(self.model.config, "commit_min_p", None),
1290
+ target_confidence=getattr(self.model.config, "commit_target_confidence", None),
1291
+ failure_budget=float(self.model.config.commit_failure_budget),
1292
  )
1293
  (
1294
  normal_commit,
 
1326
  stagnation_steps=next_stagnation,
1327
  ),
1328
  )
1329
+ if output.working_state is None:
 
 
 
1330
  raise RuntimeError("Forward did not return working trajectory features.")
1331
  next_state = self.model._write_committed_memory(
 
1332
  next_state=next_state,
1333
  working_state=output.working_state,
 
1334
  heavy_hidden=output.heavy_hidden_state,
1335
+ commit_token_ids=commit_token_ids,
1336
  commit_lengths=commit_lengths,
1337
  prefix_lengths=logical_lengths,
1338
  commit_reason=infer_commit_reason(
 
1349
  commit_lengths,
1350
  self.sampler,
1351
  generators=generators,
1352
+ remaining_lengths=remaining_lengths - commit_lengths,
1353
+ pad_token_id=self.pad_token_id,
1354
  )
1355
  shifted_states = [
1356
  _slice_rolling_state(shifted, row) for row in range(batch_size)
 
1414
  reason = "episode_watchdog"
1415
 
1416
  elapsed = time.perf_counter() - started
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1417
 
1418
  if reason is not None:
1419
  self._finish(state, reason)
 
1618
  return torch.tensor(
1619
  [getattr(output, name) for output in ordered],
1620
  device=input_ids.device,
 
1621
  )
1622
 
1623
  return ModilifyMk2GenerationOutput(
 
1631
  no_progress_steps=tensor("no_progress_steps", dtype=torch.long),
1632
  jump_count=tensor("jump_count", dtype=torch.long),
1633
  forced_jump_bad_count=tensor("forced_jump_bad_count", dtype=torch.long),
 
 
 
 
1634
  average_commit_len=tensor("average_commit_len", dtype=torch.float32),
1635
  state_shift_count=tensor("state_shift_count", dtype=torch.long),
 
 
1636
  )
 
 
 
 
 
 
 
 
 
 
gdn2_memory.py ADDED
@@ -0,0 +1,112 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Reference GDN2 matrix memory for trajectory-time recurrence.
2
+
3
+ The state is FP32. Leading dimensions are independent rows or canvas cells;
4
+ only the last three dimensions, [heads, key, value], belong to the rule.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import torch
10
+ from torch import nn
11
+ from torch.nn import functional as F
12
+
13
+
14
+ class _GateProjection(nn.Module):
15
+ """GDN2 gates need channel control, not a second full-width content map."""
16
+
17
+ def __init__(self, source: int, target: int, rank: int) -> None:
18
+ super().__init__()
19
+ self.down = nn.Linear(source, rank, bias=False)
20
+ self.up = nn.Linear(rank, target, bias=False)
21
+
22
+ def forward(self, source: torch.Tensor) -> torch.Tensor:
23
+ return self.up(self.down(source))
24
+
25
+
26
+ class GDN2Memory(nn.Module):
27
+ def __init__(self, input_dim: int, heads: int, key_dim: int, value_dim: int,
28
+ *, observation_dim: int | None = None) -> None:
29
+ super().__init__()
30
+ if min(input_dim, heads, key_dim, value_dim) <= 0:
31
+ raise ValueError("GDN2 dimensions must be positive.")
32
+ self.heads, self.key_dim, self.value_dim = heads, key_dim, value_dim
33
+ observation_dim = input_dim if observation_dim is None else observation_dim
34
+ if observation_dim <= 0:
35
+ raise ValueError("GDN2 observation width must be positive.")
36
+ self.q_proj = nn.Linear(input_dim, heads * key_dim, bias=False)
37
+ self.k_proj = nn.Linear(observation_dim, heads * key_dim, bias=False)
38
+ self.v_proj = nn.Linear(observation_dim, heads * value_dim, bias=False)
39
+ self.f_proj = _GateProjection(observation_dim, heads * key_dim, min(observation_dim, key_dim))
40
+ self.b_proj = nn.Linear(observation_dim, heads * key_dim, bias=False)
41
+ self.w_proj = nn.Linear(observation_dim, heads * value_dim, bias=False)
42
+ self.g_proj = _GateProjection(input_dim, heads * value_dim, min(input_dim, value_dim))
43
+ self.o_proj = nn.Linear(heads * value_dim, input_dim, bias=False)
44
+ self.a_log = nn.Parameter(torch.zeros(heads))
45
+ self.dt_bias = nn.Parameter(torch.full((heads, key_dim), -6.906255))
46
+
47
+ def _projections(self, source: torch.Tensor):
48
+ source = F.rms_norm(source.float(), (source.shape[-1],), eps=1e-6).to(source.dtype)
49
+ shape = source.shape[:-1]
50
+ key_shape = (*shape, self.heads, self.key_dim)
51
+ value_shape = (*shape, self.heads, self.value_dim)
52
+ k = F.normalize(F.silu(self.k_proj(source).float()).reshape(key_shape), dim=-1)
53
+ v = F.silu(self.v_proj(source).float()).reshape(value_shape)
54
+ head_rate = self.a_log.float().exp().reshape(*((1,) * (source.ndim - 1)), self.heads, 1)
55
+ decay = torch.exp(-head_rate * F.softplus(
56
+ self.f_proj(source).float().reshape(key_shape) + self.dt_bias.float()))
57
+ erase = torch.sigmoid(self.b_proj(source).float().reshape(key_shape))
58
+ write = torch.sigmoid(self.w_proj(source).float().reshape(value_shape))
59
+ return k, v, decay, erase, write
60
+
61
+ def _query(self, source: torch.Tensor) -> torch.Tensor:
62
+ return F.normalize(F.silu(self.q_proj(source).float()).reshape(
63
+ *source.shape[:-1], self.heads, self.key_dim), dim=-1)
64
+
65
+ def _output(self, value: torch.Tensor, source: torch.Tensor) -> torch.Tensor:
66
+ gate = F.silu(self.g_proj(source).float().reshape(value.shape))
67
+ value = F.rms_norm(value, (self.value_dim,), eps=1e-6) * gate
68
+ return self.o_proj(value.flatten(-2).to(source.dtype))
69
+
70
+ def read(self, state: torch.Tensor, source: torch.Tensor) -> torch.Tensor:
71
+ if state.shape != (*source.shape[:-1], self.heads, self.key_dim, self.value_dim):
72
+ raise ValueError("GDN2 state and query leading dimensions differ.")
73
+ value = torch.einsum("...hk,...hkv->...hv", self._query(source), state.float())
74
+ return self._output(value, source)
75
+
76
+ def read_shared(self, state: torch.Tensor, source: torch.Tensor) -> torch.Tensor:
77
+ """Read one row matrix for all queries without a canvas-sized matrix product."""
78
+ if source.ndim != 3 or state.shape != (source.shape[0], self.heads, self.key_dim, self.value_dim):
79
+ raise ValueError("Shared GDN2 read requires [batch, canvas, width] queries.")
80
+ query = self._query(source).transpose(1, 2)
81
+ value = (query @ state.float()).transpose(1, 2)
82
+ return self._output(value, source)
83
+
84
+ @staticmethod
85
+ def _transition(state, k, v, decay, erase, write, valid):
86
+ decayed = state.float() * decay.unsqueeze(-1)
87
+ old = torch.einsum("...hk,...hkv->...hv", erase * k, decayed)
88
+ candidate = decayed + k.unsqueeze(-1) * (write * v - old).unsqueeze(-2)
89
+ return candidate if valid is None else torch.where(
90
+ valid[..., None, None, None], candidate, state.float())
91
+
92
+ def transition(self, state: torch.Tensor, source: torch.Tensor,
93
+ valid: torch.Tensor | None = None) -> torch.Tensor:
94
+ if state.shape != (*source.shape[:-1], self.heads, self.key_dim, self.value_dim):
95
+ raise ValueError("GDN2 state and observation leading dimensions differ.")
96
+ k, v, decay, erase, write = self._projections(source)
97
+ if valid is not None:
98
+ if valid.shape != source.shape[:-1]:
99
+ raise ValueError("GDN2 valid mask must match observation rows.")
100
+ return self._transition(state, k, v, decay, erase, write, valid)
101
+
102
+ def write_sequence(self, state: torch.Tensor, source: torch.Tensor,
103
+ valid: torch.Tensor) -> torch.Tensor:
104
+ """Project a commit packet once, then preserve ordered GDN2 recurrence."""
105
+ if source.ndim != 3 or valid.shape != source.shape[:2]:
106
+ raise ValueError("GDN2 sequence and mask must share [batch, length].")
107
+ if state.shape != (source.shape[0], self.heads, self.key_dim, self.value_dim):
108
+ raise ValueError("GDN2 sequence state shape differs.")
109
+ projected = self._projections(source)
110
+ for index in range(source.shape[1]):
111
+ state = self._transition(state, *(part[:, index] for part in projected), valid[:, index])
112
+ return state
gdn2_trajectory.py ADDED
@@ -0,0 +1,86 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Dual-timescale matrix state with denoise and commit lifetimes.
2
+
3
+ This module is independent of the decoder and commit policy. It provides the
4
+ reference state transition used when replacing the old history and slot paths.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ from dataclasses import dataclass
10
+
11
+ import torch
12
+ from torch import nn
13
+
14
+ from .gdn2_memory import GDN2Memory
15
+
16
+
17
+ @dataclass(frozen=True)
18
+ class GDN2TrajectoryState:
19
+ cells: torch.Tensor
20
+ row: torch.Tensor
21
+ persistent: torch.Tensor
22
+ seen: torch.Tensor
23
+
24
+
25
+ def shift(self, lengths: torch.Tensor) -> "GDN2TrajectoryState":
26
+ """Shift a non-ring canvas and zero its newly filled tail."""
27
+ batch, canvas = self.seen.shape
28
+ if lengths.shape != (batch,):
29
+ raise ValueError("Commit lengths must be per row.")
30
+ physical = torch.arange(canvas, device=self.seen.device)[None, :] + lengths[:, None]
31
+ kept = physical < canvas
32
+ selected = physical.clamp_max(canvas - 1)
33
+ cells = self.cells.gather(
34
+ 1, selected[..., None, None, None].expand_as(self.cells)
35
+ )
36
+ seen = self.seen.gather(1, selected)
37
+ return GDN2TrajectoryState(
38
+ cells.masked_fill(~kept[..., None, None, None], 0.0),
39
+ self.row, self.persistent, seen & kept,
40
+ )
41
+
42
+
43
+ class GDN2TrajectoryMemory(nn.Module):
44
+ def __init__(self, width: int, *, probes: int = 4,
45
+ working_heads: int = 16, working_key: int = 64,
46
+ working_value: int = 64, persistent_heads: int = 16,
47
+ persistent_key: int = 128, persistent_value: int = 128,
48
+ persistent_observation_dim: int | None = None) -> None:
49
+ super().__init__()
50
+ if probes <= 0:
51
+ raise ValueError("Probe count must be positive.")
52
+ self.probes = probes
53
+ self.cell = GDN2Memory(width, working_heads, working_key, working_value)
54
+ self.row = GDN2Memory(width, working_heads, working_key, working_value)
55
+ self.persistent = GDN2Memory(width, persistent_heads, persistent_key, persistent_value,
56
+ observation_dim=persistent_observation_dim)
57
+ self.probe_embed = nn.Embedding(probes, width)
58
+ nn.init.normal_(self.probe_embed.weight, std=0.02)
59
+
60
+ def read(self, state: GDN2TrajectoryState, query: torch.Tensor) -> torch.Tensor:
61
+ result = (self.cell.read(state.cells, query)
62
+ + self.row.read_shared(state.row, query)
63
+ + self.persistent.read_shared(state.persistent, query))
64
+ return torch.where(state.seen[..., None], result, torch.zeros_like(result))
65
+
66
+ def observe(self, state: GDN2TrajectoryState, observation: torch.Tensor,
67
+ live: torch.Tensor, head: torch.Tensor) -> GDN2TrajectoryState:
68
+ batch, canvas, width = observation.shape
69
+ if live.shape != (batch, canvas) or head.shape != (batch,):
70
+ raise ValueError("Observation mask and head have incorrect shapes.")
71
+ cells = self.cell.transition(state.cells, observation, live)
72
+ seen = state.seen | live
73
+ logical_idx = (head[:, None] + torch.arange(canvas, device=head.device)[None, :]) % canvas
74
+ logical = observation.gather(1, logical_idx[..., None].expand(-1, -1, width))
75
+ logical_live = live.gather(1, logical_idx)
76
+ row_state = state.row
77
+ for probe in range(self.probes):
78
+ lo = canvas * probe // self.probes
79
+ hi = canvas * (probe + 1) // self.probes
80
+ selected = logical_live[:, lo:hi]
81
+ count = selected.sum(dim=1, keepdim=True)
82
+ pooled = (logical[:, lo:hi].float() * selected[..., None]).sum(dim=1)
83
+ pooled = (pooled / count.clamp_min(1)).to(observation.dtype)
84
+ pooled = pooled + self.probe_embed.weight[probe].to(observation.dtype)
85
+ row_state = self.row.transition(row_state, pooled, count[:, 0] > 0)
86
+ return GDN2TrajectoryState(cells, row_state, state.persistent, seen)
generation_config.json CHANGED
@@ -1,24 +1,10 @@
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
  }
 
1
  {
2
+ "eos_token_id": 1,
 
 
 
 
 
 
 
3
  "jump_on_no_progress_after": 12,
4
  "max_denoising_steps": null,
 
5
  "max_ponder_steps": 64,
6
  "min_trajectory_progress": 0.005,
7
  "repetition_penalty": 1.0,
8
  "repetition_penalty_exclude_token_ids": [],
 
 
 
 
 
 
9
  "turn_end_token_id": 106
10
  }
generation_modilify_mk2.py CHANGED
@@ -1,14 +1,11 @@
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
@@ -22,36 +19,37 @@ from transformers.models.diffusion_gemma import (
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",
@@ -62,38 +60,8 @@ _PARENT_GENERATION_KEYS = frozenset({
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:
@@ -128,21 +96,7 @@ def _add_repetition_history(
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))
@@ -151,8 +105,6 @@ class ModilifyMk2GenerationConfig(DiffusionGemmaGenerationConfig):
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 = {
@@ -166,26 +118,14 @@ class ModilifyMk2GenerationConfig(DiffusionGemmaGenerationConfig):
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:
@@ -199,8 +139,6 @@ class ModilifyMk2GenerationConfig(DiffusionGemmaGenerationConfig):
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(
@@ -211,8 +149,8 @@ class ModilifyMk2GenerationConfig(DiffusionGemmaGenerationConfig):
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):
@@ -234,12 +172,6 @@ class ModilifyMk2GenerationConfig(DiffusionGemmaGenerationConfig):
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(
@@ -252,33 +184,9 @@ class ModilifyMk2GenerationConfig(DiffusionGemmaGenerationConfig):
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]:
@@ -292,20 +200,6 @@ class ModilifyMk2GenerationConfig(DiffusionGemmaGenerationConfig):
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
@@ -318,135 +212,18 @@ class ModilifyMk2GenerationOutput(ModelOutput):
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:
@@ -659,114 +436,48 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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
@@ -775,6 +486,8 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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
 
@@ -788,12 +501,9 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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]
@@ -820,18 +530,21 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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)
@@ -839,25 +552,21 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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()
@@ -868,7 +577,6 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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)
@@ -912,8 +620,6 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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:
@@ -976,6 +682,7 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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,
@@ -1016,40 +723,21 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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(
@@ -1059,17 +747,9 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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
@@ -1090,6 +770,13 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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,
@@ -1122,7 +809,6 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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, :]
@@ -1130,8 +816,6 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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, :]
@@ -1153,13 +837,11 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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,
@@ -1186,43 +868,48 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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,
@@ -1234,15 +921,12 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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
@@ -1319,21 +1003,17 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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(
@@ -1348,9 +1028,12 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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
@@ -1388,19 +1071,6 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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())
@@ -1419,8 +1089,6 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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:
@@ -1439,17 +1107,6 @@ class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
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
- ]
 
1
+ """Rolling text generation for latent-memory ModilifyMk2."""
 
 
2
 
3
  from __future__ import annotations
4
 
5
+ from collections.abc import Sequence
6
  from contextlib import contextmanager
7
  from dataclasses import dataclass, replace
8
  import math
 
9
  from typing import Any
10
 
11
  import torch
 
19
  DiffusionGemmaGenerationMixin,
20
  )
21
  from .commit_policy import (
 
22
  fused_commit_failure_rate,
23
  select_commit_lengths,
24
  )
25
+ from .configuration_modilify_mk2 import DENOISE_TEMPERATURE
26
+ from .latent_deliberation import LatentDeliberationState, infer_commit_reason
 
 
 
 
 
 
 
 
 
 
 
 
27
 
28
 
29
+ def deterministic_episode_iteration_bound(
30
+ response_lengths: torch.LongTensor,
31
+ *,
32
+ max_ponder_steps: int,
33
+ ) -> int:
34
+ """Return a safe watchdog bound without assuming canvas-sized jumps.
35
+
36
+ Every active row must advance by at least one token no later than the
37
+ configured no-progress threshold. Normal commits can also be only one
38
+ token long, so a bound based on the number of canvases is not valid.
39
+ """
40
+ if response_lengths.numel() == 0:
41
+ raise ValueError("`response_lengths` must be non-empty.")
42
+ if max_ponder_steps <= 0:
43
+ raise ValueError("`max_ponder_steps` must be positive.")
44
+ return max(1, int(response_lengths.max()) * max_ponder_steps)
45
+
46
  _PARENT_GENERATION_KEYS = frozenset({
47
  "max_new_tokens",
48
  "max_length",
49
  "return_dict_in_generate",
50
  "max_denoising_steps",
 
51
  "t_min",
52
  "t_max",
 
 
53
  "cache_implementation",
54
  "cache_config",
55
  "disable_compile",
 
60
  "_from_model_config",
61
  "transformers_version",
62
  })
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
63
  def _flatten_token_ids(*values: object) -> set[int]:
64
  """Normalize scalar and sequence token-ID configuration values."""
 
65
  token_ids: set[int] = set()
66
  for value in values:
67
  if value is None:
 
96
 
97
  class ModilifyMk2GenerationConfig(DiffusionGemmaGenerationConfig):
98
  def __init__(self, **kwargs):
 
 
 
 
 
99
  self.turn_end_token_id: int | None = kwargs.pop("turn_end_token_id", None)
 
 
 
 
 
 
 
 
 
100
  self.max_ponder_steps: int = kwargs.pop("max_ponder_steps", 64)
101
  self.jump_on_no_progress_after: int = kwargs.pop("jump_on_no_progress_after", 12)
102
  self.min_trajectory_progress: float = float(kwargs.pop("min_trajectory_progress", 0.005))
 
105
  self.repetition_penalty_exclude_token_ids: list[int] = list(
106
  dict.fromkeys(int(token_id) for token_id in excluded_token_ids or ())
107
  )
 
 
108
  kwargs.pop("t_min", None)
109
  kwargs.pop("t_max", None)
110
  parent_kwargs = {
 
118
  confidence_threshold=None,
119
  **parent_kwargs,
120
  )
121
+ self.t_min = DENOISE_TEMPERATURE
122
+ self.t_max = DENOISE_TEMPERATURE
 
123
 
124
  def update(self, defaults_only=False, allow_custom_entries=False, **kwargs):
125
+ """Apply supported overrides while keeping temperature globally fixed."""
126
 
 
 
 
 
 
127
  if "turn_end_token_id" in kwargs:
128
  self.turn_end_token_id = kwargs.pop("turn_end_token_id")
 
 
 
 
 
 
129
  if "max_ponder_steps" in kwargs:
130
  self.max_ponder_steps = kwargs.pop("max_ponder_steps")
131
  if "jump_on_no_progress_after" in kwargs:
 
139
  self.repetition_penalty_exclude_token_ids = list(
140
  dict.fromkeys(int(token_id) for token_id in excluded_token_ids or ())
141
  )
 
 
142
  kwargs.pop("t_min", None)
143
  kwargs.pop("t_max", None)
144
  unused = super().update(
 
149
  self.sampler_config = None
150
  self.stability_threshold = None
151
  self.confidence_threshold = None
152
+ self.t_min = DENOISE_TEMPERATURE
153
+ self.t_max = DENOISE_TEMPERATURE
154
  return unused
155
 
156
  def validate(self, **kwargs):
 
172
  raise ValueError("`jump_on_no_progress_after` must be a positive integer.")
173
  if not isinstance(self.min_trajectory_progress, (int, float)) or self.min_trajectory_progress < 0:
174
  raise ValueError("`min_trajectory_progress` must be a non-negative number.")
 
 
 
 
 
 
175
  if not math.isfinite(self.repetition_penalty) or self.repetition_penalty <= 0:
176
  raise ValueError("`repetition_penalty` must be a finite positive number.")
177
  if any(
 
184
 
185
  @classmethod
186
  def from_model_config(cls, model_config):
187
+ """Build the only generation field owned by the ModilifyMk2 model config."""
188
 
189
+ return cls(turn_end_token_id=model_config.turn_end_token_id)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
190
 
191
  @staticmethod
192
  def _get_default_generation_params() -> dict[str, object]:
 
200
  }
201
 
202
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
203
  @dataclass
204
  class ModilifyMk2GenerationOutput(ModelOutput):
205
  sequences: torch.LongTensor
 
212
  no_progress_steps: int | torch.LongTensor | None = None
213
  jump_count: int | torch.LongTensor | None = None
214
  forced_jump_bad_count: int | torch.LongTensor | None = None
 
 
215
  average_commit_len: float | torch.FloatTensor | None = None
216
  state_shift_count: int | torch.LongTensor | None = None
 
 
 
 
 
217
 
218
 
219
  @dataclass
220
  class ModilifyMk2RollingState:
221
  """All real iterative state; no vocabulary-sized tensor is retained."""
 
222
  canvas: torch.LongTensor
223
  confidence: torch.FloatTensor
224
  entropy: torch.FloatTensor
 
225
  latent_state: LatentDeliberationState
226
+ head: torch.LongTensor | None = None
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
227
 
228
 
229
  class NoiseCanvasSampler:
 
436
  vocab_size=self.config.text_config.vocab_size,
437
  )
438
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
439
  def _write_committed_memory(
440
  self,
441
  *,
 
442
  next_state: ModilifyMk2RollingState,
443
  working_state: torch.Tensor,
 
444
  heavy_hidden: torch.Tensor,
445
+ commit_token_ids: torch.LongTensor,
446
  commit_lengths: torch.Tensor,
447
  prefix_lengths: torch.Tensor,
448
  commit_reason: torch.Tensor,
449
+ canvas_head: torch.Tensor | None = None,
450
  ) -> ModilifyMk2RollingState:
451
  if not bool(commit_lengths.gt(0).any()):
452
  return next_state
453
+ batch, canvas = commit_token_ids.shape
454
+ if working_state.shape[:2] != (batch, canvas):
455
+ raise ValueError("Committed token IDs must match the unshifted canvas.")
456
+ max_commit = min(int(commit_lengths.max()), canvas)
457
+ if canvas_head is None:
458
+ selected_ids = commit_token_ids[:, :max_commit]
459
+ else:
460
+ head = canvas_head.to(device=commit_token_ids.device, dtype=torch.long).view(batch, 1)
461
+ physical = (head + torch.arange(max_commit, device=commit_token_ids.device)) % canvas
462
+ selected_ids = commit_token_ids.gather(1, physical)
463
+ with torch.no_grad():
464
+ committed_token_embeddings = self.model.decoder.embed_tokens(selected_ids)
465
+ memory = self.latent_deliberation.commit_write(
466
  memory=next_state.latent_state.memory_slots,
467
  working_state=working_state,
468
+
 
469
  heavy_hidden=heavy_hidden,
470
+ committed_token_embeddings=committed_token_embeddings,
471
  commit_lengths=commit_lengths,
 
472
  commit_reason=commit_reason,
473
+ canvas_head=canvas_head,
474
  )
475
  return replace(
476
  next_state,
477
+ latent_state=replace(
478
+ next_state.latent_state, memory_slots=memory,
479
+ gdn2=replace(next_state.latent_state.gdn2, persistent=memory),
480
+ ),
481
  )
482
 
483
  @staticmethod
 
486
  commit_lengths: torch.LongTensor,
487
  sampler: NoiseCanvasSampler,
488
  generators: Sequence[torch.Generator] | None = None,
489
+ remaining_lengths: torch.LongTensor | None = None,
490
+ pad_token_id: int = 0,
491
  ) -> ModilifyMk2RollingState:
492
  """Shift every rolling row by its own committed prefix length."""
493
 
 
501
  # Seeded generation deliberately advances every active request
502
  # once per denoise step, independent of the other active rows.
503
  for generator in generators:
504
+ sampler.initialize_canvas(
505
+ 1, state.canvas.device, generators=[generator]
506
+ )
 
 
 
507
  return state
508
  positions = torch.arange(canvas_length, device=state.canvas.device)[None, :]
509
  source = positions + commit_lengths[:, None]
 
530
  for row, (commit_length, generator) in enumerate(
531
  zip(commit_lengths.detach().cpu().tolist(), generators, strict=True)
532
  ):
533
+ sampled = sampler.initialize_canvas(
534
+ 1,
535
+ state.canvas.device,
536
+ generators=[generator],
537
+ )
 
 
 
 
 
538
  tail[row] = sampled[0]
539
  canvas = torch.cat((state.canvas, tail), dim=1).gather(1, source)
540
+ if remaining_lengths is not None:
541
+ if remaining_lengths.shape != (batch_size,):
542
+ raise ValueError("Remaining lengths must have shape [batch].")
543
+ canvas = canvas.masked_fill(
544
+ (positions >= (canvas_length - commit_lengths)[:, None])
545
+ & (positions >= remaining_lengths[:, None]),
546
+ pad_token_id,
547
+ )
548
  unknown_entropy = float(sampler.initial_entropy)
549
  latent = state.latent_state
550
  committed = commit_lengths.gt(0)
 
552
  memory_slots=latent.memory_slots.clone(),
553
  confidence=shift(latent.confidence),
554
  entropy=shift(latent.entropy, unknown_entropy),
 
 
 
 
555
  ponder_steps=torch.where(
556
  committed, torch.zeros_like(latent.ponder_steps), latent.ponder_steps
557
  ),
558
  stagnation_steps=torch.where(
559
  committed, torch.zeros_like(latent.stagnation_steps), latent.stagnation_steps
560
  ),
561
+ gdn2=latent.gdn2.shift(commit_lengths),
562
  )
563
  return ModilifyMk2RollingState(
564
  canvas=canvas,
565
  confidence=shift(state.confidence),
566
  entropy=shift(state.entropy, unknown_entropy),
 
567
  latent_state=shifted_latent,
568
+
569
+ head=state.head,
570
  )
571
 
572
  @torch.inference_mode()
 
577
  streamer: BaseStreamer | None = None,
578
  generation_config: ModilifyMk2GenerationConfig | None = None,
579
  logits_processor: LogitsProcessorList | None = None,
 
580
  **kwargs,
581
  ) -> ModilifyMk2GenerationOutput:
582
  request_seeds = kwargs.pop("seeds", None)
 
620
  sampling_generators = None
621
  if batch_size > 1 and streamer is not None:
622
  raise ValueError("ModilifyMk2 streamers currently support batch size 1 only.")
 
 
623
  if batch_size > 1 and past_key_values is not None:
624
  raise ValueError("Batched ModilifyMk2 generation requires a fresh KV cache.")
625
  if batch_size > 1:
 
682
  _, max_new_tokens = self._prepare_generated_length(
683
  generation_config, cached_length + input_width
684
  )
685
+
686
  max_iterations = deterministic_episode_iteration_bound(
687
  torch.tensor([max_new_tokens]),
688
  max_ponder_steps=generation_config.max_ponder_steps,
 
723
  prompt_positions = input_mask.long().cumsum(dim=-1).sub(1).clamp_min(0).to(torch.int32)
724
  logical_lengths = cache_attention_mask.long().sum(dim=-1)
725
  if input_width:
 
 
 
 
 
 
726
  past_key_values = self.model.encoder(
727
  input_ids=input_ids,
728
  attention_mask=cache_attention_mask,
729
  past_key_values=past_key_values,
730
  position_ids=prompt_positions,
 
731
  ).past_key_values
732
 
733
  sampler = self._prepare_sampler(generation_config, canvas_length)
734
  latent = LatentDeliberationState.empty(
735
  batch_size=batch_size, canvas_length=canvas_length,
 
 
 
 
 
 
 
 
736
  device=device,
 
737
  )
738
+ initial_canvas = sampler.initialize_canvas(
739
+ batch_size, device, generators=sampling_generators
740
+ )
 
 
 
741
  state = ModilifyMk2RollingState(
742
  canvas=initial_canvas,
743
  confidence=torch.zeros(
 
747
  (batch_size, canvas_length), math.log(self.config.text_config.vocab_size),
748
  device=device, dtype=torch.float32,
749
  ),
 
 
 
750
  latent_state=latent,
751
+
752
+ head=torch.zeros(batch_size, device=device, dtype=torch.long),
 
 
 
 
 
753
  )
754
  turn_end = (
755
  self.config.turn_end_token_id
 
770
  if isinstance(pad_token_id, (list, tuple)):
771
  pad_token_id = pad_token_id[0]
772
  pad_token_id = int(0 if pad_token_id is None else pad_token_id)
773
+ state = replace(
774
+ state,
775
+ canvas=state.canvas.masked_fill(
776
+ torch.arange(canvas_length, device=device)[None, :] >= max_new_tokens,
777
+ pad_token_id,
778
+ ),
779
+ )
780
  excluded_repetition_token_ids = _flatten_token_ids(
781
  generation_config.repetition_penalty_exclude_token_ids,
782
  generation_config.pad_token_id,
 
809
  jumps = torch.zeros_like(committed)
810
  forced_jump_tokens = torch.zeros_like(committed)
811
  shifts = torch.zeros_like(committed)
 
812
  stop_codes = torch.zeros_like(committed)
813
  active_rows = torch.ones(batch_size, dtype=torch.bool, device=device)
814
  canvas_positions = torch.arange(canvas_length, device=device)[None, :]
 
816
  streamer.put(input_ids.cpu())
817
 
818
  while bool(active_rows.any()):
 
 
819
  decoder_positions = (
820
  logical_lengths[:, None]
821
  + torch.arange(canvas_length, device=device)[None, :]
 
837
  input_ids=None, past_key_values=past_key_values,
838
  decoder_input_ids=state.canvas,
839
  previous_confidence=state.confidence, previous_entropy=state.entropy,
840
+ latent_state=state.latent_state,
841
+
 
842
  decoder_position_ids=decoder_positions, decoder_read_cache=True,
843
  decoder_attention_mask=decoder_attention_mask,
844
  compact_vocab=True,
 
845
  repetition_token_mask=repetition_history,
846
  repetition_penalty=repetition_penalty,
847
  sampling_generators=sampling_generators,
 
868
  output.next_latent_state,
869
  confidence=next_confidence.detach().float(),
870
  entropy=token_entropy.detach().float(),
 
 
 
 
871
  )
872
  remaining = torch.tensor(
873
  max_new_tokens, device=device, dtype=torch.long
874
  ).sub(committed)
875
  remaining_canvas = remaining[:, None].gt(canvas_positions)
876
+ next_latent = self.latent_deliberation.observe_state(
877
+ next_latent, output.heavy_hidden_state, output.working_state,
878
+ remaining_canvas, state.head,
879
  )
880
  next_state = ModilifyMk2RollingState(
881
  canvas=next_canvas, confidence=next_confidence,
882
+ entropy=token_entropy,
883
  latent_state=next_latent,
884
+
885
+ head=state.head,
 
 
 
 
 
 
886
  )
 
887
  normal_failure_rate = fused_commit_failure_rate(
888
  proposal_confidence, token_entropy,
889
+ entropy_weight=self.config.commit_entropy_weight,
890
+ confidence_power=self.config.commit_confidence_power,
891
+ top_k=getattr(self.config, "commit_top_k", None),
892
+ min_p=getattr(self.config, "commit_min_p", None),
893
+ target_confidence=getattr(self.config, "commit_target_confidence", None),
894
+ failure_budget=float(self.config.commit_failure_budget),
895
  )
896
  jump_failure_rate = fused_commit_failure_rate(
897
  greedy_confidence, token_entropy,
898
+ entropy_weight=self.config.commit_entropy_weight,
899
+ confidence_power=self.config.commit_confidence_power,
900
+ top_k=getattr(self.config, "commit_top_k", None),
901
+ min_p=getattr(self.config, "commit_min_p", None),
902
+ target_confidence=getattr(self.config, "commit_target_confidence", None),
903
+ failure_budget=float(self.config.commit_failure_budget),
904
  )
905
  previous_failure_rate = fused_commit_failure_rate(
906
  state.confidence, state.entropy,
907
+ entropy_weight=self.config.commit_entropy_weight,
908
+ confidence_power=self.config.commit_confidence_power,
909
+ top_k=getattr(self.config, "commit_top_k", None),
910
+ min_p=getattr(self.config, "commit_min_p", None),
911
+ target_confidence=getattr(self.config, "commit_target_confidence", None),
912
+ failure_budget=float(self.config.commit_failure_budget),
913
  )
914
  policy_decision = select_commit_lengths(
915
  sampled_token_ids=proposal,
 
921
  stagnation_steps=state.latent_state.stagnation_steps,
922
  active_rows=active_rows,
923
  remaining_lengths=remaining,
924
+ failure_budget=self.config.commit_failure_budget,
 
925
  stop_token_id=stop_token_ids,
926
  max_ponder_steps=generation_config.max_ponder_steps,
927
  stagnation_threshold=generation_config.jump_on_no_progress_after,
928
  min_progress=generation_config.min_trajectory_progress,
929
  )
 
 
930
  next_ponder = policy_decision.ponder_steps
931
  next_stagnation = policy_decision.stagnation_steps
932
  commit_lengths = policy_decision.commit_lengths
 
1003
  ).past_key_values
1004
  if streamer is not None:
1005
  streamer.put(committed_block.cpu())
 
 
1006
  committed_rows = commit_lengths.gt(0)
1007
  shifts += committed_rows.long()
1008
  if (
1009
+ output.working_state is None
 
1010
  ):
1011
  raise RuntimeError("Forward did not return working trajectory features.")
1012
  next_state = self._write_committed_memory(
 
1013
  next_state=next_state,
1014
  working_state=output.working_state,
 
1015
  heavy_hidden=output.heavy_hidden_state,
1016
+ commit_token_ids=commit_token_ids,
1017
  commit_lengths=commit_lengths,
1018
  prefix_lengths=logical_lengths,
1019
  commit_reason=infer_commit_reason(
 
1028
  commit_lengths,
1029
  sampler,
1030
  generators=sampling_generators,
1031
+ remaining_lengths=remaining - commit_lengths,
1032
+ pad_token_id=pad_token_id,
1033
  )
 
1034
  state = shifted
1035
+ committed += commit_lengths
1036
+ logical_lengths += commit_lengths
1037
 
1038
  turn_hits = (
1039
  commit_token_ids.eq(turn_end) & commit_positions
 
1071
  torch.full_like(stop_codes, 5),
1072
  stop_codes,
1073
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
1074
  active_rows = stop_codes.eq(0)
1075
 
1076
  output_width = int(committed.max())
 
1089
  )
1090
  tokens_per_forward = committed.float() / denoise_steps.clamp_min(1).float()
1091
  average_commit_len = committed.float() / shifts.clamp_min(1).float()
 
 
1092
 
1093
  def scalar_or_tensor(value: torch.Tensor, *, floating: bool = False):
1094
  if batch_size > 1:
 
1107
  no_progress_steps=scalar_or_tensor(state.latent_state.stagnation_steps),
1108
  jump_count=scalar_or_tensor(jumps),
1109
  forced_jump_bad_count=scalar_or_tensor(forced_jump_tokens),
 
 
1110
  average_commit_len=scalar_or_tensor(average_commit_len, floating=True),
1111
  state_shift_count=scalar_or_tensor(shifts),
 
 
1112
  )
 
 
 
 
 
 
 
latent_deliberation.py CHANGED
@@ -1,45 +1,31 @@
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(
@@ -93,197 +79,15 @@ def _fp32_scaled_dot_product_attention(
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(
@@ -291,84 +95,35 @@ class LatentDeliberationState:
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(
@@ -377,13 +132,9 @@ def advance_trajectory_clocks(
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
@@ -411,85 +162,12 @@ def should_force_trajectory_jump(
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):
@@ -500,7 +178,7 @@ class _RMSNorm(nn.Module):
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):
@@ -514,216 +192,6 @@ class _SwiGLU(nn.Module):
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
 
@@ -767,86 +235,8 @@ class _RankAttention(nn.Module):
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,
@@ -893,20 +283,7 @@ class DecoderMemoryBus(nn.Module):
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,
@@ -915,8 +292,6 @@ class DecoderMemoryBus(nn.Module):
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.")
@@ -931,9 +306,23 @@ class DecoderMemoryBus(nn.Module):
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)
@@ -945,18 +334,28 @@ class DecoderMemoryBus(nn.Module):
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)
@@ -968,1059 +367,6 @@ class DecoderMemoryBus(nn.Module):
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:
@@ -2028,12 +374,12 @@ def slice_latent_state(
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
 
@@ -2045,12 +391,14 @@ def cat_latent_states(states: Sequence[LatentDeliberationState]) -> LatentDelibe
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
 
@@ -2060,7 +408,6 @@ def infer_commit_reason(
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
 
@@ -2072,7 +419,7 @@ def infer_commit_reason(
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:
@@ -2097,38 +444,181 @@ def infer_commit_reason(
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
- ]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Schema25 Torch GDN2 trajectory state, processor, and memory readers."""
 
 
 
 
 
 
 
2
 
3
  from __future__ import annotations
4
 
5
  from collections.abc import Sequence
6
+ from dataclasses import dataclass, replace
7
+ import weakref
8
  import math
9
 
10
  import torch
11
  from torch import nn
12
  from torch.nn import functional as F
13
 
14
+ from .gdn2_trajectory import GDN2TrajectoryMemory, GDN2TrajectoryState
15
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
16
 
17
  COMMIT_REASON_NONE = 0
18
  COMMIT_REASON_NORMAL = 1
19
  COMMIT_REASON_FORCED_JUMP = 2
20
  COMMIT_REASON_TERMINAL = 3
21
+
22
+
23
+ def _sdpa_mask_value(dtype: torch.dtype) -> float:
24
+ """Additive SDPA mask that stays finite on MPS fp16/bf16."""
25
+
26
+ if dtype in (torch.float16, torch.bfloat16):
27
+ return -1.0e4
28
+ return -1.0e9
29
 
30
 
31
  def _fp32_scaled_dot_product_attention(
 
79
  return output.to(dtype=output_dtype)
80
 
81
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
82
  @dataclass
83
  class LatentDeliberationState:
84
  """Persistent slots plus per-canvas trajectory clocks. No token latents."""
 
85
  memory_slots: torch.Tensor
86
  confidence: torch.Tensor
87
  entropy: torch.Tensor
 
 
 
 
88
  ponder_steps: torch.Tensor
89
  stagnation_steps: torch.Tensor
90
+ gdn2: GDN2TrajectoryState
91
 
92
  @classmethod
93
  def empty(
 
95
  *,
96
  batch_size: int,
97
  canvas_length: int,
 
 
98
  device: torch.device,
 
99
  ) -> "LatentDeliberationState":
100
+ persistent = torch.zeros(batch_size, 16, 128, 128, device=device, dtype=torch.float32)
101
  return cls(
102
+ memory_slots=persistent,
 
 
103
  confidence=torch.zeros(
104
  batch_size, canvas_length, device=device, dtype=torch.float32
105
  ),
106
  entropy=torch.zeros(
107
  batch_size, canvas_length, device=device, dtype=torch.float32
108
  ),
 
 
 
 
 
 
 
 
 
 
 
 
109
  ponder_steps=torch.zeros(batch_size, device=device, dtype=torch.int32),
110
  stagnation_steps=torch.zeros(batch_size, device=device, dtype=torch.int32),
111
+ gdn2=GDN2TrajectoryState(
112
+ cells=torch.zeros(batch_size, canvas_length, 16, 64, 64,
113
+ device=device, dtype=torch.float32),
114
+ row=torch.zeros(batch_size, 16, 64, 64,
115
+ device=device, dtype=torch.float32),
116
+ persistent=persistent,
117
+ seen=torch.zeros(batch_size, canvas_length,
118
+ device=device, dtype=torch.bool),
119
+ ),
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
120
  )
121
 
122
 
123
  @dataclass
124
  class LatentProcessorOutput:
125
  context: torch.Tensor
 
126
  state: LatentDeliberationState
 
127
 
128
 
129
  def advance_trajectory_clocks(
 
132
  *,
133
  commit_lengths: torch.LongTensor,
134
  active_rows: torch.BoolTensor,
 
 
135
  ) -> tuple[torch.IntTensor, torch.IntTensor]:
136
  """Advance useful-ponder and stagnation clocks for each row."""
137
 
 
 
138
  if not (
139
  ponder_steps.shape == stagnation_steps.shape == commit_lengths.shape
140
  == active_rows.shape
 
162
  ponder_steps: torch.Tensor | None = None,
163
  max_ponder_steps: int | None = None,
164
  ) -> torch.BoolTensor:
165
+ jump = stagnation_steps.ge(stagnation_threshold)
166
+ if progress_scores is not None:
167
+ jump = jump & progress_scores.le(float(min_progress))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
168
  if ponder_steps is not None and max_ponder_steps is not None and max_ponder_steps > 0:
169
+ jump = jump | ponder_steps.ge(max_ponder_steps)
170
+ return jump.to(torch.bool)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
171
 
172
 
173
  class _RMSNorm(nn.Module):
 
178
 
179
  def forward(self, hidden: torch.Tensor) -> torch.Tensor:
180
  rms = hidden.float().square().mean(dim=-1, keepdim=True).add(self.eps).rsqrt()
181
+ return (hidden.float() * rms * self.weight.float()).to(dtype=hidden.dtype)
182
 
183
 
184
  class _SwiGLU(nn.Module):
 
192
  return self.down(F.silu(self.gate(hidden)) * self.up(hidden))
193
 
194
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
195
  class _RankAttention(nn.Module):
196
  """Sequence attention in a rank-``kv_rank`` subspace, then map back to ``dim``."""
197
 
 
235
  return self.o_proj(context)
236
 
237
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
238
  class DecoderMemoryBus(nn.Module):
239
+ """Read memory through a per-head gated residual."""
 
 
 
 
 
240
 
241
  def __init__(
242
  self,
 
283
  span = max(2 * max_relative_span - 1, 1)
284
  self.rel_bias = nn.Parameter(torch.zeros(max(num_heads, 1), span))
285
  self.max_relative_span = max_relative_span
 
 
 
 
 
 
 
286
 
 
 
 
 
 
 
287
 
288
  def prepare_kv(
289
  self,
 
292
  ) -> tuple[torch.Tensor, torch.Tensor] | None:
293
  if self.num_readers <= 0:
294
  return None
 
 
295
  if self.address_with_identity:
296
  if slot_identity is None:
297
  raise ValueError("Persistent bus requires slot identity on keys.")
 
306
  values = self.v_proj(mapped_values).view(batch, slots, heads, head_dim).transpose(1, 2)
307
  return keys, values
308
 
309
+ def _relative_mask(
310
+ self,
311
+ queries: int,
312
+ keys: int,
313
+ device: torch.device,
314
+ dtype: torch.dtype,
315
+ positions: torch.Tensor | None = None,
316
+ ) -> torch.Tensor | None:
317
  if not self.relative_bias:
318
  return None
319
+ if positions is not None:
320
+ if positions.shape[1] != queries or queries != keys:
321
+ raise ValueError("Working memory positions must match both canvas axes.")
322
+ relative = (
323
+ positions[:, :, None] - positions[:, None, :] + (keys - 1)
324
+ ).clamp(0, self.rel_bias.shape[1] - 1)
325
+ return self.rel_bias[:, relative].permute(1, 0, 2, 3).to(dtype=dtype)
326
  q = torch.arange(queries, device=device)
327
  k = torch.arange(keys, device=device)
328
  rel = (q[:, None] - k[None, :] + (keys - 1)).clamp(0, self.rel_bias.shape[1] - 1)
 
334
  reader_index: int,
335
  keys: torch.Tensor,
336
  values: torch.Tensor,
337
+ positions: torch.Tensor | None = None,
338
+ *, key_seen: torch.Tensor | None = None,
339
  ) -> torch.Tensor:
340
  batch, canvas, _dim = hidden.shape
341
  heads = self.num_heads
342
  head_dim = self.head_dim
343
  query = self.q_proj[reader_index](self.q_norm(hidden))
344
  query = query.view(batch, canvas, heads, head_dim).transpose(1, 2)
345
+ bias = self._relative_mask(
346
+ canvas, keys.shape[2], hidden.device, query.dtype, positions
347
+ )
348
+ if bias is not None and bias.ndim == 3:
349
  bias = bias.unsqueeze(0)
350
+ if key_seen is not None:
351
+ key_mask = torch.zeros((batch, 1, 1, keys.shape[2]), device=hidden.device, dtype=torch.float32)
352
+ key_mask = key_mask.masked_fill(~key_seen[:, None, None, :], -1e9)
353
+ bias = key_mask if bias is None else bias.float() + key_mask
354
  context = _fp32_scaled_dot_product_attention(
355
  query, keys, values, attn_mask=bias
356
  )
357
+ if key_seen is not None:
358
+ context = torch.where(key_seen.any(-1)[:, None, None, None], context, torch.zeros_like(context))
359
  scale = torch.tanh(self.alpha[reader_index]).to(dtype=hidden.dtype).view(1, heads, 1, 1)
360
  context = context * scale
361
  context = context.transpose(1, 2).reshape(batch, canvas, self.kv_rank)
 
367
  self.rel_bias.zero_()
368
 
369
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
370
  def slice_latent_state(
371
  state: LatentDeliberationState, rows: slice | torch.Tensor
372
  ) -> LatentDeliberationState:
 
374
  memory_slots=state.memory_slots[rows],
375
  confidence=state.confidence[rows],
376
  entropy=state.entropy[rows],
 
 
 
 
377
  ponder_steps=state.ponder_steps[rows],
378
  stagnation_steps=state.stagnation_steps[rows],
379
+ gdn2=GDN2TrajectoryState(
380
+ cells=state.gdn2.cells[rows], row=state.gdn2.row[rows],
381
+ persistent=state.gdn2.persistent[rows], seen=state.gdn2.seen[rows],
382
+ ),
383
  )
384
 
385
 
 
391
  memory_slots=cat("memory_slots"),
392
  confidence=cat("confidence"),
393
  entropy=cat("entropy"),
 
 
 
 
394
  ponder_steps=cat("ponder_steps"),
395
  stagnation_steps=cat("stagnation_steps"),
396
+ gdn2=GDN2TrajectoryState(
397
+ cells=torch.cat([state.gdn2.cells for state in states], dim=0),
398
+ row=torch.cat([state.gdn2.row for state in states], dim=0),
399
+ persistent=torch.cat([state.gdn2.persistent for state in states], dim=0),
400
+ seen=torch.cat([state.gdn2.seen for state in states], dim=0),
401
+ ),
402
  )
403
 
404
 
 
408
  jump_rows: torch.Tensor | None = None,
409
  commit_token_ids: torch.Tensor | None = None,
410
  terminal_token_ids: Sequence[int] = (),
 
411
  ) -> torch.Tensor:
412
  """Return per-row commit-reason codes. No hard skip; writer sees the label."""
413
 
 
419
  )
420
  committed = commit_lengths.gt(0)
421
  default = (
422
+ COMMIT_REASON_NORMAL
423
  )
424
  reasons = torch.where(committed, torch.full_like(reasons, default), reasons)
425
  if jump_rows is not None:
 
444
  return reasons
445
 
446
 
447
+ class _CanvasBlock(nn.Module):
448
+ def __init__(self, width: int, heads: int, rank: int, ffn: int,
449
+ window: int, global_attention: bool) -> None:
450
+ super().__init__()
451
+ self.norm = _RMSNorm(width)
452
+ self.attn = _RankAttention(width, heads, rank)
453
+ self.ff_norm = _RMSNorm(width)
454
+ self.ff = _SwiGLU(width, ffn)
455
+ self.window = window
456
+ self.global_attention = global_attention
457
+
458
+ def forward(self, hidden: torch.Tensor, seen: torch.Tensor,
459
+ offsets: torch.Tensor) -> torch.Tensor:
460
+ allowed = seen[:, None, :].expand(-1, hidden.shape[1], -1)
461
+ if not self.global_attention:
462
+ allowed = allowed & ((offsets[:, :, None] - offsets[:, None, :]).abs() < self.window)
463
+ additive = torch.zeros(allowed.shape, device=hidden.device, dtype=torch.float32)
464
+ additive = additive.masked_fill(~allowed, -1e9)
465
+ normed = self.norm(hidden)
466
+ hidden = hidden + self.attn(normed, normed, normed, attn_mask=additive)
467
+ hidden = hidden + self.ff(self.ff_norm(hidden))
468
+ return torch.where(seen[..., None], hidden, torch.zeros_like(hidden))
469
+
470
+
471
+ class _WorkingBus(DecoderMemoryBus):
472
+ def prepare_kv(self, memory: tuple[torch.Tensor, torch.Tensor],
473
+ slot_identity: torch.Tensor | None = None):
474
+ del slot_identity
475
+ working, seen = memory
476
+ pair = super().prepare_kv(working)
477
+ return None if pair is None else (*pair, seen)
478
+
479
+ def read(self, hidden: torch.Tensor, reader_index: int,
480
+ keys: torch.Tensor, values: torch.Tensor, seen: torch.Tensor,
481
+ positions: torch.Tensor | None = None) -> torch.Tensor:
482
+ written = super().read(hidden, reader_index, keys, values, positions, key_seen=seen)
483
+ return torch.where(seen[..., None], written, hidden)
484
+
485
+
486
+ class _PersistentBus(nn.Module):
487
+ def __init__(self, memory: nn.Module, readers: int) -> None:
488
+ super().__init__()
489
+ object.__setattr__(self, "_memory_ref", weakref.ref(memory))
490
+ self.num_readers = readers
491
+ self.alpha = nn.Parameter(torch.zeros(max(readers, 1), 1))
492
+
493
+
494
+ def reset_identity_parameters(self) -> None:
495
+ with torch.no_grad():
496
+ self.alpha.zero_()
497
+
498
+ def prepare_kv(self, memory: torch.Tensor,
499
+ seen: torch.Tensor | None = None):
500
+ if self.num_readers <= 0:
501
+ return None
502
+ return memory, seen
503
+
504
+ def read(self, hidden: torch.Tensor, reader_index: int,
505
+ memory: torch.Tensor, seen: torch.Tensor | None) -> torch.Tensor:
506
+ delta = self._memory_ref().read_shared(memory, hidden)
507
+ if seen is not None:
508
+ delta = delta * seen[..., None].to(delta.dtype)
509
+ return hidden + torch.tanh(self.alpha[reader_index]).to(hidden.dtype) * delta
510
+
511
+
512
+ class LatentDeliberationTransformer(nn.Module):
513
+ def __init__(self, *, hidden_size: int,
514
+ latent_dim: int = 2816, ffn_dim: int = 7168,
515
+ num_layers: int = 4,
516
+ num_heads: int = 16, local_attention_window: int = 128,
517
+ tape_probes: int = 4,
518
+
519
+ history_kv_rank: int = 1024, num_memory_readers: int = 0,
520
+ num_working_readers: int | None = None,
521
+ num_persistent_readers: int | None = None,
522
+ working_last_block_global: bool = True,
523
+ commit_sequence_dim: int | None = None,
524
+ max_canvas_length: int = 256) -> None:
525
+ super().__init__()
526
+ if hidden_size != latent_dim:
527
+ raise ValueError("Schema25 GDN2 requires hidden_size == latent_dim.")
528
+ self.hidden_size = hidden_size
529
+ self.latent_dim = latent_dim
530
+ self.tape_probes = tape_probes
531
+ self.packet_dim = int(commit_sequence_dim or history_kv_rank)
532
+ if self.packet_dim % num_heads:
533
+ raise ValueError("Commit packet rank must be divisible by attention heads.")
534
+ self.trajectory = GDN2TrajectoryMemory(
535
+ hidden_size, probes=tape_probes, persistent_observation_dim=self.packet_dim,
536
+ )
537
+ self.blocks = nn.ModuleList([
538
+ _CanvasBlock(hidden_size, num_heads, history_kv_rank, ffn_dim,
539
+ local_attention_window,
540
+ bool(working_last_block_global and i == num_layers - 1))
541
+ for i in range(num_layers)
542
+ ])
543
+ self.output_norm = _RMSNorm(hidden_size)
544
+ readers = num_memory_readers if num_working_readers is None else num_working_readers
545
+ persistent_readers = (
546
+ num_memory_readers if num_persistent_readers is None else num_persistent_readers
547
+ )
548
+ bus_rank = history_kv_rank if history_kv_rank % num_heads == 0 else num_heads
549
+ self.working_memory_bus = _WorkingBus(
550
+ hidden_size, num_heads, readers, hidden_size, bus_rank,
551
+ relative_bias=True, max_relative_span=max_canvas_length,
552
+ )
553
+ self.persistent_memory_bus = _PersistentBus(
554
+ self.trajectory.persistent, persistent_readers,
555
+ )
556
+ self.experience_in = nn.Linear(hidden_size * 3, self.packet_dim, bias=False)
557
+ self.reason_embed = nn.Embedding(6, self.packet_dim)
558
+
559
+ def reset_identity_parameters(self) -> None:
560
+ self.working_memory_bus.reset_identity_parameters()
561
+ self.persistent_memory_bus.reset_identity_parameters()
562
+
563
+
564
+ def forward(self, *, token_embeddings: torch.Tensor, confidence: torch.Tensor,
565
+ entropy: torch.Tensor, state: LatentDeliberationState,
566
+ canvas_head: torch.Tensor | None = None) -> LatentProcessorOutput:
567
+ batch, canvas, width = token_embeddings.shape
568
+ if state.memory_slots.shape != state.gdn2.persistent.shape:
569
+ raise ValueError("Persistent GDN2 state shape differs from memory slots.")
570
+ memory = replace(state.gdn2, persistent=state.memory_slots)
571
+ hidden = self.trajectory.read(memory, token_embeddings)
572
+ offsets = torch.arange(canvas, device=token_embeddings.device)[None, :].expand(batch, -1)
573
+ if canvas_head is not None:
574
+ offsets = (offsets - canvas_head[:, None]) % canvas
575
+ for block in self.blocks:
576
+ hidden = block(hidden, memory.seen, offsets)
577
+ hidden = self.output_norm(hidden) * memory.seen[..., None].to(hidden.dtype)
578
+ next_state = replace(state, confidence=confidence.float(), entropy=entropy.float(),
579
+ gdn2=memory)
580
+ return LatentProcessorOutput(hidden, next_state)
581
+
582
+ def observe_state(self, state: LatentDeliberationState,
583
+ heavy: torch.Tensor, working: torch.Tensor,
584
+ live: torch.Tensor, head: torch.Tensor) -> LatentDeliberationState:
585
+ source = heavy.detach() + working
586
+ updated = self.trajectory.observe(state.gdn2, source, live, head)
587
+ return replace(state, gdn2=updated)
588
+
589
+
590
+ def commit_write(self, *, memory: torch.Tensor, working_state: torch.Tensor,
591
+ heavy_hidden: torch.Tensor,
592
+ committed_token_embeddings: torch.Tensor,
593
+ commit_lengths: torch.Tensor,
594
+ commit_reason: torch.Tensor | None = None,
595
+ canvas_head: torch.Tensor | None = None):
596
+ batch, canvas, width = working_state.shape
597
+ count = int(commit_lengths.max().item())
598
+ if count <= 0:
599
+ return memory
600
+ index = torch.arange(count, device=working_state.device)[None, :].expand(batch, -1)
601
+ if canvas_head is not None:
602
+ index = (index + canvas_head[:, None]) % canvas
603
+ selected_working = working_state.gather(
604
+ 1, index[..., None].expand(-1, -1, width)
605
+ )
606
+ selected_heavy = heavy_hidden.detach().gather(
607
+ 1, index[..., None].expand(-1, -1, width)
608
+ )
609
+ if committed_token_embeddings.shape != selected_working.shape:
610
+ raise ValueError("Committed embeddings do not match the prefix.")
611
+ # Canvas processing already supplies bidirectional spatial context.
612
+ # Separate normalized role channels feed the ordered GDN2 writer directly.
613
+ # Unit-floor normalization keeps a zero Working state at zero without
614
+ # amplifying its derivative by 1/sqrt(eps) on the first denoise.
615
+ roles = tuple(F.rms_norm(value.float(), (width,), eps=1.0).to(value.dtype)
616
+ for value in (selected_heavy, selected_working,
617
+ committed_token_embeddings.detach()))
618
+ packet = self.experience_in(torch.cat(roles, dim=-1))
619
+ reason = torch.zeros(batch, device=packet.device, dtype=torch.long) if commit_reason is None else commit_reason.long()
620
+ reason = torch.where((reason == 1) | (reason == 2), 5, reason).clamp(0, 5)
621
+ packet = packet + self.reason_embed(reason)[:, None]
622
+ valid = torch.arange(count, device=packet.device)[None, :] < commit_lengths[:, None]
623
+ written = self.trajectory.persistent.write_sequence(memory, packet, valid)
624
+ return written
lora.py ADDED
@@ -0,0 +1,96 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Unfused dense and expert inference adapters."""
2
+ from __future__ import annotations
3
+ from dataclasses import dataclass
4
+ import math
5
+ import torch
6
+ from torch import nn
7
+ from transformers.models.diffusion_gemma.modeling_diffusion_gemma import DiffusionGemmaTextExperts
8
+ from .mps_ops import attach_expert_lora
9
+
10
+ @dataclass
11
+ class LoRAConfig:
12
+ r: int
13
+ alpha: int
14
+ target_modules: tuple[str, ...]
15
+ expert_r: int = 8
16
+ expert_alpha: int = 8
17
+
18
+
19
+ class LoRALinear(nn.Module):
20
+ def __init__(self, base: nn.Linear, config: LoRAConfig):
21
+ super().__init__()
22
+ if config.r <= 0:
23
+ raise ValueError("LoRA rank must be positive.")
24
+ self.base = base
25
+ self.scaling = config.alpha / config.r
26
+ self.lora_a = nn.Parameter(
27
+ torch.empty(config.r, base.in_features, device=base.weight.device, dtype=base.weight.dtype)
28
+ )
29
+ self.lora_b = nn.Parameter(
30
+ torch.zeros(base.out_features, config.r, device=base.weight.device, dtype=base.weight.dtype)
31
+ )
32
+ nn.init.kaiming_uniform_(self.lora_a, a=math.sqrt(5))
33
+ self.base.requires_grad_(False)
34
+
35
+ def forward(self, inputs: torch.Tensor) -> torch.Tensor:
36
+ update = torch.nn.functional.linear(
37
+ torch.nn.functional.linear(inputs, self.lora_a), self.lora_b,
38
+ )
39
+ return self.base(inputs) + self.scaling * update
40
+
41
+ def get_parent_module(root: nn.Module, module_name: str) -> tuple[nn.Module, str]:
42
+ parent = root
43
+ parts = module_name.split(".")
44
+ for part in parts[:-1]:
45
+ parent = getattr(parent, part)
46
+ return parent, parts[-1]
47
+
48
+ def inject_lora(model: nn.Module, config: LoRAConfig) -> int:
49
+ targets = set(config.target_modules)
50
+ replacements = []
51
+ for name, module in model.named_modules():
52
+ if not isinstance(module, nn.Linear):
53
+ continue
54
+ if ".router." in name:
55
+ continue
56
+ in_trunk = ".layers." in name
57
+ if name.rsplit(".", 1)[-1] not in targets or not in_trunk:
58
+ continue
59
+ replacements.append((name, module))
60
+ for name, module in replacements:
61
+ parent, child = get_parent_module(model, name)
62
+ setattr(parent, child, LoRALinear(module, config))
63
+ expert_modules = [
64
+ module for module in model.modules()
65
+ if isinstance(module, DiffusionGemmaTextExperts)
66
+ ]
67
+ for module in expert_modules:
68
+ attach_expert_lora(
69
+ module,
70
+ rank=config.expert_r,
71
+ alpha=config.expert_alpha,
72
+ )
73
+ share_tied_lora_parameters(model)
74
+ if not replacements and not expert_modules:
75
+ raise RuntimeError("No LoRA target modules were found.")
76
+ return len(replacements) + len(expert_modules)
77
+
78
+ def share_tied_lora_parameters(model: nn.Module) -> None:
79
+ """Use one adapter for the encoder/decoder pair whose base weights are tied."""
80
+ trunk = getattr(model, "model", model)
81
+ encoder = trunk.encoder.language_model
82
+ decoder = trunk.decoder
83
+ expert_names = ("lora_gate_up_a", "lora_gate_up_b", "lora_down_a", "lora_down_b")
84
+ for name, decoder_module in decoder.named_modules():
85
+ if not name:
86
+ continue
87
+ try:
88
+ encoder_module = encoder.get_submodule(name)
89
+ except AttributeError:
90
+ continue
91
+ if isinstance(decoder_module, LoRALinear) and isinstance(encoder_module, LoRALinear):
92
+ encoder_module.lora_a = decoder_module.lora_a
93
+ encoder_module.lora_b = decoder_module.lora_b
94
+ for parameter_name in expert_names:
95
+ if hasattr(decoder_module, parameter_name) and hasattr(encoder_module, parameter_name):
96
+ setattr(encoder_module, parameter_name, getattr(decoder_module, parameter_name))
model-00001-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:f68ecb8a12a1841220beb7abf37d9c01997898a7c35d17421e8fc55c87885378
3
+ size 776579538
model-00002-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:a7845246469006a6104991410f98103f3553a2f856b03e15b1b9174dc0af7e9b
3
+ size 1991115264
model-00003-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:4b7a0562d798b0bca197081b96998cfb5ca7bc78b80f851bd341bed7d5f7582d
3
+ size 1645281386
model-00004-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:75cb16be33e56252a1e60f4ae509aa9fc68d88c384c6a3ffcbc151717f0a0ad9
3
+ size 1645281386
model-00005-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:26195c1830ed972490d87b8d2925acf68cb8055150983e1b18b18115734c2e2a
3
+ size 1645281426
model-00006-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:5a379d44ba91bdc640e93d1ba938b07fdb2ad8d78befc3d7ce72fb2c92a45103
3
+ size 1674191634
model-00007-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c7cf0185a888b2760eb02d88e729886b986dfb2c983bcd12eb4755f31ffa8142
3
+ size 1645281426
model-00008-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:3219fcc299b5bc7274eae56aa764285a6816f693560ebaad4d3bc9d844dfaf6e
3
+ size 1645281426
model-00009-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e013c6c6150825fd566767c9371043d3803bfa18062ee08512c3e8f948f8d0bb
3
+ size 1645281426
model-00010-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:ec236ad44475058d30ae259f426d142f0585d11025febd01d3e6acea108d3518
3
+ size 1645281426
model-00011-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:cc005df4bc904483b63d0e0708681ca06f70ebe998c189de76a04ebd055bf769
3
+ size 1645281426
model-00012-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:f3ed5d1fc09e199e1484f4814ac8aa9191599e316c3be0cf497f0f0234b0bbfb
3
+ size 1674191634
model-00013-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:2962e644f6573492005e9fc6e27aff0c8b8f70817e9d057283e47f0aa0c88a79
3
+ size 1645281426
model-00014-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:90aff74d0d96d117c7bab8b1db3dcad7b098c06d5a279fef25dfb9d324dd1040
3
+ size 1645281418
model-00015-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:851680d90aa6e88311695b8595c966c82e12037a8d465be80b9c0b2776ee4de3
3
+ size 1645281386
model-00016-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:87f1b7dafd98f6a4b84aaf4ad0b12f33c162fc51d0586a1ffbc67181869b08c9
3
+ size 1645281426
model-00017-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:30eaee452145152d9d386ecd66b3725ec43f0fbf98c2edb977cd14df3f2eddb7
3
+ size 1645281426
model-00018-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e23f04834cf29b0a9be22469df57214b5c57eb3eb2e7daf29b4d9f14aad0b0f1
3
+ size 1645281426
model-00019-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:828884749f80e06e9e0dfdedd59fb7b11b03fc21a1c275865ae00ee59089957b
3
+ size 1674191634
model-00020-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:de812a62a470609ae9c905d21ef956605c40961e324f2c9f5008c6998aa9c281
3
+ size 1645281426
model-00021-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:55a980eeaf71a20471c0c8d947736d19592fbd727994088bd5a1ab895f757671
3
+ size 1645281426
model-00022-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:0bef98fc4713c3064afb0b0a2c2f078a475f95f846848ab9f08527c9ba82fff8
3
+ size 1645281426
model-00023-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:7dc71ce1fefedc8db3d055a316978d851e252cf98700f60153fdb496e6534c02
3
+ size 1645281426
model-00024-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:cfe3f829c42df7fd42e722c4569eb81441f43f55e10179455ffe4b774c7b0193
3
+ size 1645281426
model-00025-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e234c850e1c9f3d8cb9647252707beae73d9f6f9fb89ae3d9bf04d8a619b52cb
3
+ size 1674191634
model-00026-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:46afc6cb1b4927ebeb8a3791219f4d7f2d87a345374809c30a08958c3a3acd17
3
+ size 1645281386
model-00027-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:64c3778eb609039f2c753af4c056f01910c1bc7234d1ad158e41b560e70ce949
3
+ size 1645281386
model-00028-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d6b688f675bd7a78a55409cf4ce0493f27c504a6a94de7bdda3051496f9602a0
3
+ size 1674191602
model-00029-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:0987505395c67fed81d79850a4bb108c60993ed2078c3645323cfae9c1dea2c6
3
+ size 1645281386
model-00030-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:f26c9e6176125c92d9ef39fd237b4cd06132f93ab479545cf50a4f4c16353cb6
3
+ size 1645281386
model-00031-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:fed5d3814fb1b292f06de9690224fea1541feffbc9aeab2199093b75e1e231ef
3
+ size 1645281386
model-00032-of-00032.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:62f60aff697f4befac196f2f790f2597f0a7fe8e19530dd09a3915d518d93f86
3
+ size 1166261198
model.safetensors.index.json CHANGED
The diff for this file is too large to render. See raw diff
 
modeling_modilify_mk2.py CHANGED
@@ -1,11 +1,9 @@
1
- # Copyright 2026 Modilify
2
- # SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.0
3
- """Multimodal Modilify Mk2 model with recurrent latent deliberation."""
4
 
5
  from __future__ import annotations
6
 
7
  from collections.abc import Sequence
8
- from dataclasses import dataclass, replace
9
  from typing import Any
10
 
11
  import torch
@@ -13,7 +11,6 @@ from torch import nn
13
  from torch.nn import functional as F
14
  from transformers.cache_utils import Cache
15
  from transformers.modeling_outputs import BaseModelOutputWithPast
16
- from transformers.utils import ModelOutput
17
 
18
  from transformers.masking_utils import (
19
  ALL_MASK_ATTENTION_FUNCTIONS,
@@ -21,7 +18,7 @@ from transformers.masking_utils import (
21
  )
22
  from transformers.models.diffusion_gemma import (
23
  DiffusionGemmaDecoderModel,
24
- DiffusionGemmaEncoderModel,
25
  DiffusionGemmaPreTrainedModel,
26
  )
27
  from transformers.models.diffusion_gemma.modeling_diffusion_gemma import (
@@ -29,61 +26,58 @@ from transformers.models.diffusion_gemma.modeling_diffusion_gemma import (
29
  DiffusionGemmaTextRouter,
30
  )
31
 
32
- from .mps_ops import mps_segmented_experts_forward as _mps_segmented_experts_forward # noqa: F401
33
- from .configuration_modilify_mk2 import ModilifyMk2Config
 
34
  from .generation_modilify_mk2 import ModilifyMk2GenerationConfig, ModilifyMk2GenerationMixin
35
- from .latent_deliberation import (
36
- LatentDeliberationState,
37
- LatentDeliberationTransformer,
38
- LatentProcessorOutput,
39
- TrajectoryHistory,
40
- TrajectoryTape,
41
- empty_trajectory_tape,
42
- )
43
  from .vocab_ops import chunked_vocab_statistics
44
 
45
 
46
- @dataclass
47
- class ModilifyMk2DecoderOutput(BaseModelOutputWithPast):
48
- token_embeddings: torch.FloatTensor | None = None
49
- latent_residual_diagnostics: dict[str, torch.Tensor] | None = None
 
 
 
 
 
 
 
50
 
51
 
52
  @dataclass
53
  class ModilifyMk2ModelOutput(BaseModelOutputWithPast):
54
- token_embeddings: torch.FloatTensor | None = None
55
  encoder_last_hidden_state: torch.FloatTensor | None = None
56
- latent_residual_diagnostics: dict[str, torch.Tensor] | None = None
57
 
58
 
59
  @dataclass
60
- class ModilifyMk2BlockDiffusionOutput(ModelOutput):
61
- """Inference output used by the rolling diffusion generator."""
62
-
63
- logits: torch.FloatTensor | None = None
64
- heavy_hidden_state: torch.FloatTensor | None = None
65
- next_latent_state: LatentDeliberationState | None = None
66
  past_key_values: Cache | None = None
 
 
67
  encoder_last_hidden_state: torch.FloatTensor | None = None
68
- temporal_context: torch.FloatTensor | None = None
69
- history_projected: torch.FloatTensor | None = None
70
  working_state: torch.FloatTensor | None = None
71
- latent_residual_diagnostics: dict[str, torch.Tensor] | None = None
72
  proposal: torch.LongTensor | None = None
73
  proposal_confidence: torch.FloatTensor | None = None
74
  token_entropy: torch.FloatTensor | None = None
75
  greedy_proposal: torch.LongTensor | None = None
76
  greedy_confidence: torch.FloatTensor | None = None
77
 
 
 
 
 
 
 
78
 
79
- class ModilifyMk2RMSNorm(DiffusionGemmaRMSNorm):
80
- """Official RMSNorm parameters, pre-fusion same-dtype forward."""
81
-
82
- def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
83
- normed_output = self._norm(hidden_states)
84
- if self.with_scale:
85
- normed_output = normed_output * self.weight.to(dtype=normed_output.dtype)
86
- return normed_output.type_as(hidden_states)
87
 
88
 
89
  class ModilifyMk2TextRouter(DiffusionGemmaTextRouter):
@@ -91,7 +85,7 @@ class ModilifyMk2TextRouter(DiffusionGemmaTextRouter):
91
 
92
  def __init__(self, config: Any) -> None:
93
  super().__init__(config)
94
- self.norm = ModilifyMk2RMSNorm(self.hidden_size, eps=self.eps, with_scale=False)
95
 
96
  def forward(
97
  self, hidden_states: torch.Tensor
@@ -114,18 +108,26 @@ class ModilifyMk2TextRouter(DiffusionGemmaTextRouter):
114
  return router_probabilities, top_k_weights, top_k_index
115
 
116
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
117
  def install_modilify_mk2_trunk_semantics(module: nn.Module) -> None:
118
  """Swap official leaf modules on this instance. Never patch Transformers classes."""
119
 
120
  for name, child in list(module.named_children()):
121
- if type(child) is DiffusionGemmaRMSNorm:
122
- dim = int(child.weight.shape[0]) if child.with_scale else 1
123
- replacement = ModilifyMk2RMSNorm(
124
- dim, eps=child.eps, with_scale=child.with_scale
125
- )
126
- replacement.load_state_dict(child.state_dict())
127
- setattr(module, name, replacement)
128
- elif type(child) is DiffusionGemmaTextRouter:
129
  replacement = ModilifyMk2TextRouter(child.config)
130
  replacement.load_state_dict(child.state_dict())
131
  setattr(module, name, replacement)
@@ -133,11 +135,26 @@ def install_modilify_mk2_trunk_semantics(module: nn.Module) -> None:
133
  install_modilify_mk2_trunk_semantics(child)
134
 
135
 
136
- class ModilifyMk2EncoderModel(DiffusionGemmaEncoderModel):
137
- """Unmodified Transformers DiffusionGemma multimodal encoder."""
138
 
139
  config_class = ModilifyMk2Config
140
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
141
 
142
  class ModilifyMk2DecoderModel(DiffusionGemmaDecoderModel):
143
  """Diffusion decoder accepting compact self-conditioning embeddings."""
@@ -171,7 +188,12 @@ class ModilifyMk2DecoderModel(DiffusionGemmaDecoderModel):
171
  if isinstance(decoder_attention_mask, dict) and all(
172
  mask.ndim == 4 for mask in decoder_attention_mask.values()
173
  ):
174
- return decoder_attention_mask
 
 
 
 
 
175
 
176
  text_config = config.get_text_config() if hasattr(config, "get_text_config") else config
177
  q_length = inputs_embeds.shape[1]
@@ -193,23 +215,24 @@ class ModilifyMk2DecoderModel(DiffusionGemmaDecoderModel):
193
  max_length = sliding_layer.get_max_length() + additional_kv_length
194
  if kv_length >= max_length:
195
  kv_length = max_length
196
- mask_mapping[layer_pattern] = ALL_MASK_ATTENTION_FUNCTIONS[
197
- config._attn_implementation
198
- ](
199
- batch_size=inputs_embeds.shape[0],
200
- q_length=q_length,
201
- kv_length=kv_length,
202
- q_offset=q_offset,
203
- kv_offset=kv_offset,
204
- mask_function=bidirectional_mask_function,
205
- attention_mask=decoder_attention_mask,
206
- allow_is_causal_skip=False,
207
- allow_is_bidirectional_skip=True,
208
- local_size=getattr(text_config, "sliding_window", None),
 
 
 
 
209
  dtype=inputs_embeds.dtype,
210
- config=text_config,
211
- use_vmap=False,
212
- device=inputs_embeds.device,
213
  )
214
  return mask_mapping
215
 
@@ -217,9 +240,7 @@ class ModilifyMk2DecoderModel(DiffusionGemmaDecoderModel):
217
  self,
218
  token_embeddings: torch.Tensor,
219
  latent_context: torch.Tensor | None,
220
- *,
221
- collect_diagnostics: bool = True,
222
- ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
223
  """Map latent context through the frozen native self-conditioning bridge."""
224
 
225
  if latent_context is None:
@@ -235,10 +256,6 @@ class ModilifyMk2DecoderModel(DiffusionGemmaDecoderModel):
235
  * mapper.up_proj(normalized_context)
236
  )
237
  mapped_fp32 = mapped_context.float()
238
- # Use the energy directly in the cap denominator. Computing
239
- # sqrt(E[x²]) and immediately squaring it again has an undefined
240
- # backward at the identity-init point x=0 (0/0 in d(sqrt)/dx), which
241
- # poisoned the detached temporal-context VJP on every denoise.
242
  mapped_token_energy = mapped_fp32.square().mean(dim=-1, keepdim=True)
243
  token_rms_per_token = (
244
  token_embeddings.float().square().mean(dim=-1, keepdim=True).sqrt()
@@ -249,21 +266,7 @@ class ModilifyMk2DecoderModel(DiffusionGemmaDecoderModel):
249
  )
250
  mapped_context = (mapped_fp32 * soft_cap_scale).to(dtype=mapped_context.dtype)
251
  combined = token_embeddings + mapped_context
252
- diagnostics: dict[str, torch.Tensor] = {}
253
- if collect_diagnostics:
254
- token_rms = token_embeddings.detach().float().square().mean().sqrt()
255
- mapped_rms = mapped_context.detach().float().square().mean().sqrt()
256
- diagnostics = {
257
- "token_embedding_rms": token_rms,
258
- "latent_context_rms": context.detach().float().square().mean().sqrt(),
259
- "mapped_sc_rms": mapped_rms,
260
- "actual_residual_rms": mapped_rms,
261
- "latent_to_embedding_rms_ratio": (
262
- mapped_rms / token_rms.clamp_min(1.0e-12)
263
- ),
264
- "merged_input_rms": combined.detach().float().square().mean().sqrt(),
265
- }
266
- return mapper.post_norm(combined), diagnostics
267
 
268
  def _run_stack(
269
  self,
@@ -288,10 +291,8 @@ class ModilifyMk2DecoderModel(DiffusionGemmaDecoderModel):
288
  )
289
  working_bus = kwargs.pop("working_bus", None)
290
  working_state = kwargs.pop("working_state", None)
 
291
  persistent_bus = kwargs.pop("persistent_bus", None)
292
- memory_bus = kwargs.pop("memory_bus", None)
293
- if memory_bus is None:
294
- memory_bus = persistent_bus
295
  memory_slots = kwargs.pop("memory_slots", None)
296
  slot_identity = kwargs.pop("slot_identity", None)
297
  cache = past_key_values
@@ -304,8 +305,8 @@ class ModilifyMk2DecoderModel(DiffusionGemmaDecoderModel):
304
  if working_bus is not None and working_state is not None:
305
  working_kv = working_bus.prepare_kv(working_state)
306
  memory_kv = None
307
- if memory_bus is not None and memory_slots is not None:
308
- memory_kv = memory_bus.prepare_kv(memory_slots, slot_identity)
309
  working_reader = 0
310
  reader_index = 0
311
  for index in range(self.text_config.num_hidden_layers):
@@ -325,14 +326,17 @@ class ModilifyMk2DecoderModel(DiffusionGemmaDecoderModel):
325
  and working_reader < working_bus.num_readers
326
  ):
327
  hidden = working_bus.read(
328
- hidden, working_reader, working_kv[0], working_kv[1]
 
 
 
329
  )
330
  working_reader += 1
331
  if (
332
  memory_kv is not None
333
- and reader_index < memory_bus.num_readers
334
  ):
335
- hidden = memory_bus.read(
336
  hidden, reader_index, memory_kv[0], memory_kv[1]
337
  )
338
  reader_index += 1
@@ -346,15 +350,13 @@ class ModilifyMk2DecoderModel(DiffusionGemmaDecoderModel):
346
  temporal_context_embeddings: torch.FloatTensor | None = None,
347
  decoder_attention_mask: torch.Tensor | dict | None = None,
348
  decoder_position_ids: torch.LongTensor | None = None,
349
- collect_latent_diagnostics: bool = True,
350
- memory_bus: Any | None = None,
351
  memory_slots: torch.Tensor | None = None,
352
  working_bus: Any | None = None,
353
  working_state: torch.Tensor | None = None,
354
  persistent_bus: Any | None = None,
355
  slot_identity: torch.Tensor | None = None,
356
  **kwargs: Any,
357
- ) -> ModilifyMk2DecoderOutput:
358
  if "use_cache" in kwargs:
359
  raise ValueError("The diffusion decoder always reads the supplied cache.")
360
  if decoder_token_embeddings is None:
@@ -364,22 +366,15 @@ class ModilifyMk2DecoderModel(DiffusionGemmaDecoderModel):
364
  expected = (*decoder_input_ids.shape, self.text_config.hidden_size)
365
  if token_embeddings.shape != expected:
366
  raise ValueError("Precomputed decoder embeddings have the wrong shape.")
367
- context_embeddings = (
368
- torch.zeros_like(token_embeddings)
369
- if temporal_context_embeddings is None
370
- else temporal_context_embeddings.to(token_embeddings)
371
- )
372
- inputs_embeds, diagnostics = self.merge_latent_context(
373
  token_embeddings,
374
- context_embeddings if temporal_context_embeddings is not None else None,
375
- collect_diagnostics=collect_latent_diagnostics,
376
  )
377
  hidden = self._run_stack(
378
  inputs_embeds,
379
  past_key_values=past_key_values,
380
  decoder_attention_mask=decoder_attention_mask,
381
  decoder_position_ids=decoder_position_ids,
382
- memory_bus=memory_bus,
383
  memory_slots=memory_slots,
384
  working_bus=working_bus,
385
  working_state=working_state,
@@ -387,11 +382,9 @@ class ModilifyMk2DecoderModel(DiffusionGemmaDecoderModel):
387
  slot_identity=slot_identity,
388
  **kwargs,
389
  )
390
- return ModilifyMk2DecoderOutput(
391
  last_hidden_state=hidden,
392
  past_key_values=past_key_values,
393
- token_embeddings=token_embeddings,
394
- latent_residual_diagnostics=diagnostics,
395
  )
396
 
397
 
@@ -412,6 +405,16 @@ class ModilifyMk2Model(DiffusionGemmaPreTrainedModel):
412
  self.encoder = ModilifyMk2EncoderModel(config)
413
  self.decoder = ModilifyMk2DecoderModel(config)
414
  install_modilify_mk2_trunk_semantics(self)
 
 
 
 
 
 
 
 
 
 
415
  self.post_init()
416
 
417
  def get_encoder(self):
@@ -440,8 +443,6 @@ class ModilifyMk2Model(DiffusionGemmaPreTrainedModel):
440
  decoder_attention_mask: torch.Tensor | dict | None = None,
441
  decoder_position_ids: torch.LongTensor | None = None,
442
  return_encoder_outputs: bool = True,
443
- collect_latent_diagnostics: bool = True,
444
- memory_bus: Any | None = None,
445
  memory_slots: torch.Tensor | None = None,
446
  working_bus: Any | None = None,
447
  working_state: torch.Tensor | None = None,
@@ -450,15 +451,13 @@ class ModilifyMk2Model(DiffusionGemmaPreTrainedModel):
450
  **kwargs: Any,
451
  ) -> ModilifyMk2ModelOutput:
452
  encoder_hidden = None
453
- encoder_keys = ("pixel_values", "mm_token_type_ids", "image_position_ids", "inputs_embeds")
454
- encoder_kwargs = {key: kwargs.pop(key) for key in encoder_keys if key in kwargs}
455
  if input_ids is not None:
456
  encoded = self.encoder(
457
  input_ids=input_ids,
458
  attention_mask=attention_mask,
459
  past_key_values=past_key_values,
460
  position_ids=position_ids,
461
- **encoder_kwargs,
462
  )
463
  past_key_values = encoded.past_key_values
464
  if return_encoder_outputs:
@@ -472,8 +471,6 @@ class ModilifyMk2Model(DiffusionGemmaPreTrainedModel):
472
  temporal_context_embeddings=temporal_context_embeddings,
473
  decoder_attention_mask=decoder_attention_mask,
474
  decoder_position_ids=decoder_position_ids,
475
- collect_latent_diagnostics=collect_latent_diagnostics,
476
- memory_bus=memory_bus,
477
  memory_slots=memory_slots,
478
  working_bus=working_bus,
479
  working_state=working_state,
@@ -484,9 +481,7 @@ class ModilifyMk2Model(DiffusionGemmaPreTrainedModel):
484
  return ModilifyMk2ModelOutput(
485
  last_hidden_state=decoded.last_hidden_state,
486
  past_key_values=past_key_values,
487
- token_embeddings=decoded.token_embeddings,
488
  encoder_last_hidden_state=encoder_hidden,
489
- latent_residual_diagnostics=decoded.latent_residual_diagnostics,
490
  )
491
 
492
 
@@ -494,14 +489,13 @@ class ModilifyMk2ForBlockDiffusion(DiffusionGemmaPreTrainedModel, ModilifyMk2Gen
494
  config_class = ModilifyMk2Config
495
  _tied_weights_keys = {"lm_head.weight": "model.decoder.embed_tokens.weight"}
496
  generation_config_class = ModilifyMk2GenerationConfig
 
497
 
498
  @torch.no_grad()
499
  def _init_weights(self, module: nn.Module) -> None:
500
  super()._init_weights(module)
501
  if isinstance(module, LatentDeliberationTransformer):
502
  module.reset_identity_parameters()
503
- module.working_memory_bus.freeze()
504
- module.persistent_memory_bus.freeze()
505
 
506
  def __init__(self, config: ModilifyMk2Config):
507
  super().__init__(config)
@@ -509,15 +503,11 @@ class ModilifyMk2ForBlockDiffusion(DiffusionGemmaPreTrainedModel, ModilifyMk2Gen
509
  layer_types = tuple(getattr(config.text_config, "layer_types", None) or ())
510
  self.latent_deliberation = LatentDeliberationTransformer(
511
  hidden_size=config.text_config.hidden_size,
512
- vocab_size=config.text_config.vocab_size,
513
  latent_dim=config.latent_dim,
514
  ffn_dim=config.latent_ffn_dim,
515
- memory_slots=config.latent_memory_slots,
516
  num_layers=config.latent_num_layers,
517
  num_heads=config.latent_num_heads,
518
  local_attention_window=config.latent_local_attention_window,
519
- dropout=config.latent_dropout,
520
- history_length=config.latent_history_length,
521
  tape_probes=config.latent_tape_probes,
522
  history_kv_rank=config.latent_history_kv_rank,
523
  num_memory_readers=sum(layer_type == "full_attention" for layer_type in layer_types),
@@ -530,13 +520,17 @@ class ModilifyMk2ForBlockDiffusion(DiffusionGemmaPreTrainedModel, ModilifyMk2Gen
530
  if config.persistent_memory_bus else 0
531
  ),
532
  working_last_block_global=config.latent_working_last_block_global,
533
- experience_roles=config.experience_roles,
534
- commit_sequence_layers=config.commit_sequence_layers,
535
  commit_sequence_dim=config.commit_sequence_dim,
536
  max_canvas_length=config.canvas_length,
537
  )
538
  self.lm_head = nn.Linear(config.text_config.hidden_size, config.text_config.vocab_size, bias=False)
539
  self.final_logit_softcapping = config.text_config.final_logit_softcapping
 
 
 
 
 
 
540
  self.post_init()
541
  install_modilify_mk2_trunk_semantics(self)
542
 
@@ -544,192 +538,101 @@ class ModilifyMk2ForBlockDiffusion(DiffusionGemmaPreTrainedModel, ModilifyMk2Gen
544
  logits = self.lm_head(hidden)
545
  return torch.tanh(logits / self.final_logit_softcapping) * self.final_logit_softcapping
546
 
 
547
  def _prepare_latent_context(
548
  self,
549
  decoder_input_ids: torch.LongTensor,
550
  *,
551
- history: TrajectoryHistory | None,
552
- tape: TrajectoryTape | None,
553
  confidence: torch.Tensor | None,
554
  entropy: torch.Tensor | None,
555
- age: torch.Tensor | None,
556
  latent_state: LatentDeliberationState | None,
557
- ) -> tuple[torch.Tensor, LatentDeliberationState, torch.Tensor, torch.Tensor]:
 
558
  batch, canvas = decoder_input_ids.shape
559
- dtype = self.model.decoder.embed_tokens.weight.dtype
560
  if latent_state is None:
561
  latent_state = LatentDeliberationState.empty(
562
  batch_size=batch, canvas_length=canvas,
563
- latent_dim=self.config.latent_dim, memory_slots=self.config.latent_memory_slots,
564
- device=decoder_input_ids.device, dtype=dtype,
565
- )
566
- if history is None:
567
- history = TrajectoryHistory.empty(
568
- batch_size=batch,
569
- canvas_length=canvas,
570
- hidden_size=self.config.text_config.hidden_size,
571
- history_length=self.config.latent_history_length,
572
  device=decoder_input_ids.device,
573
- dtype=dtype,
574
- )
575
- if tape is None:
576
- tape = empty_trajectory_tape(
577
- batch_size=batch,
578
- config=self.config,
579
- device=decoder_input_ids.device,
580
- dtype=dtype,
581
  )
582
  confidence = latent_state.confidence if confidence is None else confidence.squeeze(-1).float()
583
  entropy = latent_state.entropy if entropy is None else entropy.squeeze(-1).float()
584
- if age is not None:
585
- latent_state = replace(
586
- latent_state,
587
- age=age.to(device=decoder_input_ids.device, dtype=torch.int32),
588
- )
589
  token_embeddings = self.model.decoder.embed_tokens(decoder_input_ids)
590
  processed: LatentProcessorOutput = self.latent_deliberation(
591
  token_embeddings=token_embeddings,
592
  confidence=confidence,
593
  entropy=entropy,
594
  state=latent_state,
595
- history=history,
596
- tape=tape,
597
  )
 
 
598
  return (
599
- processed.context,
600
- processed.state,
601
  token_embeddings,
602
- processed.history_projected,
603
  )
604
 
605
  def forward(
606
- self,
607
- *,
608
- input_ids: torch.LongTensor | None = None,
609
  attention_mask: torch.Tensor | dict | None = None,
610
  past_key_values: Cache | None = None,
611
  position_ids: torch.LongTensor | None = None,
612
  decoder_input_ids: torch.LongTensor,
613
  previous_confidence: torch.FloatTensor | None = None,
614
  previous_entropy: torch.FloatTensor | None = None,
615
- token_age: torch.Tensor | None = None,
616
  latent_state: LatentDeliberationState | None = None,
617
- history: TrajectoryHistory | None = None,
618
- tape: TrajectoryTape | None = None,
619
- history_hidden_state: TrajectoryHistory | torch.FloatTensor | None = None,
620
  decoder_attention_mask: torch.Tensor | dict | None = None,
621
  decoder_position_ids: torch.LongTensor | None = None,
622
  return_encoder_outputs: bool = True,
623
  compact_vocab: bool = False,
624
- denoise_temperature: float | None = None,
625
  repetition_token_mask: torch.BoolTensor | None = None,
626
  repetition_penalty: float = 1.0,
627
  sampling_generators: Sequence[torch.Generator] | None = None,
628
- collect_latent_diagnostics: bool = True,
629
  **kwargs: Any,
630
  ) -> ModilifyMk2BlockDiffusionOutput:
631
- if history is None and isinstance(history_hidden_state, TrajectoryHistory):
632
- history = history_hidden_state
633
- (
634
- latent_context,
635
- next_state,
636
- decoder_token_embeddings,
637
- history_projected,
638
- ) = self._prepare_latent_context(
639
- decoder_input_ids,
640
- history=history,
641
- tape=tape,
642
- confidence=previous_confidence,
643
- entropy=previous_entropy,
644
- age=token_age,
645
- latent_state=latent_state,
646
  )
647
  working_bus = self.latent_deliberation.working_memory_bus
648
  persistent_bus = self.latent_deliberation.persistent_memory_bus
649
- working_readers_active = working_bus.num_readers > 0
650
- memory_readers_active = persistent_bus.num_readers > 0
651
  outputs = self.model(
652
  input_ids=input_ids, attention_mask=attention_mask,
653
  past_key_values=past_key_values, position_ids=position_ids,
654
- decoder_input_ids=decoder_input_ids,
655
- decoder_token_embeddings=decoder_token_embeddings,
656
  temporal_context_embeddings=latent_context,
657
  decoder_attention_mask=decoder_attention_mask,
658
  decoder_position_ids=decoder_position_ids,
659
  return_encoder_outputs=return_encoder_outputs,
660
- collect_latent_diagnostics=collect_latent_diagnostics,
661
- working_bus=working_bus if working_readers_active else None,
662
- working_state=latent_context if working_readers_active else None,
663
- persistent_bus=persistent_bus if memory_readers_active else None,
664
- memory_slots=next_state.memory_slots if memory_readers_active else None,
665
- slot_identity=(
666
- self.latent_deliberation.scaled_memory_slot_identity(
667
- batch_size=decoder_input_ids.shape[0],
668
- device=decoder_input_ids.device,
669
- dtype=latent_context.dtype,
670
- )
671
- if memory_readers_active else None
672
- ),
673
- **kwargs,
674
- )
675
- temperature = (
676
- self.config.denoise_temperature
677
- if denoise_temperature is None
678
- else float(denoise_temperature)
679
  )
680
- proposal = proposal_confidence = token_entropy = None
681
- greedy_proposal = greedy_confidence = None
682
- logits = None
683
  if compact_vocab:
684
- (
685
- proposal,
686
- proposal_confidence,
687
- token_entropy,
688
- greedy_proposal,
689
- greedy_confidence,
690
- ) = chunked_vocab_statistics(
691
- outputs.last_hidden_state.detach(),
692
- self.lm_head.weight.detach(),
693
- softcap=self.final_logit_softcapping,
694
- temperature=temperature,
695
  chunk_size=self.config.vocab_chunk_size,
696
- repetition_token_mask=repetition_token_mask,
697
- repetition_penalty=repetition_penalty,
698
  sampling_generators=sampling_generators,
 
699
  )
700
  else:
701
  logits = self._finalize_logits(outputs.last_hidden_state)
702
  return ModilifyMk2BlockDiffusionOutput(
703
- logits=logits,
704
- past_key_values=outputs.past_key_values,
705
  encoder_last_hidden_state=outputs.encoder_last_hidden_state,
706
- heavy_hidden_state=outputs.last_hidden_state,
707
- next_latent_state=next_state,
708
- temporal_context=latent_context,
709
- history_projected=history_projected,
710
- working_state=latent_context,
711
- latent_residual_diagnostics=outputs.latent_residual_diagnostics,
712
- proposal=proposal,
713
- proposal_confidence=proposal_confidence,
714
- token_entropy=token_entropy,
715
- greedy_proposal=greedy_proposal,
716
  greedy_confidence=greedy_confidence,
717
  )
718
-
719
-
720
- ModilifyMk2Model.register_for_auto_class("AutoModel")
721
- ModilifyMk2ForBlockDiffusion.register_for_auto_class("AutoModelForCausalLM")
722
- ModilifyMk2ForBlockDiffusion.register_for_auto_class("AutoModelForMultimodalLM")
723
-
724
-
725
- __all__ = [
726
- "ModilifyMk2BlockDiffusionOutput",
727
- "ModilifyMk2Config",
728
- "ModilifyMk2DecoderModel",
729
- "ModilifyMk2EncoderModel",
730
- "ModilifyMk2ForBlockDiffusion",
731
- "ModilifyMk2Model",
732
- "ModilifyMk2RMSNorm",
733
- "ModilifyMk2TextRouter",
734
- "install_modilify_mk2_trunk_semantics",
735
- ]
 
1
+ """Text-only ModilifyMk2 model with recurrent latent deliberation."""
 
 
2
 
3
  from __future__ import annotations
4
 
5
  from collections.abc import Sequence
6
+ from dataclasses import dataclass
7
  from typing import Any
8
 
9
  import torch
 
11
  from torch.nn import functional as F
12
  from transformers.cache_utils import Cache
13
  from transformers.modeling_outputs import BaseModelOutputWithPast
 
14
 
15
  from transformers.masking_utils import (
16
  ALL_MASK_ATTENTION_FUNCTIONS,
 
18
  )
19
  from transformers.models.diffusion_gemma import (
20
  DiffusionGemmaDecoderModel,
21
+ DiffusionGemmaEncoderTextModel,
22
  DiffusionGemmaPreTrainedModel,
23
  )
24
  from transformers.models.diffusion_gemma.modeling_diffusion_gemma import (
 
26
  DiffusionGemmaTextRouter,
27
  )
28
 
29
+ from .mps_ops import register_mps_backends
30
+ from .lora import LoRAConfig, inject_lora
31
+ from .configuration_modilify_mk2 import ModilifyMk2Config, DENOISE_TEMPERATURE
32
  from .generation_modilify_mk2 import ModilifyMk2GenerationConfig, ModilifyMk2GenerationMixin
33
+ from .latent_deliberation import LatentDeliberationState, LatentProcessorOutput, _sdpa_mask_value
34
+ from .latent_deliberation import LatentDeliberationTransformer
 
 
 
 
 
 
35
  from .vocab_ops import chunked_vocab_statistics
36
 
37
 
38
+ PRECISION_POLICY = "gdn2_small_fp32_v1"
39
+
40
+
41
+ def small_parameter_fp32(name: str) -> bool:
42
+ return name.startswith("latent_deliberation.") and (
43
+ name.endswith((".dt_bias", ".a_log"))
44
+ or (name.endswith(".weight") and "norm" in name.rsplit(".", 2)[-2])
45
+ )
46
+
47
+
48
+ register_mps_backends()
49
 
50
 
51
  @dataclass
52
  class ModilifyMk2ModelOutput(BaseModelOutputWithPast):
 
53
  encoder_last_hidden_state: torch.FloatTensor | None = None
 
54
 
55
 
56
  @dataclass
57
+ class ModilifyMk2BlockDiffusionOutput:
58
+ logits: torch.FloatTensor | None
59
+ heavy_hidden_state: torch.FloatTensor
60
+ next_latent_state: LatentDeliberationState
 
 
61
  past_key_values: Cache | None = None
62
+ hidden_states: tuple[torch.FloatTensor, ...] | None = None
63
+ attentions: tuple[torch.FloatTensor, ...] | None = None
64
  encoder_last_hidden_state: torch.FloatTensor | None = None
 
 
65
  working_state: torch.FloatTensor | None = None
 
66
  proposal: torch.LongTensor | None = None
67
  proposal_confidence: torch.FloatTensor | None = None
68
  token_entropy: torch.FloatTensor | None = None
69
  greedy_proposal: torch.LongTensor | None = None
70
  greedy_confidence: torch.FloatTensor | None = None
71
 
72
+ def to_tuple(self) -> tuple[Any, ...]:
73
+ return tuple(value for value in (
74
+ self.logits, self.past_key_values, self.hidden_states,
75
+ self.attentions, self.encoder_last_hidden_state,
76
+ self.heavy_hidden_state, self.next_latent_state,
77
+ ) if value is not None)
78
 
79
+ def __getitem__(self, key: str | int) -> Any:
80
+ return getattr(self, key) if isinstance(key, str) else self.to_tuple()[key]
 
 
 
 
 
 
81
 
82
 
83
  class ModilifyMk2TextRouter(DiffusionGemmaTextRouter):
 
85
 
86
  def __init__(self, config: Any) -> None:
87
  super().__init__(config)
88
+ self.norm = DiffusionGemmaRMSNorm(self.hidden_size, eps=self.eps, with_scale=False)
89
 
90
  def forward(
91
  self, hidden_states: torch.Tensor
 
108
  return router_probabilities, top_k_weights, top_k_index
109
 
110
 
111
+ def _finite_decoder_attention_mask(
112
+ mask: torch.Tensor | None,
113
+ *,
114
+ dtype: torch.dtype,
115
+ ) -> torch.Tensor | None:
116
+ """Keep boolean masks and finite values for fully blocked query rows."""
117
+
118
+ del dtype
119
+ if mask is None:
120
+ return None
121
+ if mask.dtype == torch.bool:
122
+ return mask | ~mask.any(dim=-1, keepdim=True)
123
+ return mask.clamp_min(_sdpa_mask_value(mask.dtype))
124
+
125
+
126
  def install_modilify_mk2_trunk_semantics(module: nn.Module) -> None:
127
  """Swap official leaf modules on this instance. Never patch Transformers classes."""
128
 
129
  for name, child in list(module.named_children()):
130
+ if type(child) is DiffusionGemmaTextRouter:
 
 
 
 
 
 
 
131
  replacement = ModilifyMk2TextRouter(child.config)
132
  replacement.load_state_dict(child.state_dict())
133
  setattr(module, name, replacement)
 
135
  install_modilify_mk2_trunk_semantics(child)
136
 
137
 
138
+ class ModilifyMk2EncoderModel(DiffusionGemmaPreTrainedModel):
139
+ """Checkpoint-compatible wrapper around the text encoder trunk."""
140
 
141
  config_class = ModilifyMk2Config
142
 
143
+ def __init__(self, config: ModilifyMk2Config):
144
+ super().__init__(config)
145
+ self.language_model = DiffusionGemmaEncoderTextModel(config.text_config)
146
+ install_modilify_mk2_trunk_semantics(self)
147
+ self.post_init()
148
+
149
+ def get_input_embeddings(self) -> nn.Module:
150
+ return self.language_model.embed_tokens
151
+
152
+ def set_input_embeddings(self, value: nn.Module) -> None:
153
+ self.language_model.embed_tokens = value
154
+
155
+ def forward(self, **kwargs: Any) -> BaseModelOutputWithPast:
156
+ return self.language_model(**kwargs)
157
+
158
 
159
  class ModilifyMk2DecoderModel(DiffusionGemmaDecoderModel):
160
  """Diffusion decoder accepting compact self-conditioning embeddings."""
 
188
  if isinstance(decoder_attention_mask, dict) and all(
189
  mask.ndim == 4 for mask in decoder_attention_mask.values()
190
  ):
191
+ return {
192
+ name: _finite_decoder_attention_mask(
193
+ mask, dtype=inputs_embeds.dtype
194
+ )
195
+ for name, mask in decoder_attention_mask.items()
196
+ }
197
 
198
  text_config = config.get_text_config() if hasattr(config, "get_text_config") else config
199
  q_length = inputs_embeds.shape[1]
 
215
  max_length = sliding_layer.get_max_length() + additional_kv_length
216
  if kv_length >= max_length:
217
  kv_length = max_length
218
+ mask_mapping[layer_pattern] = _finite_decoder_attention_mask(
219
+ ALL_MASK_ATTENTION_FUNCTIONS[config._attn_implementation](
220
+ batch_size=inputs_embeds.shape[0],
221
+ q_length=q_length,
222
+ kv_length=kv_length,
223
+ q_offset=q_offset,
224
+ kv_offset=kv_offset,
225
+ mask_function=bidirectional_mask_function,
226
+ attention_mask=decoder_attention_mask,
227
+ allow_is_causal_skip=False,
228
+ allow_is_bidirectional_skip=True,
229
+ local_size=getattr(text_config, "sliding_window", None),
230
+ dtype=inputs_embeds.dtype,
231
+ config=text_config,
232
+ use_vmap=False,
233
+ device=inputs_embeds.device,
234
+ ),
235
  dtype=inputs_embeds.dtype,
 
 
 
236
  )
237
  return mask_mapping
238
 
 
240
  self,
241
  token_embeddings: torch.Tensor,
242
  latent_context: torch.Tensor | None,
243
+ ) -> torch.Tensor:
 
 
244
  """Map latent context through the frozen native self-conditioning bridge."""
245
 
246
  if latent_context is None:
 
256
  * mapper.up_proj(normalized_context)
257
  )
258
  mapped_fp32 = mapped_context.float()
 
 
 
 
259
  mapped_token_energy = mapped_fp32.square().mean(dim=-1, keepdim=True)
260
  token_rms_per_token = (
261
  token_embeddings.float().square().mean(dim=-1, keepdim=True).sqrt()
 
266
  )
267
  mapped_context = (mapped_fp32 * soft_cap_scale).to(dtype=mapped_context.dtype)
268
  combined = token_embeddings + mapped_context
269
+ return mapper.post_norm(combined)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
270
 
271
  def _run_stack(
272
  self,
 
291
  )
292
  working_bus = kwargs.pop("working_bus", None)
293
  working_state = kwargs.pop("working_state", None)
294
+ working_canvas_head = kwargs.pop("working_canvas_head", None)
295
  persistent_bus = kwargs.pop("persistent_bus", None)
 
 
 
296
  memory_slots = kwargs.pop("memory_slots", None)
297
  slot_identity = kwargs.pop("slot_identity", None)
298
  cache = past_key_values
 
305
  if working_bus is not None and working_state is not None:
306
  working_kv = working_bus.prepare_kv(working_state)
307
  memory_kv = None
308
+ if persistent_bus is not None and memory_slots is not None:
309
+ memory_kv = persistent_bus.prepare_kv(memory_slots, slot_identity)
310
  working_reader = 0
311
  reader_index = 0
312
  for index in range(self.text_config.num_hidden_layers):
 
326
  and working_reader < working_bus.num_readers
327
  ):
328
  hidden = working_bus.read(
329
+ hidden, working_reader, *working_kv,
330
+ positions=(
331
+ decoder_position_ids if working_canvas_head is not None else None
332
+ ),
333
  )
334
  working_reader += 1
335
  if (
336
  memory_kv is not None
337
+ and reader_index < persistent_bus.num_readers
338
  ):
339
+ hidden = persistent_bus.read(
340
  hidden, reader_index, memory_kv[0], memory_kv[1]
341
  )
342
  reader_index += 1
 
350
  temporal_context_embeddings: torch.FloatTensor | None = None,
351
  decoder_attention_mask: torch.Tensor | dict | None = None,
352
  decoder_position_ids: torch.LongTensor | None = None,
 
 
353
  memory_slots: torch.Tensor | None = None,
354
  working_bus: Any | None = None,
355
  working_state: torch.Tensor | None = None,
356
  persistent_bus: Any | None = None,
357
  slot_identity: torch.Tensor | None = None,
358
  **kwargs: Any,
359
+ ) -> BaseModelOutputWithPast:
360
  if "use_cache" in kwargs:
361
  raise ValueError("The diffusion decoder always reads the supplied cache.")
362
  if decoder_token_embeddings is None:
 
366
  expected = (*decoder_input_ids.shape, self.text_config.hidden_size)
367
  if token_embeddings.shape != expected:
368
  raise ValueError("Precomputed decoder embeddings have the wrong shape.")
369
+ inputs_embeds = self.merge_latent_context(
 
 
 
 
 
370
  token_embeddings,
371
+ temporal_context_embeddings,
 
372
  )
373
  hidden = self._run_stack(
374
  inputs_embeds,
375
  past_key_values=past_key_values,
376
  decoder_attention_mask=decoder_attention_mask,
377
  decoder_position_ids=decoder_position_ids,
 
378
  memory_slots=memory_slots,
379
  working_bus=working_bus,
380
  working_state=working_state,
 
382
  slot_identity=slot_identity,
383
  **kwargs,
384
  )
385
+ return BaseModelOutputWithPast(
386
  last_hidden_state=hidden,
387
  past_key_values=past_key_values,
 
 
388
  )
389
 
390
 
 
405
  self.encoder = ModilifyMk2EncoderModel(config)
406
  self.decoder = ModilifyMk2DecoderModel(config)
407
  install_modilify_mk2_trunk_semantics(self)
408
+ if getattr(config, "lora_config", None):
409
+ config.text_config._experts_implementation = "mps_segmented"
410
+ inject_lora(self, LoRAConfig(**config.lora_config))
411
+ self._tied_weights_keys = {
412
+ **self._tied_weights_keys,
413
+ r"encoder.language_model.layers\.(?:[^.]+\.)*lora_[ab]$":
414
+ r"decoder.layers\.(?:[^.]+\.)*lora_[ab]$",
415
+ r"encoder.language_model.layers\.(?:[^.]+\.)*lora_(?:gate_up|down)_[ab]$":
416
+ r"decoder.layers\.(?:[^.]+\.)*lora_(?:gate_up|down)_[ab]$",
417
+ }
418
  self.post_init()
419
 
420
  def get_encoder(self):
 
443
  decoder_attention_mask: torch.Tensor | dict | None = None,
444
  decoder_position_ids: torch.LongTensor | None = None,
445
  return_encoder_outputs: bool = True,
 
 
446
  memory_slots: torch.Tensor | None = None,
447
  working_bus: Any | None = None,
448
  working_state: torch.Tensor | None = None,
 
451
  **kwargs: Any,
452
  ) -> ModilifyMk2ModelOutput:
453
  encoder_hidden = None
 
 
454
  if input_ids is not None:
455
  encoded = self.encoder(
456
  input_ids=input_ids,
457
  attention_mask=attention_mask,
458
  past_key_values=past_key_values,
459
  position_ids=position_ids,
460
+ **kwargs,
461
  )
462
  past_key_values = encoded.past_key_values
463
  if return_encoder_outputs:
 
471
  temporal_context_embeddings=temporal_context_embeddings,
472
  decoder_attention_mask=decoder_attention_mask,
473
  decoder_position_ids=decoder_position_ids,
 
 
474
  memory_slots=memory_slots,
475
  working_bus=working_bus,
476
  working_state=working_state,
 
481
  return ModilifyMk2ModelOutput(
482
  last_hidden_state=decoded.last_hidden_state,
483
  past_key_values=past_key_values,
 
484
  encoder_last_hidden_state=encoder_hidden,
 
485
  )
486
 
487
 
 
489
  config_class = ModilifyMk2Config
490
  _tied_weights_keys = {"lm_head.weight": "model.decoder.embed_tokens.weight"}
491
  generation_config_class = ModilifyMk2GenerationConfig
492
+ supports_gradient_checkpointing = False
493
 
494
  @torch.no_grad()
495
  def _init_weights(self, module: nn.Module) -> None:
496
  super()._init_weights(module)
497
  if isinstance(module, LatentDeliberationTransformer):
498
  module.reset_identity_parameters()
 
 
499
 
500
  def __init__(self, config: ModilifyMk2Config):
501
  super().__init__(config)
 
503
  layer_types = tuple(getattr(config.text_config, "layer_types", None) or ())
504
  self.latent_deliberation = LatentDeliberationTransformer(
505
  hidden_size=config.text_config.hidden_size,
 
506
  latent_dim=config.latent_dim,
507
  ffn_dim=config.latent_ffn_dim,
 
508
  num_layers=config.latent_num_layers,
509
  num_heads=config.latent_num_heads,
510
  local_attention_window=config.latent_local_attention_window,
 
 
511
  tape_probes=config.latent_tape_probes,
512
  history_kv_rank=config.latent_history_kv_rank,
513
  num_memory_readers=sum(layer_type == "full_attention" for layer_type in layer_types),
 
520
  if config.persistent_memory_bus else 0
521
  ),
522
  working_last_block_global=config.latent_working_last_block_global,
 
 
523
  commit_sequence_dim=config.commit_sequence_dim,
524
  max_canvas_length=config.canvas_length,
525
  )
526
  self.lm_head = nn.Linear(config.text_config.hidden_size, config.text_config.vocab_size, bias=False)
527
  self.final_logit_softcapping = config.text_config.final_logit_softcapping
528
+ policy = getattr(config, "precision_policy", PRECISION_POLICY)
529
+ if policy != PRECISION_POLICY:
530
+ raise ValueError(f"Unsupported export precision policy: {policy}")
531
+ self._keep_in_fp32_modules_strict = {
532
+ name for name, _ in self.named_parameters() if small_parameter_fp32(name)
533
+ }
534
  self.post_init()
535
  install_modilify_mk2_trunk_semantics(self)
536
 
 
538
  logits = self.lm_head(hidden)
539
  return torch.tanh(logits / self.final_logit_softcapping) * self.final_logit_softcapping
540
 
541
+
542
  def _prepare_latent_context(
543
  self,
544
  decoder_input_ids: torch.LongTensor,
545
  *,
546
+
 
547
  confidence: torch.Tensor | None,
548
  entropy: torch.Tensor | None,
 
549
  latent_state: LatentDeliberationState | None,
550
+ canvas_head: torch.Tensor | None,
551
+ ) -> tuple[torch.Tensor, LatentDeliberationState, torch.Tensor]:
552
  batch, canvas = decoder_input_ids.shape
 
553
  if latent_state is None:
554
  latent_state = LatentDeliberationState.empty(
555
  batch_size=batch, canvas_length=canvas,
 
 
 
 
 
 
 
 
 
556
  device=decoder_input_ids.device,
 
 
 
 
 
 
 
 
557
  )
558
  confidence = latent_state.confidence if confidence is None else confidence.squeeze(-1).float()
559
  entropy = latent_state.entropy if entropy is None else entropy.squeeze(-1).float()
 
 
 
 
 
560
  token_embeddings = self.model.decoder.embed_tokens(decoder_input_ids)
561
  processed: LatentProcessorOutput = self.latent_deliberation(
562
  token_embeddings=token_embeddings,
563
  confidence=confidence,
564
  entropy=entropy,
565
  state=latent_state,
566
+
567
+ canvas_head=canvas_head,
568
  )
569
+ latent_context = processed.context
570
+ next_state = processed.state
571
  return (
572
+ latent_context,
573
+ next_state,
574
  token_embeddings,
 
575
  )
576
 
577
  def forward(
578
+ self, *, input_ids: torch.LongTensor | None = None,
 
 
579
  attention_mask: torch.Tensor | dict | None = None,
580
  past_key_values: Cache | None = None,
581
  position_ids: torch.LongTensor | None = None,
582
  decoder_input_ids: torch.LongTensor,
583
  previous_confidence: torch.FloatTensor | None = None,
584
  previous_entropy: torch.FloatTensor | None = None,
 
585
  latent_state: LatentDeliberationState | None = None,
586
+ canvas_head: torch.Tensor | None = None,
 
 
587
  decoder_attention_mask: torch.Tensor | dict | None = None,
588
  decoder_position_ids: torch.LongTensor | None = None,
589
  return_encoder_outputs: bool = True,
590
  compact_vocab: bool = False,
 
591
  repetition_token_mask: torch.BoolTensor | None = None,
592
  repetition_penalty: float = 1.0,
593
  sampling_generators: Sequence[torch.Generator] | None = None,
 
594
  **kwargs: Any,
595
  ) -> ModilifyMk2BlockDiffusionOutput:
596
+ latent_context, next_state, token_embeddings = self._prepare_latent_context(
597
+ decoder_input_ids, confidence=previous_confidence, entropy=previous_entropy,
598
+ latent_state=latent_state, canvas_head=canvas_head,
 
 
 
 
 
 
 
 
 
 
 
 
599
  )
600
  working_bus = self.latent_deliberation.working_memory_bus
601
  persistent_bus = self.latent_deliberation.persistent_memory_bus
 
 
602
  outputs = self.model(
603
  input_ids=input_ids, attention_mask=attention_mask,
604
  past_key_values=past_key_values, position_ids=position_ids,
605
+ decoder_input_ids=decoder_input_ids, decoder_token_embeddings=token_embeddings,
 
606
  temporal_context_embeddings=latent_context,
607
  decoder_attention_mask=decoder_attention_mask,
608
  decoder_position_ids=decoder_position_ids,
609
  return_encoder_outputs=return_encoder_outputs,
610
+ working_bus=working_bus,
611
+ working_state=(latent_context, next_state.gdn2.seen),
612
+ working_canvas_head=canvas_head,
613
+ persistent_bus=persistent_bus, memory_slots=next_state.memory_slots,
614
+ slot_identity=next_state.gdn2.seen,
615
+ **kwargs,
 
 
 
 
 
 
 
 
 
 
 
 
 
616
  )
617
+ proposal = proposal_confidence = token_entropy = greedy_proposal = greedy_confidence = None
 
 
618
  if compact_vocab:
619
+ logits = None
620
+ proposal, proposal_confidence, token_entropy, greedy_proposal, greedy_confidence = chunked_vocab_statistics(
621
+ outputs.last_hidden_state, self.lm_head.weight,
622
+ softcap=self.final_logit_softcapping, temperature=DENOISE_TEMPERATURE,
 
 
 
 
 
 
 
623
  chunk_size=self.config.vocab_chunk_size,
624
+ repetition_token_mask=repetition_token_mask, repetition_penalty=repetition_penalty,
 
625
  sampling_generators=sampling_generators,
626
+ top_k=self.config.commit_top_k, min_p=self.config.commit_min_p,
627
  )
628
  else:
629
  logits = self._finalize_logits(outputs.last_hidden_state)
630
  return ModilifyMk2BlockDiffusionOutput(
631
+ logits=logits, past_key_values=outputs.past_key_values,
 
632
  encoder_last_hidden_state=outputs.encoder_last_hidden_state,
633
+ heavy_hidden_state=outputs.last_hidden_state, next_latent_state=next_state,
634
+ working_state=latent_context,
635
+ proposal=proposal, proposal_confidence=proposal_confidence,
636
+ token_entropy=token_entropy, greedy_proposal=greedy_proposal,
 
 
 
 
 
 
637
  greedy_confidence=greedy_confidence,
638
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
mps_ops.py CHANGED
@@ -1,5 +1,3 @@
1
- # Copyright 2026 Modilify
2
- # SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.0
3
  """MPS-specific kernels that preserve the model's mathematical operations."""
4
 
5
  from __future__ import annotations
@@ -8,10 +6,75 @@ import math
8
 
9
  import torch
10
  from torch import nn
11
- from torch.nn import functional as F
12
  from transformers.integrations.moe import ALL_EXPERTS_FUNCTIONS, _grouped_linear
13
 
14
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
15
  def lora_aware_eager_experts_forward(
16
  self: nn.Module,
17
  hidden_states: torch.Tensor,
@@ -24,7 +87,6 @@ def lora_aware_eager_experts_forward(
24
  expert_mask = expert_mask.permute(2, 1, 0)
25
  expert_hit = torch.greater(expert_mask.sum(dim=(-1, -2)), 0).nonzero()
26
 
27
- lora_dropout = getattr(self, "lora_dropout", nn.Identity())
28
  lora_scaling = getattr(self, "lora_scaling", 1.0)
29
 
30
  for expert_idx in expert_hit:
@@ -36,7 +98,7 @@ def lora_aware_eager_experts_forward(
36
  gate_up = nn.functional.linear(current_state, self.gate_up_proj[expert_idx])
37
  if hasattr(self, "lora_gate_up_a"):
38
  update = nn.functional.linear(
39
- nn.functional.linear(lora_dropout(current_state), self.lora_gate_up_a[expert_idx]),
40
  self.lora_gate_up_b[expert_idx],
41
  )
42
  gate_up = gate_up + lora_scaling * update
@@ -45,7 +107,7 @@ def lora_aware_eager_experts_forward(
45
  expert_output = nn.functional.linear(current_hidden_states, self.down_proj[expert_idx])
46
  if hasattr(self, "lora_down_a"):
47
  update = nn.functional.linear(
48
- nn.functional.linear(lora_dropout(current_hidden_states), self.lora_down_a[expert_idx]),
49
  self.lora_down_b[expert_idx],
50
  )
51
  expert_output = expert_output + lora_scaling * update
@@ -55,120 +117,15 @@ def lora_aware_eager_experts_forward(
55
  return final_hidden_states
56
 
57
 
58
- def grouped_expert_offsets(
59
- sorted_expert_ids: torch.Tensor,
60
- num_experts: int,
61
- ) -> torch.Tensor:
62
- """Build grouped-mm offsets without MPS ``histc`` host synchronization."""
63
- counts = torch.bincount(sorted_expert_ids, minlength=num_experts)[:num_experts]
64
- return torch.cumsum(counts, dim=0, dtype=torch.int32)
65
-
66
-
67
- def mps_grouped_mm_experts_forward(
68
- self: torch.nn.Module,
69
- hidden_states: torch.Tensor,
70
- top_k_index: torch.Tensor,
71
- top_k_weights: torch.Tensor,
72
- ) -> torch.Tensor:
73
- """Transformers grouped-mm MoE with histogram-free MPS routing offsets."""
74
- num_top_k = top_k_index.size(-1)
75
- num_tokens = hidden_states.size(0)
76
- hidden_dim = hidden_states.size(-1)
77
- sample_weights = top_k_weights.reshape(-1)
78
- expert_ids = top_k_index.reshape(-1)
79
-
80
- expert_ids_g, perm = torch.sort(expert_ids)
81
- selected_hidden_states_g = hidden_states[perm // num_top_k]
82
- sample_weights_g = sample_weights[perm]
83
- offsets = grouped_expert_offsets(expert_ids_g, self.num_experts)
84
-
85
- sentinel_mask = (expert_ids_g >= self.num_experts).unsqueeze(-1)
86
- expert_ids_g.clamp_(max=self.num_experts - 1)
87
- selected_hidden_states_g.masked_fill_(sentinel_mask, 0.0)
88
-
89
- gate_up_weights = self.gate_up_proj if self.has_gate else self.up_proj
90
- gate_up_biases = (
91
- self.gate_up_proj_bias[expert_ids_g]
92
- if self.has_gate and self.has_bias
93
- else self.up_proj_bias[expert_ids_g]
94
- if self.has_bias
95
- else None
96
- )
97
- projected = _grouped_linear(
98
- selected_hidden_states_g,
99
- gate_up_weights,
100
- offsets,
101
- bias=gate_up_biases,
102
- is_transposed=self.is_transposed,
103
- )
104
- projected = self._apply_gate(projected) if self.has_gate else self.act_fn(projected)
105
-
106
- down_biases = self.down_proj_bias[expert_ids_g] if self.has_bias else None
107
- projected = _grouped_linear(
108
- projected,
109
- self.down_proj,
110
- offsets,
111
- bias=down_biases,
112
- is_transposed=self.is_transposed,
113
- )
114
- weighted = projected * sample_weights_g.unsqueeze(-1)
115
- weighted.masked_fill_(sentinel_mask, 0.0)
116
-
117
- # Scatter directly back to the original top-k order. Constructing the
118
- # inverse permutation and then gathering performs the same permutation
119
- # with an extra index tensor and an extra read pass over `weighted`.
120
- reordered = torch.empty_like(weighted)
121
- reordered.index_copy_(0, perm, weighted)
122
- weighted = reordered
123
- return weighted.view(num_tokens, num_top_k, hidden_dim).sum(dim=1).to(hidden_states.dtype)
124
-
125
-
126
- ALL_EXPERTS_FUNCTIONS.register("mps_grouped_mm", mps_grouped_mm_experts_forward)
127
-
128
-
129
- class _IndependentExpertLoRA(torch.autograd.Function):
130
- """Capacity-padded reference path for independent expert LoRA."""
131
-
132
- @staticmethod
133
- def forward(
134
- ctx,
135
- inputs: torch.Tensor,
136
- weight_a: torch.Tensor,
137
- weight_b: torch.Tensor,
138
- expert_ids: torch.LongTensor,
139
- counts: torch.LongTensor,
140
- capacity: int,
141
- ) -> torch.Tensor:
142
- offsets = counts.cumsum(0)
143
- starts = offsets - counts
144
- slots = torch.arange(inputs.shape[0], device=inputs.device) - starts[expert_ids]
145
- padded = inputs.new_zeros((weight_a.shape[0], capacity, inputs.shape[-1]))
146
- padded[expert_ids, slots] = inputs
147
- low_rank = torch.bmm(padded, weight_a.transpose(1, 2))
148
- updates = torch.bmm(low_rank, weight_b.transpose(1, 2))
149
- ctx.save_for_backward(
150
- padded,
151
- low_rank,
152
- weight_a,
153
- weight_b,
154
- expert_ids,
155
- slots,
156
- )
157
- return updates[expert_ids, slots]
158
-
159
- @staticmethod
160
- def backward(ctx, grad_output: torch.Tensor):
161
- padded, low_rank, weight_a, weight_b, expert_ids, slots = ctx.saved_tensors
162
- grad_padded = grad_output.new_zeros(
163
- (weight_a.shape[0], padded.shape[1], grad_output.shape[-1])
164
- )
165
- grad_padded[expert_ids, slots] = grad_output
166
- grad_weight_b = torch.bmm(grad_padded.transpose(1, 2), low_rank)
167
- grad_low_rank = torch.bmm(grad_padded, weight_b)
168
- grad_weight_a = torch.bmm(grad_low_rank.transpose(1, 2), padded)
169
- grad_inputs_padded = torch.bmm(grad_low_rank, weight_a)
170
- grad_inputs = grad_inputs_padded[expert_ids, slots]
171
- return grad_inputs, grad_weight_a, grad_weight_b, None, None, None
172
 
173
 
174
  _INDEPENDENT_EXPERT_LORA_METAL_SOURCE = r"""
@@ -238,116 +195,6 @@ kernel void expert_lora_forward_b(
238
  output[index] = bfloat(value);
239
  }
240
 
241
- kernel void expert_lora_backward_low_rank(
242
- device const bfloat* grad_output [[buffer(0)]],
243
- device const bfloat* weight_b [[buffer(1)]],
244
- device const long* expert_ids [[buffer(2)]],
245
- device bfloat* grad_low_rank [[buffer(3)]],
246
- constant uint& row_count [[buffer(4)]],
247
- constant uint& output_dim [[buffer(5)]],
248
- uint row [[threadgroup_position_in_grid]],
249
- uint lane [[thread_index_in_threadgroup]]) {
250
- if (row >= row_count) return;
251
- const uint expert = uint(expert_ids[row]);
252
- threadgroup float partial[kThreads * kRank];
253
- for (uint rank = 0; rank < kRank; ++rank) {
254
- float value = 0.0f;
255
- const uint grad_base = row * output_dim;
256
- for (uint column = lane; column < output_dim; column += kThreads) {
257
- const uint b_index =
258
- (expert * output_dim + column) * kRank + rank;
259
- value += float(grad_output[grad_base + column])
260
- * float(weight_b[b_index]);
261
- }
262
- partial[lane * kRank + rank] = value;
263
- }
264
- threadgroup_barrier(mem_flags::mem_threadgroup);
265
- for (uint stride = kThreads / 2; stride > 0; stride >>= 1) {
266
- if (lane < stride) {
267
- for (uint rank = 0; rank < kRank; ++rank) {
268
- partial[lane * kRank + rank] +=
269
- partial[(lane + stride) * kRank + rank];
270
- }
271
- }
272
- threadgroup_barrier(mem_flags::mem_threadgroup);
273
- }
274
- if (lane == 0) {
275
- for (uint rank = 0; rank < kRank; ++rank) {
276
- grad_low_rank[row * kRank + rank] = bfloat(partial[rank]);
277
- }
278
- }
279
- }
280
-
281
- kernel void expert_lora_backward_inputs(
282
- device const bfloat* grad_low_rank [[buffer(0)]],
283
- device const bfloat* weight_a [[buffer(1)]],
284
- device const long* expert_ids [[buffer(2)]],
285
- device bfloat* grad_inputs [[buffer(3)]],
286
- constant uint& row_count [[buffer(4)]],
287
- constant uint& input_dim [[buffer(5)]],
288
- uint index [[thread_position_in_grid]]) {
289
- const uint total = row_count * input_dim;
290
- if (index >= total) return;
291
- const uint row = index / input_dim;
292
- const uint column = index - row * input_dim;
293
- const uint expert = uint(expert_ids[row]);
294
- float value = 0.0f;
295
- for (uint rank = 0; rank < kRank; ++rank) {
296
- const uint a_index =
297
- (expert * kRank + rank) * input_dim + column;
298
- value += float(grad_low_rank[row * kRank + rank])
299
- * float(weight_a[a_index]);
300
- }
301
- grad_inputs[index] = bfloat(value);
302
- }
303
-
304
- kernel void expert_lora_backward_a(
305
- device const bfloat* grad_low_rank [[buffer(0)]],
306
- device const bfloat* inputs [[buffer(1)]],
307
- device const long* offsets [[buffer(2)]],
308
- device bfloat* grad_weight_a [[buffer(3)]],
309
- constant uint& expert_count [[buffer(4)]],
310
- constant uint& input_dim [[buffer(5)]],
311
- uint index [[thread_position_in_grid]]) {
312
- const uint total = expert_count * kRank * input_dim;
313
- if (index >= total) return;
314
- const uint column = index % input_dim;
315
- const uint rank_expert = index / input_dim;
316
- const uint rank = rank_expert % kRank;
317
- const uint expert = rank_expert / kRank;
318
- const uint start = uint(offsets[expert]);
319
- const uint stop = uint(offsets[expert + 1]);
320
- float value = 0.0f;
321
- for (uint row = start; row < stop; ++row) {
322
- value += float(grad_low_rank[row * kRank + rank])
323
- * float(inputs[row * input_dim + column]);
324
- }
325
- grad_weight_a[index] = bfloat(value);
326
- }
327
-
328
- kernel void expert_lora_backward_b(
329
- device const bfloat* grad_output [[buffer(0)]],
330
- device const bfloat* low_rank [[buffer(1)]],
331
- device const long* offsets [[buffer(2)]],
332
- device bfloat* grad_weight_b [[buffer(3)]],
333
- constant uint& expert_count [[buffer(4)]],
334
- constant uint& output_dim [[buffer(5)]],
335
- uint index [[thread_position_in_grid]]) {
336
- const uint total = expert_count * output_dim * kRank;
337
- if (index >= total) return;
338
- const uint rank = index % kRank;
339
- const uint output_expert = index / kRank;
340
- const uint column = output_expert % output_dim;
341
- const uint expert = output_expert / output_dim;
342
- const uint start = uint(offsets[expert]);
343
- const uint stop = uint(offsets[expert + 1]);
344
- float value = 0.0f;
345
- for (uint row = start; row < stop; ++row) {
346
- value += float(grad_output[row * output_dim + column])
347
- * float(low_rank[row * kRank + rank]);
348
- }
349
- grad_weight_b[index] = bfloat(value);
350
- }
351
  """
352
 
353
 
@@ -363,118 +210,20 @@ def _get_independent_expert_lora_metal_library():
363
  return _independent_expert_lora_metal_library
364
 
365
 
366
- class _MetalIndependentExpertLoRA(torch.autograd.Function):
367
- """Independent rank-8 expert A/B with direct BF16 Metal forward/backward."""
368
-
369
- @staticmethod
370
- def forward(
371
- ctx,
372
- inputs: torch.Tensor,
373
- weight_a: torch.Tensor,
374
- weight_b: torch.Tensor,
375
- expert_ids: torch.LongTensor,
376
- counts: torch.LongTensor,
377
- ) -> torch.Tensor:
378
- inputs = inputs.contiguous()
379
- weight_a = weight_a.contiguous()
380
- weight_b = weight_b.contiguous()
381
- expert_ids = expert_ids.contiguous()
382
- expert_count, rank, input_dim = weight_a.shape
383
- row_count = inputs.shape[0]
384
- output_dim = weight_b.shape[1]
385
- if rank != 8:
386
- raise ValueError("The direct Metal expert LoRA kernel requires rank 8.")
387
- offsets = torch.cat(
388
- (
389
- counts.new_zeros(1, dtype=torch.int64),
390
- counts.cumsum(0, dtype=torch.int64),
391
- )
392
- ).contiguous()
393
- low_rank = inputs.new_empty((row_count, rank))
394
- output = inputs.new_empty((row_count, output_dim))
395
- library = _get_independent_expert_lora_metal_library()
396
- library.expert_lora_forward_a(
397
- inputs,
398
- weight_a,
399
- expert_ids,
400
- low_rank,
401
- row_count,
402
- input_dim,
403
- threads=(row_count * 256, 1, 1),
404
- group_size=(256, 1, 1),
405
- )
406
- library.expert_lora_forward_b(
407
- low_rank,
408
- weight_b,
409
- expert_ids,
410
- output,
411
- row_count,
412
- output_dim,
413
- threads=(row_count * output_dim, 1, 1),
414
- group_size=(256, 1, 1),
415
- )
416
- ctx.save_for_backward(
417
- inputs,
418
- low_rank,
419
- weight_a,
420
- weight_b,
421
- expert_ids,
422
- offsets,
423
- )
424
- ctx.dimensions = (expert_count, row_count, input_dim, output_dim)
425
- return output
426
-
427
- @staticmethod
428
- def backward(ctx, grad_output: torch.Tensor):
429
- inputs, low_rank, weight_a, weight_b, expert_ids, offsets = ctx.saved_tensors
430
- expert_count, row_count, input_dim, output_dim = ctx.dimensions
431
- grad_output = grad_output.contiguous()
432
- grad_low_rank = low_rank.new_empty(low_rank.shape)
433
- grad_inputs = inputs.new_empty(inputs.shape)
434
- grad_weight_a = weight_a.new_empty(weight_a.shape)
435
- grad_weight_b = weight_b.new_empty(weight_b.shape)
436
- library = _get_independent_expert_lora_metal_library()
437
- library.expert_lora_backward_low_rank(
438
- grad_output,
439
- weight_b,
440
- expert_ids,
441
- grad_low_rank,
442
- row_count,
443
- output_dim,
444
- threads=(row_count * 256, 1, 1),
445
- group_size=(256, 1, 1),
446
- )
447
- library.expert_lora_backward_inputs(
448
- grad_low_rank,
449
- weight_a,
450
- expert_ids,
451
- grad_inputs,
452
- row_count,
453
- input_dim,
454
- threads=(row_count * input_dim, 1, 1),
455
- group_size=(256, 1, 1),
456
- )
457
- library.expert_lora_backward_a(
458
- grad_low_rank,
459
- inputs,
460
- offsets,
461
- grad_weight_a,
462
- expert_count,
463
- input_dim,
464
- threads=(expert_count * 8 * input_dim, 1, 1),
465
- group_size=(256, 1, 1),
466
- )
467
- library.expert_lora_backward_b(
468
- grad_output,
469
- low_rank,
470
- offsets,
471
- grad_weight_b,
472
- expert_count,
473
- output_dim,
474
- threads=(expert_count * output_dim * 8, 1, 1),
475
- group_size=(256, 1, 1),
476
- )
477
- return grad_inputs, grad_weight_a, grad_weight_b, None, None
478
 
479
 
480
  def _independent_expert_lora(
@@ -496,14 +245,13 @@ def _independent_expert_lora(
496
  and weight_a.shape[1] == 8
497
  and hasattr(torch.mps, "compile_shader")
498
  ):
499
- return _MetalIndependentExpertLoRA.apply(
500
  inputs,
501
  weight_a,
502
  weight_b,
503
  expert_ids,
504
- counts,
505
  )
506
- return _IndependentExpertLoRA.apply(
507
  inputs,
508
  weight_a,
509
  weight_b,
@@ -553,16 +301,7 @@ def mps_segmented_experts_forward(
553
  bias=gate_up_bias,
554
  is_transposed=self.is_transposed,
555
  )
556
- if hasattr(self, "lora_gate_up_a") and self.training:
557
- gate_up_update = _independent_expert_lora(
558
- self.lora_dropout(selected_hidden_g),
559
- self.lora_gate_up_a,
560
- self.lora_gate_up_b,
561
- expert_ids_g,
562
- counts,
563
- )
564
- gate_up_g.add_(gate_up_update, alpha=self.lora_scaling)
565
- elif hasattr(self, "lora_gate_up_a"):
566
  gate_up_update = _independent_expert_lora(
567
  selected_hidden_g,
568
  self.lora_gate_up_a,
@@ -570,7 +309,7 @@ def mps_segmented_experts_forward(
570
  expert_ids_g,
571
  counts,
572
  )
573
- gate_up_g.add_(gate_up_update, alpha=self.lora_scaling)
574
  activated_g = (
575
  self._apply_gate(gate_up_g)
576
  if self.has_gate
@@ -584,16 +323,7 @@ def mps_segmented_experts_forward(
584
  bias=down_bias,
585
  is_transposed=self.is_transposed,
586
  )
587
- if hasattr(self, "lora_down_a") and self.training:
588
- down_update = _independent_expert_lora(
589
- self.lora_dropout(activated_g),
590
- self.lora_down_a,
591
- self.lora_down_b,
592
- expert_ids_g,
593
- counts,
594
- )
595
- projected_g.add_(down_update, alpha=self.lora_scaling)
596
- elif hasattr(self, "lora_down_a"):
597
  down_update = _independent_expert_lora(
598
  activated_g,
599
  self.lora_down_a,
@@ -601,18 +331,12 @@ def mps_segmented_experts_forward(
601
  expert_ids_g,
602
  counts,
603
  )
604
- projected_g.add_(down_update, alpha=self.lora_scaling)
605
  weighted_g = projected_g * sample_weights_g[:, None]
606
  weighted = torch.empty_like(weighted_g)
607
  weighted.index_copy_(0, permutation, weighted_g)
608
  return weighted.view(num_tokens, num_top_k, hidden_dim).sum(dim=1).to(hidden_states.dtype)
609
 
610
 
611
- ALL_EXPERTS_FUNCTIONS.register("mps_segmented", mps_segmented_experts_forward)
612
-
613
-
614
- __all__ = [
615
- "grouped_expert_offsets",
616
- "mps_grouped_mm_experts_forward",
617
- "mps_segmented_experts_forward",
618
- ]
 
 
 
1
  """MPS-specific kernels that preserve the model's mathematical operations."""
2
 
3
  from __future__ import annotations
 
6
 
7
  import torch
8
  from torch import nn
 
9
  from transformers.integrations.moe import ALL_EXPERTS_FUNCTIONS, _grouped_linear
10
 
11
 
12
+ def attach_expert_lora(
13
+ experts: nn.Module,
14
+ *,
15
+ rank: int,
16
+ alpha: int,
17
+ ) -> None:
18
+ """Attach packed LoRA parameters to one expert module instance."""
19
+ if rank <= 0:
20
+ raise ValueError("Expert LoRA rank must be positive.")
21
+ if alpha <= 0:
22
+ raise ValueError("Expert LoRA alpha must be positive.")
23
+ parameter_names = (
24
+ "lora_gate_up_a",
25
+ "lora_gate_up_b",
26
+ "lora_down_a",
27
+ "lora_down_b",
28
+ )
29
+ existing = tuple(hasattr(experts, name) for name in parameter_names)
30
+ if any(existing):
31
+ if not all(existing):
32
+ raise RuntimeError("Expert LoRA parameters are only partially initialized.")
33
+ return
34
+
35
+ experts.lora_scaling = float(alpha) / rank
36
+ experts.lora_gate_up_a = nn.Parameter(
37
+ torch.empty(
38
+ experts.num_experts,
39
+ rank,
40
+ experts.hidden_dim,
41
+ device=experts.gate_up_proj.device,
42
+ dtype=experts.gate_up_proj.dtype,
43
+ )
44
+ )
45
+ experts.lora_gate_up_b = nn.Parameter(
46
+ torch.zeros(
47
+ experts.num_experts,
48
+ 2 * experts.intermediate_dim,
49
+ rank,
50
+ device=experts.gate_up_proj.device,
51
+ dtype=experts.gate_up_proj.dtype,
52
+ )
53
+ )
54
+ experts.lora_down_a = nn.Parameter(
55
+ torch.empty(
56
+ experts.num_experts,
57
+ rank,
58
+ experts.intermediate_dim,
59
+ device=experts.down_proj.device,
60
+ dtype=experts.down_proj.dtype,
61
+ )
62
+ )
63
+ experts.lora_down_b = nn.Parameter(
64
+ torch.zeros(
65
+ experts.num_experts,
66
+ experts.hidden_dim,
67
+ rank,
68
+ device=experts.down_proj.device,
69
+ dtype=experts.down_proj.dtype,
70
+ )
71
+ )
72
+ nn.init.kaiming_uniform_(experts.lora_gate_up_a, a=math.sqrt(5))
73
+ nn.init.kaiming_uniform_(experts.lora_down_a, a=math.sqrt(5))
74
+ experts.gate_up_proj.requires_grad_(False)
75
+ experts.down_proj.requires_grad_(False)
76
+
77
+
78
  def lora_aware_eager_experts_forward(
79
  self: nn.Module,
80
  hidden_states: torch.Tensor,
 
87
  expert_mask = expert_mask.permute(2, 1, 0)
88
  expert_hit = torch.greater(expert_mask.sum(dim=(-1, -2)), 0).nonzero()
89
 
 
90
  lora_scaling = getattr(self, "lora_scaling", 1.0)
91
 
92
  for expert_idx in expert_hit:
 
98
  gate_up = nn.functional.linear(current_state, self.gate_up_proj[expert_idx])
99
  if hasattr(self, "lora_gate_up_a"):
100
  update = nn.functional.linear(
101
+ nn.functional.linear(current_state, self.lora_gate_up_a[expert_idx]),
102
  self.lora_gate_up_b[expert_idx],
103
  )
104
  gate_up = gate_up + lora_scaling * update
 
107
  expert_output = nn.functional.linear(current_hidden_states, self.down_proj[expert_idx])
108
  if hasattr(self, "lora_down_a"):
109
  update = nn.functional.linear(
110
+ nn.functional.linear(current_hidden_states, self.lora_down_a[expert_idx]),
111
  self.lora_down_b[expert_idx],
112
  )
113
  expert_output = expert_output + lora_scaling * update
 
117
  return final_hidden_states
118
 
119
 
120
+ def _padded_expert_lora(inputs: torch.Tensor, weight_a: torch.Tensor, weight_b: torch.Tensor, expert_ids: torch.LongTensor, counts: torch.LongTensor, capacity: int) -> torch.Tensor:
121
+ offsets = counts.cumsum(0)
122
+ starts = offsets - counts
123
+ slots = torch.arange(inputs.shape[0], device=inputs.device) - starts[expert_ids]
124
+ padded = inputs.new_zeros((weight_a.shape[0], capacity, inputs.shape[-1]))
125
+ padded[expert_ids, slots] = inputs
126
+ low_rank = torch.bmm(padded, weight_a.transpose(1, 2))
127
+ updates = torch.bmm(low_rank, weight_b.transpose(1, 2))
128
+ return updates[expert_ids, slots]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
129
 
130
 
131
  _INDEPENDENT_EXPERT_LORA_METAL_SOURCE = r"""
 
195
  output[index] = bfloat(value);
196
  }
197
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
198
  """
199
 
200
 
 
210
  return _independent_expert_lora_metal_library
211
 
212
 
213
+ def _metal_expert_lora(inputs: torch.Tensor, weight_a: torch.Tensor, weight_b: torch.Tensor, expert_ids: torch.LongTensor) -> torch.Tensor:
214
+ inputs = inputs.contiguous()
215
+ weight_a = weight_a.contiguous()
216
+ weight_b = weight_b.contiguous()
217
+ expert_ids = expert_ids.contiguous()
218
+ expert_count, rank, input_dim = weight_a.shape
219
+ row_count = inputs.shape[0]
220
+ output_dim = weight_b.shape[1]
221
+ low_rank = inputs.new_empty((row_count, rank))
222
+ output = inputs.new_empty((row_count, output_dim))
223
+ library = _get_independent_expert_lora_metal_library()
224
+ library.expert_lora_forward_a(inputs, weight_a, expert_ids, low_rank, row_count, input_dim, threads=(row_count * 256, 1, 1), group_size=(256, 1, 1))
225
+ library.expert_lora_forward_b(low_rank, weight_b, expert_ids, output, row_count, output_dim, threads=(row_count * output_dim, 1, 1), group_size=(256, 1, 1))
226
+ return output
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
227
 
228
 
229
  def _independent_expert_lora(
 
245
  and weight_a.shape[1] == 8
246
  and hasattr(torch.mps, "compile_shader")
247
  ):
248
+ return _metal_expert_lora(
249
  inputs,
250
  weight_a,
251
  weight_b,
252
  expert_ids,
 
253
  )
254
+ return _padded_expert_lora(
255
  inputs,
256
  weight_a,
257
  weight_b,
 
301
  bias=gate_up_bias,
302
  is_transposed=self.is_transposed,
303
  )
304
+ if hasattr(self, "lora_gate_up_a"):
 
 
 
 
 
 
 
 
 
305
  gate_up_update = _independent_expert_lora(
306
  selected_hidden_g,
307
  self.lora_gate_up_a,
 
309
  expert_ids_g,
310
  counts,
311
  )
312
+ gate_up_g = gate_up_g + gate_up_update * self.lora_scaling
313
  activated_g = (
314
  self._apply_gate(gate_up_g)
315
  if self.has_gate
 
323
  bias=down_bias,
324
  is_transposed=self.is_transposed,
325
  )
326
+ if hasattr(self, "lora_down_a"):
 
 
 
 
 
 
 
 
 
327
  down_update = _independent_expert_lora(
328
  activated_g,
329
  self.lora_down_a,
 
331
  expert_ids_g,
332
  counts,
333
  )
334
+ projected_g = projected_g + down_update * self.lora_scaling
335
  weighted_g = projected_g * sample_weights_g[:, None]
336
  weighted = torch.empty_like(weighted_g)
337
  weighted.index_copy_(0, permutation, weighted_g)
338
  return weighted.view(num_tokens, num_top_k, hidden_dim).sum(dim=1).to(hidden_states.dtype)
339
 
340
 
341
+ def register_mps_backends() -> None:
342
+ ALL_EXPERTS_FUNCTIONS.register("mps_segmented", mps_segmented_experts_forward)
 
 
 
 
 
 
vocab_ops.py CHANGED
@@ -1,16 +1,10 @@
1
- # Copyright 2026 Modilify
2
- # SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.0
3
- """Exact memory-bounded vocabulary projection and sampling."""
4
-
5
  from __future__ import annotations
6
-
7
  import math
8
  from collections.abc import Sequence
9
-
10
  import torch
11
  from torch.nn import functional as F
12
 
13
-
14
  def _stable_max(values: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
15
  """Return ``max`` indices that stay in-range on MPS NaN/-inf rows."""
16
 
@@ -20,322 +14,6 @@ def _stable_max(values: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
20
  return best, index
21
  return best, index.clamp(0, width - 1)
22
 
23
-
24
- class _ChunkedVocabStatistics(torch.autograd.Function):
25
- """One projection pass for CE and exact Gumbel-max categorical sampling."""
26
-
27
- @staticmethod
28
- def forward(
29
- ctx,
30
- hidden: torch.Tensor,
31
- weight: torch.Tensor,
32
- labels: torch.LongTensor,
33
- softcap: float,
34
- temperature: float,
35
- chunk_size: int,
36
- ) -> tuple[
37
- torch.Tensor,
38
- torch.Tensor,
39
- torch.LongTensor,
40
- torch.Tensor,
41
- torch.Tensor,
42
- torch.LongTensor,
43
- torch.Tensor,
44
- ]:
45
- if hidden.ndim != 2 or weight.ndim != 2:
46
- raise ValueError("Chunked vocabulary tensors must be matrices.")
47
- if hidden.shape[1] != weight.shape[1]:
48
- raise ValueError("Hidden and vocabulary projection dimensions differ.")
49
- if temperature <= 0 or chunk_size <= 0:
50
- raise ValueError("Temperature and vocabulary chunk size must be positive.")
51
- labels = labels.to(device=hidden.device, dtype=torch.long)
52
- if labels.shape != (hidden.shape[0],):
53
- raise ValueError("Labels must contain one value per hidden row.")
54
-
55
- rows = hidden.shape[0]
56
- row_indices = torch.arange(rows, device=hidden.device)
57
- valid = labels.ge(0)
58
- raw_log_z = torch.full(
59
- (rows,), -torch.inf, device=hidden.device, dtype=torch.float32
60
- )
61
- sample_log_z = torch.full_like(raw_log_z, -torch.inf)
62
- gold_score = torch.zeros_like(raw_log_z)
63
- best_gumbel = torch.full_like(raw_log_z, -torch.inf)
64
- selected_score = torch.zeros_like(raw_log_z)
65
- selected = torch.zeros(rows, device=hidden.device, dtype=torch.long)
66
- greedy_score = torch.full_like(raw_log_z, -torch.inf)
67
- greedy = torch.zeros(rows, device=hidden.device, dtype=torch.long)
68
- moment_max = torch.full_like(raw_log_z, -torch.inf)
69
- moment_sum = torch.zeros_like(raw_log_z)
70
- moment_weighted = torch.zeros_like(raw_log_z)
71
- vocab_size = weight.shape[0]
72
-
73
- with torch.no_grad():
74
- for start in range(0, vocab_size, int(chunk_size)):
75
- stop = min(start + int(chunk_size), vocab_size)
76
- raw = F.linear(hidden, weight[start:stop])
77
- scores = (
78
- torch.tanh(raw.float() / float(softcap))
79
- * float(softcap)
80
- )
81
- raw_log_z = torch.logaddexp(
82
- raw_log_z,
83
- torch.logsumexp(scores, dim=-1),
84
- )
85
- sample_scores = scores / float(temperature)
86
- sample_log_z = torch.logaddexp(
87
- sample_log_z,
88
- torch.logsumexp(sample_scores, dim=-1),
89
- )
90
-
91
- in_chunk = valid & labels.ge(start) & labels.lt(stop)
92
- local_gold = (labels - start).clamp(0, stop - start - 1)
93
- gold_score = torch.where(
94
- in_chunk,
95
- scores[row_indices, local_gold],
96
- gold_score,
97
- )
98
-
99
- uniform = torch.rand(
100
- sample_scores.shape,
101
- device=sample_scores.device,
102
- dtype=torch.float32,
103
- ).clamp_(
104
- min=torch.finfo(torch.float32).tiny,
105
- max=1.0 - torch.finfo(torch.float32).eps,
106
- )
107
- gumbel_scores = sample_scores - torch.log(-torch.log(uniform))
108
- chunk_best, chunk_index = _stable_max(gumbel_scores)
109
- replace_best = chunk_best.gt(best_gumbel)
110
- candidate_score = sample_scores.gather(
111
- 1, chunk_index[:, None]
112
- ).squeeze(-1)
113
- best_gumbel = torch.maximum(best_gumbel, chunk_best)
114
- selected = torch.where(
115
- replace_best,
116
- chunk_index + start,
117
- selected,
118
- )
119
- selected_score = torch.where(
120
- replace_best,
121
- candidate_score,
122
- selected_score,
123
- )
124
-
125
- chunk_max, chunk_argmax = _stable_max(sample_scores)
126
- replace_greedy = chunk_max.gt(greedy_score)
127
- greedy_score = torch.maximum(greedy_score, chunk_max)
128
- greedy = torch.where(
129
- replace_greedy,
130
- chunk_argmax + start,
131
- greedy,
132
- )
133
- shifted = torch.exp(sample_scores - chunk_max[:, None])
134
- chunk_sum = shifted.sum(dim=-1)
135
- chunk_weighted = (shifted * sample_scores).sum(dim=-1)
136
- merged_max = torch.maximum(moment_max, chunk_max)
137
- previous_scale = torch.exp(moment_max - merged_max)
138
- chunk_scale = torch.exp(chunk_max - merged_max)
139
- moment_sum = (
140
- moment_sum * previous_scale + chunk_sum * chunk_scale
141
- )
142
- moment_weighted = (
143
- moment_weighted * previous_scale
144
- + chunk_weighted * chunk_scale
145
- )
146
- moment_max = merged_max
147
-
148
- raw_nll = torch.where(
149
- valid,
150
- raw_log_z - gold_score,
151
- torch.zeros_like(raw_log_z),
152
- )
153
- temperature_nll = torch.where(
154
- valid,
155
- sample_log_z - gold_score / float(temperature),
156
- torch.zeros_like(sample_log_z),
157
- )
158
- confidence = torch.exp(selected_score - sample_log_z).clamp_(0.0, 1.0)
159
- greedy_confidence = torch.exp(greedy_score - sample_log_z).clamp_(0.0, 1.0)
160
- entropy = sample_log_z - moment_weighted / moment_sum.clamp_min(
161
- torch.finfo(torch.float32).tiny
162
- )
163
-
164
- ctx.save_for_backward(
165
- hidden,
166
- weight,
167
- labels,
168
- selected,
169
- raw_log_z,
170
- sample_log_z,
171
- confidence,
172
- entropy,
173
- )
174
- ctx.softcap = float(softcap)
175
- ctx.temperature = float(temperature)
176
- ctx.chunk_size = int(chunk_size)
177
- ctx.mark_non_differentiable(selected, greedy, greedy_confidence)
178
- return (
179
- raw_nll,
180
- temperature_nll,
181
- selected,
182
- confidence,
183
- entropy,
184
- greedy,
185
- greedy_confidence,
186
- )
187
-
188
- @staticmethod
189
- def backward(
190
- ctx,
191
- grad_raw_nll: torch.Tensor | None,
192
- grad_temperature_nll: torch.Tensor | None,
193
- grad_selected: torch.Tensor | None,
194
- grad_confidence: torch.Tensor | None,
195
- grad_entropy: torch.Tensor | None,
196
- grad_greedy: torch.Tensor | None,
197
- grad_greedy_confidence: torch.Tensor | None,
198
- ):
199
- del grad_selected, grad_greedy, grad_greedy_confidence
200
- (
201
- hidden,
202
- weight,
203
- labels,
204
- selected,
205
- raw_log_z,
206
- sample_log_z,
207
- confidence,
208
- entropy,
209
- ) = ctx.saved_tensors
210
- grad_raw_nll = (
211
- torch.zeros_like(raw_log_z)
212
- if grad_raw_nll is None
213
- else grad_raw_nll.float()
214
- )
215
- grad_temperature_nll = (
216
- torch.zeros_like(sample_log_z)
217
- if grad_temperature_nll is None
218
- else grad_temperature_nll.float()
219
- )
220
- grad_confidence = (
221
- torch.zeros_like(confidence)
222
- if grad_confidence is None
223
- else grad_confidence.float()
224
- )
225
- valid = labels.ge(0)
226
- row_indices = torch.arange(hidden.shape[0], device=hidden.device)
227
- grad_hidden = torch.zeros_like(hidden, dtype=torch.float32)
228
- confidence_scale = (
229
- grad_confidence * confidence / ctx.temperature
230
- )
231
- for start in range(0, weight.shape[0], ctx.chunk_size):
232
- stop = min(start + ctx.chunk_size, weight.shape[0])
233
- raw = F.linear(hidden, weight[start:stop])
234
- scores = torch.tanh(raw.float() / ctx.softcap) * ctx.softcap
235
- raw_probability = torch.exp(scores - raw_log_z[:, None])
236
- sample_probability = torch.exp(
237
- scores / ctx.temperature - sample_log_z[:, None]
238
- )
239
- score_gradient = (
240
- raw_probability * (grad_raw_nll * valid)[:, None]
241
- + sample_probability
242
- * (grad_temperature_nll * valid)[:, None]
243
- / ctx.temperature
244
- - sample_probability * confidence_scale[:, None]
245
- )
246
- if grad_entropy is not None:
247
- grad_entropy_f = grad_entropy.float()
248
- sample_scores = scores / ctx.temperature
249
- entropy_grad = -(
250
- grad_entropy_f / ctx.temperature
251
- )[:, None] * sample_probability * (
252
- sample_scores - sample_log_z[:, None] + entropy[:, None]
253
- )
254
- score_gradient = score_gradient + entropy_grad
255
-
256
- gold_in_chunk = valid & labels.ge(start) & labels.lt(stop)
257
- gold_local = (labels - start).clamp(0, stop - start - 1)
258
- score_gradient[row_indices, gold_local] -= (
259
- grad_raw_nll * gold_in_chunk
260
- + grad_temperature_nll
261
- * gold_in_chunk
262
- / ctx.temperature
263
- )
264
- selected_in_chunk = selected.ge(start) & selected.lt(stop)
265
- selected_local = (selected - start).clamp(0, stop - start - 1)
266
- score_gradient[row_indices, selected_local] += (
267
- confidence_scale * selected_in_chunk
268
- )
269
- score_gradient *= 1.0 - (scores / ctx.softcap).square()
270
- # The B16 shape feeds 4,096 rows through a 32,768-wide vocabulary
271
- # chunk. MPS BF16 GEMM can return NaNs for this backward-only
272
- # projection even when every score gradient is finite. Accumulate
273
- # the exact VJP in FP32; casting the finished hidden gradient back
274
- # to the model dtype happens only once below.
275
- chunk_gradient = F.linear(
276
- score_gradient,
277
- weight[start:stop].float().transpose(0, 1),
278
- )
279
- grad_hidden.add_(chunk_gradient)
280
-
281
- return (
282
- grad_hidden.to(hidden.dtype),
283
- None,
284
- None,
285
- None,
286
- None,
287
- None,
288
- )
289
-
290
-
291
- def chunked_vocab_training_statistics(
292
- hidden: torch.Tensor,
293
- weight: torch.Tensor,
294
- labels: torch.LongTensor,
295
- *,
296
- softcap: float,
297
- temperature: float,
298
- chunk_size: int,
299
- ) -> tuple[
300
- torch.Tensor,
301
- torch.Tensor,
302
- torch.LongTensor,
303
- torch.Tensor,
304
- torch.Tensor,
305
- torch.LongTensor,
306
- torch.Tensor,
307
- ]:
308
- """Return exact losses plus sampled and greedy proposal statistics."""
309
-
310
- shape = labels.shape
311
- outputs = _ChunkedVocabStatistics.apply(
312
- hidden.reshape(-1, hidden.shape[-1]),
313
- weight,
314
- labels.reshape(-1),
315
- float(softcap),
316
- float(temperature),
317
- int(chunk_size),
318
- )
319
- (
320
- raw_nll,
321
- temperature_nll,
322
- proposal,
323
- confidence,
324
- entropy,
325
- greedy,
326
- greedy_confidence,
327
- ) = outputs
328
- return (
329
- raw_nll.view(shape),
330
- temperature_nll.view(shape),
331
- proposal.view(shape),
332
- confidence.view(shape),
333
- entropy.view(shape),
334
- greedy.view(shape),
335
- greedy_confidence.view(shape),
336
- )
337
-
338
-
339
  @torch.no_grad()
340
  def chunked_vocab_statistics(
341
  hidden: torch.Tensor,
@@ -347,6 +25,8 @@ def chunked_vocab_statistics(
347
  repetition_token_mask: torch.BoolTensor | None = None,
348
  repetition_penalty: float = 1.0,
349
  sampling_generators: Sequence[torch.Generator] | None = None,
 
 
350
  ) -> tuple[
351
  torch.LongTensor,
352
  torch.Tensor,
@@ -359,71 +39,77 @@ def chunked_vocab_statistics(
359
  if not math.isfinite(repetition_penalty) or repetition_penalty <= 0:
360
  raise ValueError("`repetition_penalty` must be a finite positive number.")
361
 
362
- # Preserve the original custom-autograd path exactly when disabled. In
363
- # particular, this keeps the same chunk-local RNG draws and default output.
364
- if repetition_penalty != 1.0 or sampling_generators is not None:
365
- if hidden.ndim < 2 or weight.ndim != 2:
366
- raise ValueError("Chunked vocabulary tensors must have matrix features.")
367
- if hidden.shape[-1] != weight.shape[1]:
368
- raise ValueError("Hidden and vocabulary projection dimensions differ.")
369
- if temperature <= 0 or chunk_size <= 0:
370
- raise ValueError("Temperature and vocabulary chunk size must be positive.")
371
- expected_mask_shape = (hidden.shape[0], weight.shape[0])
372
- if repetition_penalty != 1.0:
373
- if repetition_token_mask is None or repetition_token_mask.shape != expected_mask_shape:
374
- raise ValueError(
375
- "`repetition_token_mask` must have shape [batch, vocabulary]."
376
- )
377
- if repetition_token_mask.dtype != torch.bool:
378
- raise ValueError("`repetition_token_mask` must be a boolean tensor.")
379
- if sampling_generators is not None and len(sampling_generators) != hidden.shape[0]:
380
  raise ValueError(
381
- "`sampling_generators` must contain one generator per batch row."
382
  )
383
-
384
- output_shape = hidden.shape[:-1]
385
- flat_hidden = hidden.reshape(-1, hidden.shape[-1])
386
- rows_per_batch = math.prod(hidden.shape[1:-1])
387
- batch_indices = torch.arange(
388
- flat_hidden.shape[0], device=hidden.device
389
- ).div(rows_per_batch, rounding_mode="floor")
390
- sample_log_z = torch.full(
391
- (flat_hidden.shape[0],),
392
- -torch.inf,
393
- device=hidden.device,
394
- dtype=torch.float32,
395
- )
396
- best_gumbel = torch.full_like(sample_log_z, -torch.inf)
397
- selected_score = torch.zeros_like(sample_log_z)
398
- selected = torch.zeros(
399
- flat_hidden.shape[0], device=hidden.device, dtype=torch.long
400
  )
401
- greedy_score = torch.full_like(sample_log_z, -torch.inf)
402
- greedy = torch.zeros_like(selected)
403
- moment_max = torch.full_like(sample_log_z, -torch.inf)
404
- moment_sum = torch.zeros_like(sample_log_z)
405
- moment_weighted = torch.zeros_like(sample_log_z)
406
 
407
- for start in range(0, weight.shape[0], int(chunk_size)):
408
- stop = min(start + int(chunk_size), weight.shape[0])
409
- raw = F.linear(flat_hidden, weight[start:stop])
410
- scores = torch.tanh(raw.float() / float(softcap)) * float(softcap)
411
- if repetition_penalty != 1.0:
412
- assert repetition_token_mask is not None
413
- seen = repetition_token_mask[:, start:stop].index_select(
414
- 0, batch_indices
415
- )
416
- penalized = torch.where(
417
- scores < 0,
418
- scores * float(repetition_penalty),
419
- scores / float(repetition_penalty),
420
- )
421
- scores = torch.where(seen, penalized, scores)
422
- sample_scores = scores / float(temperature)
423
- sample_log_z = torch.logaddexp(
424
- sample_log_z, torch.logsumexp(sample_scores, dim=-1)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
425
  )
 
 
 
 
 
426
 
 
 
 
 
 
 
427
  if sampling_generators is None:
428
  uniform = torch.rand(
429
  sample_scores.shape,
@@ -431,8 +117,6 @@ def chunked_vocab_statistics(
431
  dtype=torch.float32,
432
  )
433
  else:
434
- # Drawing each request from its own generator makes sampling
435
- # invariant to dynamic admission, slot changes, and batch order.
436
  per_request_shape = (
437
  rows_per_batch,
438
  sample_scores.shape[-1],
@@ -465,61 +149,75 @@ def chunked_vocab_statistics(
465
  replace_best, candidate_score, selected_score
466
  )
467
 
468
- chunk_max, chunk_argmax = _stable_max(sample_scores)
469
- replace_greedy = chunk_max.gt(greedy_score)
470
- greedy_score = torch.maximum(greedy_score, chunk_max)
471
- greedy = torch.where(replace_greedy, chunk_argmax + start, greedy)
472
- shifted = torch.exp(sample_scores - chunk_max[:, None])
473
- chunk_sum = shifted.sum(dim=-1)
474
- chunk_weighted = (shifted * sample_scores).sum(dim=-1)
475
- merged_max = torch.maximum(moment_max, chunk_max)
476
- previous_scale = torch.exp(moment_max - merged_max)
477
- chunk_scale = torch.exp(chunk_max - merged_max)
478
- moment_sum = moment_sum * previous_scale + chunk_sum * chunk_scale
479
- moment_weighted = (
480
- moment_weighted * previous_scale + chunk_weighted * chunk_scale
481
- )
482
- moment_max = merged_max
483
-
484
- confidence = torch.exp(selected_score - sample_log_z).clamp_(0.0, 1.0)
485
- greedy_confidence = torch.exp(greedy_score - sample_log_z).clamp_(0.0, 1.0)
486
- entropy = sample_log_z - moment_weighted / moment_sum.clamp_min(
487
- torch.finfo(torch.float32).tiny
488
  )
489
- return (
490
- selected.view(output_shape),
491
- confidence.view(output_shape),
492
- entropy.view(output_shape),
493
- greedy.view(output_shape),
494
- greedy_confidence.view(output_shape),
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
495
  )
496
-
497
- labels = torch.full(
498
- hidden.shape[:-1],
499
- -100,
500
- device=hidden.device,
501
- dtype=torch.long,
 
 
 
502
  )
503
- (
504
- _,
505
- _,
506
- proposal,
507
- confidence,
508
- entropy,
509
- greedy,
510
- greedy_confidence,
511
- ) = chunked_vocab_training_statistics(
512
- hidden,
513
- weight,
514
- labels,
515
- softcap=softcap,
516
- temperature=temperature,
517
- chunk_size=chunk_size,
518
  )
519
- return proposal, confidence, entropy, greedy, greedy_confidence
520
-
521
-
522
- __all__ = [
523
- "chunked_vocab_statistics",
524
- "chunked_vocab_training_statistics",
525
- ]
 
1
+ """Memory-bounded vocabulary projection and exact inference sampling."""
 
 
 
2
  from __future__ import annotations
 
3
  import math
4
  from collections.abc import Sequence
 
5
  import torch
6
  from torch.nn import functional as F
7
 
 
8
  def _stable_max(values: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
9
  """Return ``max`` indices that stay in-range on MPS NaN/-inf rows."""
10
 
 
14
  return best, index
15
  return best, index.clamp(0, width - 1)
16
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
17
  @torch.no_grad()
18
  def chunked_vocab_statistics(
19
  hidden: torch.Tensor,
 
25
  repetition_token_mask: torch.BoolTensor | None = None,
26
  repetition_penalty: float = 1.0,
27
  sampling_generators: Sequence[torch.Generator] | None = None,
28
+ top_k: int | None = 40,
29
+ min_p: float | None = 0.05,
30
  ) -> tuple[
31
  torch.LongTensor,
32
  torch.Tensor,
 
39
  if not math.isfinite(repetition_penalty) or repetition_penalty <= 0:
40
  raise ValueError("`repetition_penalty` must be a finite positive number.")
41
 
42
+ if hidden.ndim < 2 or weight.ndim != 2:
43
+ raise ValueError("Chunked vocabulary tensors must have matrix features.")
44
+ if hidden.shape[-1] != weight.shape[1]:
45
+ raise ValueError("Hidden and vocabulary projection dimensions differ.")
46
+ if temperature <= 0 or chunk_size <= 0:
47
+ raise ValueError("Temperature and vocabulary chunk size must be positive.")
48
+ expected_mask_shape = (hidden.shape[0], weight.shape[0])
49
+ if repetition_penalty != 1.0:
50
+ if repetition_token_mask is None or repetition_token_mask.shape != expected_mask_shape:
 
 
 
 
 
 
 
 
 
51
  raise ValueError(
52
+ "`repetition_token_mask` must have shape [batch, vocabulary]."
53
  )
54
+ if repetition_token_mask.dtype != torch.bool:
55
+ raise ValueError("`repetition_token_mask` must be a boolean tensor.")
56
+ if sampling_generators is not None and len(sampling_generators) != hidden.shape[0]:
57
+ raise ValueError(
58
+ "`sampling_generators` must contain one generator per batch row."
 
 
 
 
 
 
 
 
 
 
 
 
59
  )
 
 
 
 
 
60
 
61
+ output_shape = hidden.shape[:-1]
62
+ flat_hidden = hidden.reshape(-1, hidden.shape[-1])
63
+ rows_per_batch = math.prod(hidden.shape[1:-1])
64
+ batch_indices = torch.arange(
65
+ flat_hidden.shape[0], device=hidden.device
66
+ ).div(rows_per_batch, rounding_mode="floor")
67
+ sample_log_z = torch.full(
68
+ (flat_hidden.shape[0],),
69
+ -torch.inf,
70
+ device=hidden.device,
71
+ dtype=torch.float32,
72
+ )
73
+ best_gumbel = torch.full_like(sample_log_z, -torch.inf)
74
+ selected_score = torch.zeros_like(sample_log_z)
75
+ selected = torch.zeros(
76
+ flat_hidden.shape[0], device=hidden.device, dtype=torch.long
77
+ )
78
+ greedy_score = torch.full_like(sample_log_z, -torch.inf)
79
+ greedy = torch.zeros_like(selected)
80
+ moment_max = torch.full_like(sample_log_z, -torch.inf)
81
+ moment_sum = torch.zeros_like(sample_log_z)
82
+ moment_weighted = torch.zeros_like(sample_log_z)
83
+
84
+ use_constrained = top_k is not None and top_k > 0
85
+ chunk_cand_scores = []
86
+ chunk_cand_tokens = []
87
+
88
+ for start in range(0, weight.shape[0], int(chunk_size)):
89
+ stop = min(start + int(chunk_size), weight.shape[0])
90
+ raw = F.linear(flat_hidden, weight[start:stop])
91
+ scores = torch.tanh(raw.float() / float(softcap)) * float(softcap)
92
+ if repetition_penalty != 1.0:
93
+ seen = repetition_token_mask[:, start:stop].index_select(
94
+ 0, batch_indices
95
+ )
96
+ penalized = torch.where(
97
+ scores < 0,
98
+ scores * float(repetition_penalty),
99
+ scores / float(repetition_penalty),
100
  )
101
+ scores = torch.where(seen, penalized, scores)
102
+ sample_scores = scores / float(temperature)
103
+ sample_log_z = torch.logaddexp(
104
+ sample_log_z, torch.logsumexp(sample_scores, dim=-1)
105
+ )
106
 
107
+ if use_constrained:
108
+ k = min(int(top_k), stop - start)
109
+ c_score, c_idx = torch.topk(sample_scores, k=k, dim=-1)
110
+ chunk_cand_scores.append(c_score)
111
+ chunk_cand_tokens.append(c_idx + start)
112
+ else:
113
  if sampling_generators is None:
114
  uniform = torch.rand(
115
  sample_scores.shape,
 
117
  dtype=torch.float32,
118
  )
119
  else:
 
 
120
  per_request_shape = (
121
  rows_per_batch,
122
  sample_scores.shape[-1],
 
149
  replace_best, candidate_score, selected_score
150
  )
151
 
152
+ chunk_max, chunk_argmax = _stable_max(sample_scores)
153
+ replace_greedy = chunk_max.gt(greedy_score)
154
+ greedy_score = torch.maximum(greedy_score, chunk_max)
155
+ greedy = torch.where(replace_greedy, chunk_argmax + start, greedy)
156
+ shifted = torch.exp(sample_scores - chunk_max[:, None])
157
+ chunk_sum = shifted.sum(dim=-1)
158
+ chunk_weighted = (shifted * sample_scores).sum(dim=-1)
159
+ merged_max = torch.maximum(moment_max, chunk_max)
160
+ previous_scale = torch.exp(moment_max - merged_max)
161
+ chunk_scale = torch.exp(chunk_max - merged_max)
162
+ moment_sum = moment_sum * previous_scale + chunk_sum * chunk_scale
163
+ moment_weighted = (
164
+ moment_weighted * previous_scale + chunk_weighted * chunk_scale
 
 
 
 
 
 
 
165
  )
166
+ moment_max = merged_max
167
+
168
+ if use_constrained:
169
+ all_scores = torch.cat(chunk_cand_scores, dim=-1)
170
+ all_tokens = torch.cat(chunk_cand_tokens, dim=-1)
171
+ global_k = min(int(top_k), all_scores.shape[-1])
172
+ cand_scores, global_idx = torch.topk(all_scores, k=global_k, dim=-1)
173
+ cand_tokens = all_tokens.gather(dim=-1, index=global_idx)
174
+ if min_p is not None and min_p > 0:
175
+ thresh = greedy_score + math.log(float(min_p))
176
+ valid_mask = cand_scores >= thresh[:, None]
177
+ eligible = cand_scores.masked_fill(~valid_mask, -float("inf"))
178
+ else:
179
+ eligible = cand_scores
180
+ if sampling_generators is None:
181
+ uniform = torch.rand(
182
+ eligible.shape,
183
+ device=eligible.device,
184
+ dtype=torch.float32,
185
+ )
186
+ else:
187
+ per_request_shape = (
188
+ rows_per_batch,
189
+ eligible.shape[-1],
190
+ )
191
+ uniform = torch.cat(
192
+ [
193
+ torch.rand(
194
+ per_request_shape,
195
+ device=eligible.device,
196
+ dtype=torch.float32,
197
+ generator=generator,
198
+ )
199
+ for generator in sampling_generators
200
+ ],
201
+ dim=0,
202
+ )
203
+ uniform = uniform.clamp_(
204
+ min=torch.finfo(torch.float32).tiny,
205
+ max=1.0 - torch.finfo(torch.float32).eps,
206
  )
207
+ gumbel = eligible - torch.log(-torch.log(uniform))
208
+ _, chosen_idx = _stable_max(gumbel)
209
+ selected = cand_tokens.gather(dim=-1, index=chosen_idx[:, None]).squeeze(-1)
210
+ selected_score = cand_scores.gather(dim=-1, index=chosen_idx[:, None]).squeeze(-1)
211
+
212
+ confidence = torch.exp(selected_score - sample_log_z).clamp_(0.0, 1.0)
213
+ greedy_confidence = torch.exp(greedy_score - sample_log_z).clamp_(0.0, 1.0)
214
+ entropy = sample_log_z - moment_weighted / moment_sum.clamp_min(
215
+ torch.finfo(torch.float32).tiny
216
  )
217
+ return (
218
+ selected.view(output_shape),
219
+ confidence.view(output_shape),
220
+ entropy.view(output_shape),
221
+ greedy.view(output_shape),
222
+ greedy_confidence.view(output_shape),
 
 
 
 
 
 
 
 
 
223
  )