Text Generation
Transformers
Safetensors
modilify_mk2
diffusion
mixture-of-experts
trust-remote-code
conversational
custom_code
Instructions to use modilify/Modilify-Mk2-preview with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use modilify/Modilify-Mk2-preview with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="modilify/Modilify-Mk2-preview", trust_remote_code=True) messages = [ {"role": "user", "content": "Who are you?"}, ] pipe(messages)# pip install -U transformers accelerate # Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("modilify/Modilify-Mk2-preview", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use modilify/Modilify-Mk2-preview with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "modilify/Modilify-Mk2-preview" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "modilify/Modilify-Mk2-preview", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker
docker model run hf.co/modilify/Modilify-Mk2-preview
- SGLang
How to use modilify/Modilify-Mk2-preview with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "modilify/Modilify-Mk2-preview" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "modilify/Modilify-Mk2-preview", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "modilify/Modilify-Mk2-preview" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "modilify/Modilify-Mk2-preview", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }' - Docker Model Runner
How to use modilify/Modilify-Mk2-preview with Docker Model Runner:
docker model run hf.co/modilify/Modilify-Mk2-preview
Publish Modilify Mk2 Preview step 1250 schema25
Browse files- .gitattributes +0 -1
- NOTICE.md +7 -6
- README.md +30 -248
- chat_template.jinja +4 -4
- commit_policy.py +40 -30
- config.json +38 -81
- configuration_modilify_mk2.py +73 -297
- continuous_batching.py +64 -154
- gdn2_memory.py +112 -0
- gdn2_trajectory.py +86 -0
- generation_config.json +1 -15
- generation_modilify_mk2.py +118 -461
- latent_deliberation.py +247 -1757
- lora.py +96 -0
- model-00001-of-00032.safetensors +3 -0
- model-00002-of-00032.safetensors +3 -0
- model-00003-of-00032.safetensors +3 -0
- model-00004-of-00032.safetensors +3 -0
- model-00005-of-00032.safetensors +3 -0
- model-00006-of-00032.safetensors +3 -0
- model-00007-of-00032.safetensors +3 -0
- model-00008-of-00032.safetensors +3 -0
- model-00009-of-00032.safetensors +3 -0
- model-00010-of-00032.safetensors +3 -0
- model-00011-of-00032.safetensors +3 -0
- model-00012-of-00032.safetensors +3 -0
- model-00013-of-00032.safetensors +3 -0
- model-00014-of-00032.safetensors +3 -0
- model-00015-of-00032.safetensors +3 -0
- model-00016-of-00032.safetensors +3 -0
- model-00017-of-00032.safetensors +3 -0
- model-00018-of-00032.safetensors +3 -0
- model-00019-of-00032.safetensors +3 -0
- model-00020-of-00032.safetensors +3 -0
- model-00021-of-00032.safetensors +3 -0
- model-00022-of-00032.safetensors +3 -0
- model-00023-of-00032.safetensors +3 -0
- model-00024-of-00032.safetensors +3 -0
- model-00025-of-00032.safetensors +3 -0
- model-00026-of-00032.safetensors +3 -0
- model-00027-of-00032.safetensors +3 -0
- model-00028-of-00032.safetensors +3 -0
- model-00029-of-00032.safetensors +3 -0
- model-00030-of-00032.safetensors +3 -0
- model-00031-of-00032.safetensors +3 -0
- model-00032-of-00032.safetensors +3 -0
- model.safetensors.index.json +0 -0
- modeling_modilify_mk2.py +157 -254
- mps_ops.py +99 -375
- 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.
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
|
|
|
|
| 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:
|
| 7 |
tags:
|
| 8 |
- diffusion
|
| 9 |
-
- multimodal
|
| 10 |
-
- image-text-to-text
|
| 11 |
- mixture-of-experts
|
| 12 |
- trust-remote-code
|
|
|
|
| 13 |
---
|
| 14 |
|
| 15 |

|
| 16 |
|
| 17 |
-
# Modilify Mk2 Preview
|
| 18 |
|
| 19 |
-
|
|
|
|
|
|
|
|
|
|
| 20 |
|
| 21 |
-
|
|
|
|
| 22 |
|
| 23 |
-
|
| 24 |
|
| 25 |
-
|
| 26 |
-
|
| 27 |
-
|
| 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
|
| 115 |
|
| 116 |
-
|
| 117 |
-
|
| 118 |
-
model =
|
| 119 |
-
|
| 120 |
trust_remote_code=True,
|
| 121 |
dtype=torch.bfloat16,
|
| 122 |
device_map="auto",
|
| 123 |
-
|
| 124 |
-
|
| 125 |
-
|
| 126 |
-
|
| 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
|
| 136 |
-
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |

|
| 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:
|
| 3 |
-
Author:
|
| 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 |
-
|
| 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
|
| 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 |
-
|
| 56 |
-
|
|
|
|
| 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 |
-
"""
|
| 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 |
-
"
|
|
|
|
| 15 |
"commit_failure_budget": 0.2,
|
|
|
|
| 16 |
"commit_sequence_dim": 1024,
|
| 17 |
-
"
|
| 18 |
-
"
|
| 19 |
-
"
|
| 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":
|
| 42 |
"latent_working_last_block_global": true,
|
| 43 |
-
"
|
| 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 |
-
"
|
| 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":
|
| 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 |
-
"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 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 |
-
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
|
| 72 |
-
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
|
| 76 |
-
|
| 77 |
-
|
| 78 |
-
|
| 79 |
-
|
| 80 |
-
|
| 81 |
-
|
| 82 |
-
|
| 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 |
-
|
|
|
|
|
|
|
| 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 |
-
|
| 129 |
-
|
| 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 |
-
|
| 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]
|
| 144 |
-
|
| 145 |
-
|
| 146 |
-
|
| 147 |
-
|
| 148 |
-
|
| 149 |
-
|
| 150 |
-
|
| 151 |
**kwargs: Any,
|
| 152 |
) -> None:
|
| 153 |
-
|
| 154 |
-
|
| 155 |
-
|
| 156 |
-
|
| 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 |
-
|
| 164 |
-
|
| 165 |
-
|
| 166 |
-
|
| 167 |
-
|
| 168 |
-
|
| 169 |
-
|
| 170 |
-
|
| 171 |
-
self.
|
| 172 |
-
self.
|
| 173 |
-
self.
|
| 174 |
-
self.
|
| 175 |
-
self.
|
| 176 |
-
self.
|
| 177 |
-
self.
|
| 178 |
-
self.
|
| 179 |
-
self.
|
| 180 |
-
self.
|
| 181 |
-
self.
|
| 182 |
-
self.
|
| 183 |
-
self.
|
| 184 |
-
self.
|
| 185 |
-
self.
|
| 186 |
-
|
| 187 |
-
|
| 188 |
-
|
| 189 |
-
|
| 190 |
-
|
| 191 |
-
|
| 192 |
-
|
| 193 |
-
|
| 194 |
-
|
| 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 |
-
|
| 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 |
-
"""
|
| 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 |
-
|
| 184 |
-
|
| 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 |
-
|
| 196 |
-
|
| 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 |
-
|
| 1056 |
-
|
| 1057 |
-
|
| 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 |
-
|
| 1083 |
-
|
| 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.
|
| 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 |
-
|
| 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 |
-
|
| 1304 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1305 |
)
|
| 1306 |
-
|
| 1307 |
-
output.heavy_hidden_state,
|
|
|
|
| 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 |
-
|
| 1316 |
-
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1328 |
)
|
| 1329 |
jump_failure_rate = fused_commit_failure_rate(
|
| 1330 |
greedy_confidence,
|
| 1331 |
token_entropy,
|
| 1332 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1333 |
)
|
| 1334 |
previous_failure_rate = fused_commit_failure_rate(
|
| 1335 |
rolling.confidence,
|
| 1336 |
rolling.entropy,
|
| 1337 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1338 |
)
|
| 1339 |
(
|
| 1340 |
normal_commit,
|
|
@@ -1372,17 +1326,13 @@ class ModilifyMk2ContinuousBatchingManager:
|
|
| 1372 |
stagnation_steps=next_stagnation,
|
| 1373 |
),
|
| 1374 |
)
|
| 1375 |
-
|
| 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 |
-
"
|
| 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 |
-
|
| 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
|
| 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 |
-
|
| 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 =
|
| 170 |
-
self.t_max =
|
| 171 |
-
self.validate()
|
| 172 |
|
| 173 |
def update(self, defaults_only=False, allow_custom_entries=False, **kwargs):
|
| 174 |
-
"""Apply supported overrides
|
| 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 =
|
| 215 |
-
self.t_max =
|
| 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
|
| 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 |
-
|
| 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 |
-
|
| 754 |
) -> ModilifyMk2RollingState:
|
| 755 |
if not bool(commit_lengths.gt(0).any()):
|
| 756 |
return next_state
|
| 757 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 758 |
memory=next_state.latent_state.memory_slots,
|
| 759 |
working_state=working_state,
|
| 760 |
-
|
| 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(
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 792 |
-
|
| 793 |
-
|
| 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 |
-
|
| 824 |
-
|
| 825 |
-
|
| 826 |
-
|
| 827 |
-
|
| 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 |
-
|
| 860 |
-
|
| 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 |
-
|
| 1048 |
-
|
| 1049 |
-
|
| 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 |
-
|
| 1067 |
-
|
| 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 |
-
|
| 1157 |
-
|
| 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 |
-
|
| 1199 |
-
output.heavy_hidden_state,
|
|
|
|
| 1200 |
)
|
| 1201 |
next_state = ModilifyMk2RollingState(
|
| 1202 |
canvas=next_canvas, confidence=next_confidence,
|
| 1203 |
-
entropy=token_entropy,
|
| 1204 |
latent_state=next_latent,
|
| 1205 |
-
|
| 1206 |
-
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1218 |
)
|
| 1219 |
jump_failure_rate = fused_commit_failure_rate(
|
| 1220 |
greedy_confidence, token_entropy,
|
| 1221 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1222 |
)
|
| 1223 |
previous_failure_rate = fused_commit_failure_rate(
|
| 1224 |
state.confidence, state.entropy,
|
| 1225 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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=
|
| 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.
|
| 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 |
-
|
| 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 |
-
|
| 41 |
-
|
| 42 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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=
|
| 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 |
-
|
| 326 |
-
|
| 327 |
-
|
| 328 |
-
|
| 329 |
-
|
| 330 |
-
|
| 331 |
-
|
| 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 |
-
|
| 415 |
-
|
| 416 |
-
|
| 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 |
-
|
| 436 |
-
|
| 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
|
| 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 |
-
"""
|
| 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(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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(
|
| 955 |
-
|
|
|
|
|
|
|
| 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 |
-
|
| 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 |
-
|
| 2101 |
-
|
| 2102 |
-
|
| 2103 |
-
|
| 2104 |
-
|
| 2105 |
-
|
| 2106 |
-
|
| 2107 |
-
|
| 2108 |
-
|
| 2109 |
-
|
| 2110 |
-
|
| 2111 |
-
|
| 2112 |
-
|
| 2113 |
-
|
| 2114 |
-
|
| 2115 |
-
|
| 2116 |
-
|
| 2117 |
-
|
| 2118 |
-
|
| 2119 |
-
|
| 2120 |
-
|
| 2121 |
-
|
| 2122 |
-
|
| 2123 |
-
|
| 2124 |
-
|
| 2125 |
-
|
| 2126 |
-
|
| 2127 |
-
|
| 2128 |
-
|
| 2129 |
-
|
| 2130 |
-
|
| 2131 |
-
|
| 2132 |
-
|
| 2133 |
-
|
| 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 |
-
|
| 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
|
| 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 |
-
|
| 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
|
| 33 |
-
from .
|
|
|
|
| 34 |
from .generation_modilify_mk2 import ModilifyMk2GenerationConfig, ModilifyMk2GenerationMixin
|
| 35 |
-
from .latent_deliberation import
|
| 36 |
-
|
| 37 |
-
LatentDeliberationTransformer,
|
| 38 |
-
LatentProcessorOutput,
|
| 39 |
-
TrajectoryHistory,
|
| 40 |
-
TrajectoryTape,
|
| 41 |
-
empty_trajectory_tape,
|
| 42 |
-
)
|
| 43 |
from .vocab_ops import chunked_vocab_statistics
|
| 44 |
|
| 45 |
|
| 46 |
-
|
| 47 |
-
|
| 48 |
-
|
| 49 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 61 |
-
|
| 62 |
-
|
| 63 |
-
|
| 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 |
-
|
| 80 |
-
|
| 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 =
|
| 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
|
| 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(
|
| 137 |
-
"""
|
| 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
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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] =
|
| 197 |
-
config._attn_implementation
|
| 198 |
-
|
| 199 |
-
|
| 200 |
-
|
| 201 |
-
|
| 202 |
-
|
| 203 |
-
|
| 204 |
-
|
| 205 |
-
|
| 206 |
-
|
| 207 |
-
|
| 208 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 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
|
| 308 |
-
memory_kv =
|
| 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
|
|
|
|
|
|
|
|
|
|
| 329 |
)
|
| 330 |
working_reader += 1
|
| 331 |
if (
|
| 332 |
memory_kv is not None
|
| 333 |
-
and reader_index <
|
| 334 |
):
|
| 335 |
-
hidden =
|
| 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 |
-
) ->
|
| 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 |
-
|
| 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 |
-
|
| 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
|
| 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 |
-
**
|
| 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 |
-
|
| 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 |
-
|
|
|
|
| 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 |
-
|
| 596 |
-
|
| 597 |
)
|
|
|
|
|
|
|
| 598 |
return (
|
| 599 |
-
|
| 600 |
-
|
| 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 |
-
|
| 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 |
-
|
| 632 |
-
|
| 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 |
-
|
| 661 |
-
|
| 662 |
-
|
| 663 |
-
|
| 664 |
-
|
| 665 |
-
|
| 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 |
-
|
| 686 |
-
|
| 687 |
-
|
| 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 |
-
|
| 708 |
-
|
| 709 |
-
|
| 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(
|
| 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(
|
| 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
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
|
| 62 |
-
|
| 63 |
-
|
| 64 |
-
|
| 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 |
-
|
| 367 |
-
|
| 368 |
-
|
| 369 |
-
|
| 370 |
-
|
| 371 |
-
|
| 372 |
-
|
| 373 |
-
|
| 374 |
-
|
| 375 |
-
|
| 376 |
-
|
| 377 |
-
)
|
| 378 |
-
|
| 379 |
-
|
| 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
|
| 500 |
inputs,
|
| 501 |
weight_a,
|
| 502 |
weight_b,
|
| 503 |
expert_ids,
|
| 504 |
-
counts,
|
| 505 |
)
|
| 506 |
-
return
|
| 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")
|
| 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
|
| 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")
|
| 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
|
| 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 |
-
|
| 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 |
-
|
| 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 |
-
|
| 363 |
-
|
| 364 |
-
if
|
| 365 |
-
|
| 366 |
-
|
| 367 |
-
|
| 368 |
-
|
| 369 |
-
|
| 370 |
-
|
| 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 |
-
"`
|
| 382 |
)
|
| 383 |
-
|
| 384 |
-
|
| 385 |
-
|
| 386 |
-
|
| 387 |
-
|
| 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 |
-
|
| 408 |
-
|
| 409 |
-
|
| 410 |
-
|
| 411 |
-
|
| 412 |
-
|
| 413 |
-
|
| 414 |
-
|
| 415 |
-
|
| 416 |
-
|
| 417 |
-
|
| 418 |
-
|
| 419 |
-
|
| 420 |
-
|
| 421 |
-
|
| 422 |
-
|
| 423 |
-
|
| 424 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 469 |
-
|
| 470 |
-
|
| 471 |
-
|
| 472 |
-
|
| 473 |
-
|
| 474 |
-
|
| 475 |
-
|
| 476 |
-
|
| 477 |
-
|
| 478 |
-
|
| 479 |
-
|
| 480 |
-
|
| 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 |
-
|
| 490 |
-
|
| 491 |
-
|
| 492 |
-
|
| 493 |
-
|
| 494 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 495 |
)
|
| 496 |
-
|
| 497 |
-
|
| 498 |
-
|
| 499 |
-
-
|
| 500 |
-
|
| 501 |
-
|
|
|
|
|
|
|
|
|
|
| 502 |
)
|
| 503 |
-
(
|
| 504 |
-
|
| 505 |
-
|
| 506 |
-
|
| 507 |
-
|
| 508 |
-
|
| 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 |
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|