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
Browse files- README.md +65 -23
- assets/01-LOGO.jpg +2 -2
- model-00008-of-00011.safetensors +3 -0
- model-00009-of-00011.safetensors +3 -0
- model-00010-of-00011.safetensors +3 -0
- model-00011-of-00011.safetensors +3 -0
- model.safetensors.index.json +0 -0
- modeling_modilify_mk2.py +735 -0
- mps_ops.py +618 -0
- processor_config.json +75 -0
- tokenizer.json +3 -0
- tokenizer_config.json +96 -0
- vocab_ops.py +525 -0
README.md
CHANGED
|
@@ -14,49 +14,91 @@ tags:
|
|
| 14 |
|
| 15 |

|
| 16 |
|
| 17 |
-
# Modilify Mk2 Preview
|
| 18 |
|
| 19 |
-
|
| 20 |
|
| 21 |
-
Mk2
|
| 22 |
|
| 23 |
-
|
| 24 |
|
| 25 |
-
|
| 26 |
|
| 27 |
-
|
| 28 |
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 32 |
|
| 33 |
-
|
| 34 |
|
| 35 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 36 |
|
| 37 |
## Model Summary
|
| 38 |
|
| 39 |
| | |
|
| 40 |
| --- | ---: |
|
| 41 |
| Architecture | Mixture-of-Experts block diffusion + dual-timescale latent Transformer |
|
| 42 |
-
|
|
| 43 |
-
|
|
| 44 |
-
| Vision Encoder | Gemma 4 Vision, 569.550M |
|
| 45 |
| Layers | 30 |
|
| 46 |
| Number of Experts | 128 |
|
| 47 |
| Selected Experts per Token | 8 |
|
| 48 |
| Vocabulary Size | 262,144 |
|
| 49 |
-
| Context Length | 262,144 tokens |
|
| 50 |
| Sliding Window | 1024 |
|
| 51 |
| Canvas Length | 256 |
|
| 52 |
| Latent Width | 2,816 |
|
| 53 |
-
| Latent
|
|
|
|
|
|
|
| 54 |
| Trajectory History | 16 frames, 4 views, rank 1,024 |
|
| 55 |
-
| Denoise Tape | 16 probes |
|
| 56 |
| Modality | Text, Image, Video |
|
| 57 |
-
| Preview checkpoint |
|
| 58 |
| Adaptation tokens | ~12.4 million |
|
| 59 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 60 |
## Getting Started
|
| 61 |
|
| 62 |
Transformers 5.14.1 is the minimum supported version.
|
|
@@ -157,12 +199,12 @@ output = model.generate(**inputs, max_new_tokens=256)
|
|
| 157 |
|
| 158 |
## Thinking mode
|
| 159 |
|
| 160 |
-
|
| 161 |
|
| 162 |
- `enable_thinking=True` inserts a system turn that contains `<|think|>` and still ends the prompt at `<|turn>model`.
|
| 163 |
- `enable_thinking=False` does **not** inject an empty thought channel. The prompt ends at `<|turn>model`.
|
| 164 |
|
| 165 |
-
The model may still open `<|channel>thought` on its own
|
| 166 |
|
| 167 |
## Configurable inference
|
| 168 |
|
|
@@ -208,13 +250,13 @@ model = AutoModelForMultimodalLM.from_pretrained(
|
|
| 208 |
)
|
| 209 |
```
|
| 210 |
|
| 211 |
-
|
| 212 |
|
| 213 |
Generation supports left-padded batches with independent stopping. Batch prompts of similar lengths together for the best throughput. Streaming and caller-supplied KV caches remain limited to batch size 1.
|
| 214 |
|
| 215 |
## Evaluation status, limitations, and risks
|
| 216 |
|
| 217 |
-
This
|
| 218 |
|
| 219 |
The model can hallucinate facts, citations, visual details, or temporal relationships; reproduce bias, unsafe content, personal information, or copyrighted material; and consume substantial time and memory during long iterative generation. Confidence-based commits control compute. They do not certify that a prefix is true. Visual performance can degrade with poor resolution, motion, occlusion, unusual aspect ratios, or domain shift.
|
| 220 |
|
|
@@ -228,7 +270,7 @@ Released under the [Modilify Open Model License 1.0](LICENSE), subject to its re
|
|
| 228 |
|
| 229 |
```bibtex
|
| 230 |
@software{modilify_mk2_preview_2026,
|
| 231 |
-
title = {Modilify Mk2 Preview},
|
| 232 |
author = {Modilify},
|
| 233 |
year = {2026},
|
| 234 |
note = {A multimodal dual-timescale latent-deliberation derivative of DiffusionGemma}
|
|
|
|
| 14 |
|
| 15 |

|
| 16 |
|
| 17 |
+
# Modilify Mk2 Preview · 26B-A5B
|
| 18 |
|
| 19 |
+
**Refine a whole canvas. Think in latent space. Carry memory forward.**
|
| 20 |
|
| 21 |
+
Modilify Mk2 is a **26B-A5B multimodal block-diffusion model** that brings parallel token refinement, latent deliberation, and persistent trajectory memory into one generation loop. A rolling 256-token canvas gives the model room to revise upcoming text together. A dedicated latent Transformer turns the history of those revisions into context for the next step. As output advances, a learned memory writer carries information across canvas boundaries.
|
| 22 |
|
| 23 |
+
The central idea is simple: **give each answer a workspace, a memory, and a variable compute budget.** Mk2 can refine its internal state over multiple passes before committing text, and release multiple tokens together when its confidence-and-entropy policy permits. Deliberation happens in continuous hidden states, without requiring every internal update to become a visible reasoning token.
|
| 24 |
|
| 25 |
+
Built on DiffusionGemma, Mk2 combines sparse expert routing with a dual-timescale latent architecture. Text, images, and sampled video frames feed the same generation path.
|
| 26 |
|
| 27 |
+
## What makes Mk2 different
|
| 28 |
|
| 29 |
+
| Architecture choice | What it enables |
|
| 30 |
+
| --- | --- |
|
| 31 |
+
| **Parallel block diffusion** | Revise a 256-token canvas jointly and commit a variable-length prefix, allowing multiple output tokens per denoising pass. |
|
| 32 |
+
| **Trajectory-aware latent deliberation** | Condition the next revision on how hidden states have evolved, including their changes, acceleration, and residuals relative to the latest state. |
|
| 33 |
+
| **Memory at two timescales** | Rebuild working state every pass while retaining persistent slots across rolling-window shifts within a generation. |
|
| 34 |
+
| **Memory inside the decoder** | Feed working and persistent memory directly into full-attention decoder layers so both can influence token refinement. |
|
| 35 |
+
| **Adaptive commitment** | Use confidence and entropy to decide how much text to release, with explicit budgets for continued refinement and forced progress. |
|
| 36 |
+
| **Sparse multimodal foundation** | Select 8 of 128 experts per token and bring text, image, and video context into a shared decoder. |
|
| 37 |
|
| 38 |
+
## Inside the generation loop
|
| 39 |
|
| 40 |
+
### A canvas built for revision
|
| 41 |
+
|
| 42 |
+
Mk2 maintains a rolling canvas of 256 candidate tokens. Each denoising pass updates the candidates using the encoded prompt, the current canvas, and latent context. Attention lets positions within the canvas inform one another before the output prefix is finalized. After a commit, the canvas shifts forward and opens space for new candidates.
|
| 43 |
+
|
| 44 |
+
This creates two useful degrees of freedom: **how many times to refine** and **how many tokens to release**. A pass can commit a longer prefix when the policy permits, or spend additional computation refining an uncertain frontier.
|
| 45 |
+
|
| 46 |
+
### Deliberation that reads its own trajectory
|
| 47 |
+
|
| 48 |
+
A 4-layer latent Transformer, 2,816 dimensions wide, builds working context before each decoder denoising pass. It reads the current canvas alongside three complementary sources of state:
|
| 49 |
+
|
| 50 |
+
- **Per-token trajectory history:** 16 recent frames represented through four views—hidden state, first difference, second difference, and residual from the latest state—projected to rank 1,024. These give the processor access to the direction and stability of recent revisions. History follows surviving canvas tokens; new positions start empty.
|
| 51 |
+
- **Denoising tape:** 16 pooled probes per frame summarize canvas activity in a time-indexed ring. The tape preserves recent step-level context as token positions move through the window.
|
| 52 |
+
- **Persistent memory:** 256 slots of 2,816 dimensions carry learned summaries across commits within the generation.
|
| 53 |
+
|
| 54 |
+
The resulting working context enters the decoder through its self-conditioning bridge and working-memory bus. Full-attention layers also read a separate persistent-memory bus. **The refinement history becomes an input to the next refinement.**
|
| 55 |
+
|
| 56 |
+
### Fast working state, lasting commit memory
|
| 57 |
+
|
| 58 |
+
Working state is recomputed on every denoising pass. Persistent slots update only when tokens commit, using a Transformer writer with a separate gate for each slot. That writer draws on the committed region's trajectory, working state, and final decoder representation.
|
| 59 |
+
|
| 60 |
+
This separates rapid revision from memory consolidation: the canvas can keep changing while persistent memory stays stable between commits. When the window advances, the memory slots remain available to later tokens.
|
| 61 |
+
|
| 62 |
+
### Compute that follows the commit frontier
|
| 63 |
+
|
| 64 |
+
Mk2 combines proposal confidence with an excess-entropy penalty and selects the longest prefix whose cumulative failure score stays below the configured budget. A tighter budget requires stronger evidence before normal commitment; a looser budget admits longer prefixes for the same scores. Stagnation handling and a pondering watchdog bound continued refinement.
|
| 65 |
+
|
| 66 |
+
Temperature and commit budget are exposed directly at inference time, making the generation policy adjustable per request. These scores govern commitment; they are not calibrated guarantees of factual correctness.
|
| 67 |
+
|
| 68 |
+
### One generation path for text, images, and video
|
| 69 |
+
|
| 70 |
+
The Gemma 4 vision tower supplies visual features through the DiffusionGemma multimodal encoder. Text, image, and sampled video-frame inputs are encoded into the prefix KV cache that conditions the decoder. The rolling text canvas then uses the same latent deliberation and memory loop across all three input modalities.
|
| 71 |
|
| 72 |
## Model Summary
|
| 73 |
|
| 74 |
| | |
|
| 75 |
| --- | ---: |
|
| 76 |
| Architecture | Mixture-of-Experts block diffusion + dual-timescale latent Transformer |
|
| 77 |
+
| Model Size | 26B-A5B |
|
| 78 |
+
| Vision Encoder | Gemma 4 Vision |
|
|
|
|
| 79 |
| Layers | 30 |
|
| 80 |
| Number of Experts | 128 |
|
| 81 |
| Selected Experts per Token | 8 |
|
| 82 |
| Vocabulary Size | 262,144 |
|
| 83 |
+
| Configured Context Length | 262,144 tokens |
|
| 84 |
| Sliding Window | 1024 |
|
| 85 |
| Canvas Length | 256 |
|
| 86 |
| Latent Width | 2,816 |
|
| 87 |
+
| Latent Transformer | 4 layers, 16 attention heads |
|
| 88 |
+
| Persistent Memory | 256 slots × 2,816 dimensions |
|
| 89 |
+
| Memory Writer | 2-layer commit-sequence Transformer + per-slot gated writer |
|
| 90 |
| Trajectory History | 16 frames, 4 views, rank 1,024 |
|
| 91 |
+
| Denoise Tape | 16 probes per frame |
|
| 92 |
| Modality | Text, Image, Video |
|
| 93 |
+
| Preview checkpoint | Training step 900 |
|
| 94 |
| Adaptation tokens | ~12.4 million |
|
| 95 |
|
| 96 |
+
## Preview release
|
| 97 |
+
|
| 98 |
+
This first public preview includes **merged BF16 weights, inference code, processor, and tokenizer**. The text backbone and latent stack are exported from training step 900, after approximately 12.4 million adaptation tokens. The Gemma 4 vision tower is restored from DiffusionGemma.
|
| 99 |
+
|
| 100 |
+
The release makes the architecture available for hands-on exploration and evaluation. Comprehensive benchmark results are not included; measured speed, reasoning quality, and multimodal reliability remain to be established for specific workloads.
|
| 101 |
+
|
| 102 |
## Getting Started
|
| 103 |
|
| 104 |
Transformers 5.14.1 is the minimum supported version.
|
|
|
|
| 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 |
|
|
|
|
| 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 |
|
|
|
|
| 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}
|
assets/01-LOGO.jpg
CHANGED
|
Git LFS Details
|
|
Git LFS Details
|
model-00008-of-00011.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:0b52a8c906686648eb04b3e28bc5c597e9b687de0a473ec093ad339dbfd8da77
|
| 3 |
+
size 4884578046
|
model-00009-of-00011.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:d425f9f2a800dd424e997db11933f6e4602bab648e2731a3d76fbe778554b2c4
|
| 3 |
+
size 4913414718
|
model-00010-of-00011.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:802add3f8c81ff7c44dbf96632682b90a1a0a08f35e80be867d4f0294654c03b
|
| 3 |
+
size 4884577974
|
model-00011-of-00011.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:1e3a84ff63b8e779f287a5b227aa8252383554d7b3527fc895aeb623b0c998f1
|
| 3 |
+
size 3959052840
|
model.safetensors.index.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
modeling_modilify_mk2.py
ADDED
|
@@ -0,0 +1,735 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2026 Modilify
|
| 2 |
+
# SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.0
|
| 3 |
+
"""Multimodal Modilify Mk2 model with recurrent latent deliberation."""
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
from collections.abc import Sequence
|
| 8 |
+
from dataclasses import dataclass, replace
|
| 9 |
+
from typing import Any
|
| 10 |
+
|
| 11 |
+
import torch
|
| 12 |
+
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,
|
| 20 |
+
bidirectional_mask_function,
|
| 21 |
+
)
|
| 22 |
+
from transformers.models.diffusion_gemma import (
|
| 23 |
+
DiffusionGemmaDecoderModel,
|
| 24 |
+
DiffusionGemmaEncoderModel,
|
| 25 |
+
DiffusionGemmaPreTrainedModel,
|
| 26 |
+
)
|
| 27 |
+
from transformers.models.diffusion_gemma.modeling_diffusion_gemma import (
|
| 28 |
+
DiffusionGemmaRMSNorm,
|
| 29 |
+
DiffusionGemmaTextRouter,
|
| 30 |
+
)
|
| 31 |
+
|
| 32 |
+
from .mps_ops import mps_segmented_experts_forward as _mps_segmented_experts_forward # noqa: F401
|
| 33 |
+
from .configuration_modilify_mk2 import ModilifyMk2Config
|
| 34 |
+
from .generation_modilify_mk2 import ModilifyMk2GenerationConfig, ModilifyMk2GenerationMixin
|
| 35 |
+
from .latent_deliberation import (
|
| 36 |
+
LatentDeliberationState,
|
| 37 |
+
LatentDeliberationTransformer,
|
| 38 |
+
LatentProcessorOutput,
|
| 39 |
+
TrajectoryHistory,
|
| 40 |
+
TrajectoryTape,
|
| 41 |
+
empty_trajectory_tape,
|
| 42 |
+
)
|
| 43 |
+
from .vocab_ops import chunked_vocab_statistics
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
@dataclass
|
| 47 |
+
class ModilifyMk2DecoderOutput(BaseModelOutputWithPast):
|
| 48 |
+
token_embeddings: torch.FloatTensor | None = None
|
| 49 |
+
latent_residual_diagnostics: dict[str, torch.Tensor] | None = None
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
@dataclass
|
| 53 |
+
class ModilifyMk2ModelOutput(BaseModelOutputWithPast):
|
| 54 |
+
token_embeddings: torch.FloatTensor | None = None
|
| 55 |
+
encoder_last_hidden_state: torch.FloatTensor | None = None
|
| 56 |
+
latent_residual_diagnostics: dict[str, torch.Tensor] | None = None
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
@dataclass
|
| 60 |
+
class ModilifyMk2BlockDiffusionOutput(ModelOutput):
|
| 61 |
+
"""Inference output used by the rolling diffusion generator."""
|
| 62 |
+
|
| 63 |
+
logits: torch.FloatTensor | None = None
|
| 64 |
+
heavy_hidden_state: torch.FloatTensor | None = None
|
| 65 |
+
next_latent_state: LatentDeliberationState | None = None
|
| 66 |
+
past_key_values: Cache | None = None
|
| 67 |
+
encoder_last_hidden_state: torch.FloatTensor | None = None
|
| 68 |
+
temporal_context: torch.FloatTensor | None = None
|
| 69 |
+
history_projected: torch.FloatTensor | None = None
|
| 70 |
+
working_state: torch.FloatTensor | None = None
|
| 71 |
+
latent_residual_diagnostics: dict[str, torch.Tensor] | None = None
|
| 72 |
+
proposal: torch.LongTensor | None = None
|
| 73 |
+
proposal_confidence: torch.FloatTensor | None = None
|
| 74 |
+
token_entropy: torch.FloatTensor | None = None
|
| 75 |
+
greedy_proposal: torch.LongTensor | None = None
|
| 76 |
+
greedy_confidence: torch.FloatTensor | None = None
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
class ModilifyMk2RMSNorm(DiffusionGemmaRMSNorm):
|
| 80 |
+
"""Official RMSNorm parameters, pre-fusion same-dtype forward."""
|
| 81 |
+
|
| 82 |
+
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
| 83 |
+
normed_output = self._norm(hidden_states)
|
| 84 |
+
if self.with_scale:
|
| 85 |
+
normed_output = normed_output * self.weight.to(dtype=normed_output.dtype)
|
| 86 |
+
return normed_output.type_as(hidden_states)
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
class ModilifyMk2TextRouter(DiffusionGemmaTextRouter):
|
| 90 |
+
"""Official router parameters with fp32 softmax so top-k weights stay finite."""
|
| 91 |
+
|
| 92 |
+
def __init__(self, config: Any) -> None:
|
| 93 |
+
super().__init__(config)
|
| 94 |
+
self.norm = ModilifyMk2RMSNorm(self.hidden_size, eps=self.eps, with_scale=False)
|
| 95 |
+
|
| 96 |
+
def forward(
|
| 97 |
+
self, hidden_states: torch.Tensor
|
| 98 |
+
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
| 99 |
+
hidden_states = self.norm(hidden_states)
|
| 100 |
+
hidden_states = hidden_states * self.scale * self.scalar_root_size
|
| 101 |
+
expert_scores = self.proj(hidden_states)
|
| 102 |
+
# Match the official DiffusionGemma router: bf16 softmax underflows to
|
| 103 |
+
# all-zero top-k rows, then 0/0 NaNs the MoE residual and every loss.
|
| 104 |
+
router_probabilities = F.softmax(expert_scores, dim=-1, dtype=torch.float32)
|
| 105 |
+
top_k_weights, top_k_index = torch.topk(
|
| 106 |
+
router_probabilities,
|
| 107 |
+
k=self.config.top_k_experts,
|
| 108 |
+
dim=-1,
|
| 109 |
+
)
|
| 110 |
+
top_k_weights = top_k_weights / top_k_weights.sum(dim=-1, keepdim=True).clamp_min(
|
| 111 |
+
torch.finfo(torch.float32).tiny
|
| 112 |
+
)
|
| 113 |
+
top_k_weights = top_k_weights * self.per_expert_scale[top_k_index]
|
| 114 |
+
return router_probabilities, top_k_weights, top_k_index
|
| 115 |
+
|
| 116 |
+
|
| 117 |
+
def install_modilify_mk2_trunk_semantics(module: nn.Module) -> None:
|
| 118 |
+
"""Swap official leaf modules on this instance. Never patch Transformers classes."""
|
| 119 |
+
|
| 120 |
+
for name, child in list(module.named_children()):
|
| 121 |
+
if type(child) is DiffusionGemmaRMSNorm:
|
| 122 |
+
dim = int(child.weight.shape[0]) if child.with_scale else 1
|
| 123 |
+
replacement = ModilifyMk2RMSNorm(
|
| 124 |
+
dim, eps=child.eps, with_scale=child.with_scale
|
| 125 |
+
)
|
| 126 |
+
replacement.load_state_dict(child.state_dict())
|
| 127 |
+
setattr(module, name, replacement)
|
| 128 |
+
elif type(child) is DiffusionGemmaTextRouter:
|
| 129 |
+
replacement = ModilifyMk2TextRouter(child.config)
|
| 130 |
+
replacement.load_state_dict(child.state_dict())
|
| 131 |
+
setattr(module, name, replacement)
|
| 132 |
+
else:
|
| 133 |
+
install_modilify_mk2_trunk_semantics(child)
|
| 134 |
+
|
| 135 |
+
|
| 136 |
+
class ModilifyMk2EncoderModel(DiffusionGemmaEncoderModel):
|
| 137 |
+
"""Unmodified Transformers DiffusionGemma multimodal encoder."""
|
| 138 |
+
|
| 139 |
+
config_class = ModilifyMk2Config
|
| 140 |
+
|
| 141 |
+
|
| 142 |
+
class ModilifyMk2DecoderModel(DiffusionGemmaDecoderModel):
|
| 143 |
+
"""Diffusion decoder accepting compact self-conditioning embeddings."""
|
| 144 |
+
|
| 145 |
+
config_class = ModilifyMk2Config
|
| 146 |
+
latent_residual_rms_ratio_cap = 0.5
|
| 147 |
+
|
| 148 |
+
def __init__(self, config: ModilifyMk2Config):
|
| 149 |
+
super().__init__(config)
|
| 150 |
+
install_modilify_mk2_trunk_semantics(self)
|
| 151 |
+
|
| 152 |
+
@staticmethod
|
| 153 |
+
def create_diffusion_decoder_attention_mask(
|
| 154 |
+
config: Any,
|
| 155 |
+
inputs_embeds: torch.Tensor,
|
| 156 |
+
past_key_values: Cache,
|
| 157 |
+
decoder_attention_mask: torch.Tensor | dict | None = None,
|
| 158 |
+
) -> dict[str, torch.Tensor | None]:
|
| 159 |
+
"""Official mask builder without the all-True sliding-window skip."""
|
| 160 |
+
|
| 161 |
+
if past_key_values is None:
|
| 162 |
+
raise ValueError(
|
| 163 |
+
"The diffusion mask requires `past_key_values` to construct the next attention mask correctly"
|
| 164 |
+
)
|
| 165 |
+
if (
|
| 166 |
+
decoder_attention_mask is None
|
| 167 |
+
or config._attn_implementation
|
| 168 |
+
not in ALL_MASK_ATTENTION_FUNCTIONS._global_mapping
|
| 169 |
+
):
|
| 170 |
+
return {"full_attention": None, "sliding_attention": None}
|
| 171 |
+
if isinstance(decoder_attention_mask, dict) and all(
|
| 172 |
+
mask.ndim == 4 for mask in decoder_attention_mask.values()
|
| 173 |
+
):
|
| 174 |
+
return decoder_attention_mask
|
| 175 |
+
|
| 176 |
+
text_config = config.get_text_config() if hasattr(config, "get_text_config") else config
|
| 177 |
+
q_length = inputs_embeds.shape[1]
|
| 178 |
+
q_offset = past_key_values.get_seq_length()
|
| 179 |
+
if isinstance(q_offset, torch.Tensor):
|
| 180 |
+
q_offset = q_offset.to(inputs_embeds.device)
|
| 181 |
+
additional_kv_length = (
|
| 182 |
+
getattr(config, "canvas_length", 0) if past_key_values.is_compileable else 0
|
| 183 |
+
)
|
| 184 |
+
mask_mapping: dict[str, torch.Tensor | None] = {}
|
| 185 |
+
for layer_pattern in set(text_config.layer_types):
|
| 186 |
+
layer_idx = past_key_values.is_sliding.index(
|
| 187 |
+
layer_pattern == "sliding_attention"
|
| 188 |
+
)
|
| 189 |
+
kv_length, kv_offset = past_key_values.get_mask_sizes(q_length, layer_idx)
|
| 190 |
+
kv_length += additional_kv_length
|
| 191 |
+
if layer_pattern == "sliding_attention" and past_key_values.is_compileable:
|
| 192 |
+
sliding_layer = past_key_values.layers[layer_idx]
|
| 193 |
+
max_length = sliding_layer.get_max_length() + additional_kv_length
|
| 194 |
+
if kv_length >= max_length:
|
| 195 |
+
kv_length = max_length
|
| 196 |
+
mask_mapping[layer_pattern] = ALL_MASK_ATTENTION_FUNCTIONS[
|
| 197 |
+
config._attn_implementation
|
| 198 |
+
](
|
| 199 |
+
batch_size=inputs_embeds.shape[0],
|
| 200 |
+
q_length=q_length,
|
| 201 |
+
kv_length=kv_length,
|
| 202 |
+
q_offset=q_offset,
|
| 203 |
+
kv_offset=kv_offset,
|
| 204 |
+
mask_function=bidirectional_mask_function,
|
| 205 |
+
attention_mask=decoder_attention_mask,
|
| 206 |
+
allow_is_causal_skip=False,
|
| 207 |
+
allow_is_bidirectional_skip=True,
|
| 208 |
+
local_size=getattr(text_config, "sliding_window", None),
|
| 209 |
+
dtype=inputs_embeds.dtype,
|
| 210 |
+
config=text_config,
|
| 211 |
+
use_vmap=False,
|
| 212 |
+
device=inputs_embeds.device,
|
| 213 |
+
)
|
| 214 |
+
return mask_mapping
|
| 215 |
+
|
| 216 |
+
def merge_latent_context(
|
| 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:
|
| 226 |
+
context = torch.zeros_like(token_embeddings)
|
| 227 |
+
else:
|
| 228 |
+
if latent_context.shape != token_embeddings.shape:
|
| 229 |
+
raise ValueError("Temporal context must match canvas hidden-state shape.")
|
| 230 |
+
context = latent_context.to(token_embeddings)
|
| 231 |
+
mapper = self.self_conditioning
|
| 232 |
+
normalized_context = mapper.pre_norm(context)
|
| 233 |
+
mapped_context = mapper.down_proj(
|
| 234 |
+
mapper.act_fn(mapper.gate_proj(normalized_context))
|
| 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()
|
| 245 |
+
)
|
| 246 |
+
residual_cap = self.latent_residual_rms_ratio_cap * token_rms_per_token
|
| 247 |
+
soft_cap_scale = residual_cap / torch.sqrt(
|
| 248 |
+
mapped_token_energy + residual_cap.square() + 1.0e-12
|
| 249 |
+
)
|
| 250 |
+
mapped_context = (mapped_fp32 * soft_cap_scale).to(dtype=mapped_context.dtype)
|
| 251 |
+
combined = token_embeddings + mapped_context
|
| 252 |
+
diagnostics: dict[str, torch.Tensor] = {}
|
| 253 |
+
if collect_diagnostics:
|
| 254 |
+
token_rms = token_embeddings.detach().float().square().mean().sqrt()
|
| 255 |
+
mapped_rms = mapped_context.detach().float().square().mean().sqrt()
|
| 256 |
+
diagnostics = {
|
| 257 |
+
"token_embedding_rms": token_rms,
|
| 258 |
+
"latent_context_rms": context.detach().float().square().mean().sqrt(),
|
| 259 |
+
"mapped_sc_rms": mapped_rms,
|
| 260 |
+
"actual_residual_rms": mapped_rms,
|
| 261 |
+
"latent_to_embedding_rms_ratio": (
|
| 262 |
+
mapped_rms / token_rms.clamp_min(1.0e-12)
|
| 263 |
+
),
|
| 264 |
+
"merged_input_rms": combined.detach().float().square().mean().sqrt(),
|
| 265 |
+
}
|
| 266 |
+
return mapper.post_norm(combined), diagnostics
|
| 267 |
+
|
| 268 |
+
def _run_stack(
|
| 269 |
+
self,
|
| 270 |
+
inputs_embeds: torch.Tensor,
|
| 271 |
+
*,
|
| 272 |
+
past_key_values: Cache | None,
|
| 273 |
+
decoder_attention_mask: torch.Tensor | dict | None,
|
| 274 |
+
decoder_position_ids: torch.LongTensor | None,
|
| 275 |
+
**kwargs: Any,
|
| 276 |
+
) -> torch.Tensor:
|
| 277 |
+
if decoder_position_ids is None:
|
| 278 |
+
prefix = past_key_values.get_seq_length(layer_idx=0) if past_key_values is not None else 0
|
| 279 |
+
decoder_position_ids = torch.arange(
|
| 280 |
+
prefix, prefix + inputs_embeds.shape[1], device=inputs_embeds.device
|
| 281 |
+
).unsqueeze(0)
|
| 282 |
+
if not isinstance(mask_mapping := decoder_attention_mask, dict):
|
| 283 |
+
mask_mapping = self.create_diffusion_decoder_attention_mask(
|
| 284 |
+
config=self.text_config,
|
| 285 |
+
inputs_embeds=inputs_embeds,
|
| 286 |
+
past_key_values=past_key_values,
|
| 287 |
+
decoder_attention_mask=decoder_attention_mask,
|
| 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
|
| 298 |
+
hidden = inputs_embeds
|
| 299 |
+
positions = {
|
| 300 |
+
layer_type: self.rotary_emb(hidden, decoder_position_ids, layer_type)
|
| 301 |
+
for layer_type in self.unique_layer_types
|
| 302 |
+
}
|
| 303 |
+
working_kv = None
|
| 304 |
+
if working_bus is not None and working_state is not None:
|
| 305 |
+
working_kv = working_bus.prepare_kv(working_state)
|
| 306 |
+
memory_kv = None
|
| 307 |
+
if memory_bus is not None and memory_slots is not None:
|
| 308 |
+
memory_kv = memory_bus.prepare_kv(memory_slots, slot_identity)
|
| 309 |
+
working_reader = 0
|
| 310 |
+
reader_index = 0
|
| 311 |
+
for index in range(self.text_config.num_hidden_layers):
|
| 312 |
+
layer = self.layers[index]
|
| 313 |
+
layer_type = self.text_config.layer_types[index]
|
| 314 |
+
hidden = layer(
|
| 315 |
+
hidden,
|
| 316 |
+
position_embeddings=positions[layer_type],
|
| 317 |
+
attention_mask=mask_mapping[layer_type],
|
| 318 |
+
position_ids=decoder_position_ids,
|
| 319 |
+
past_key_values=cache,
|
| 320 |
+
**kwargs,
|
| 321 |
+
)
|
| 322 |
+
if layer_type == "full_attention":
|
| 323 |
+
if (
|
| 324 |
+
working_kv is not None
|
| 325 |
+
and working_reader < working_bus.num_readers
|
| 326 |
+
):
|
| 327 |
+
hidden = working_bus.read(
|
| 328 |
+
hidden, working_reader, working_kv[0], working_kv[1]
|
| 329 |
+
)
|
| 330 |
+
working_reader += 1
|
| 331 |
+
if (
|
| 332 |
+
memory_kv is not None
|
| 333 |
+
and reader_index < memory_bus.num_readers
|
| 334 |
+
):
|
| 335 |
+
hidden = memory_bus.read(
|
| 336 |
+
hidden, reader_index, memory_kv[0], memory_kv[1]
|
| 337 |
+
)
|
| 338 |
+
reader_index += 1
|
| 339 |
+
return self.norm(hidden)
|
| 340 |
+
|
| 341 |
+
def forward(
|
| 342 |
+
self,
|
| 343 |
+
decoder_input_ids: torch.LongTensor,
|
| 344 |
+
past_key_values: Cache | None = None,
|
| 345 |
+
decoder_token_embeddings: torch.FloatTensor | None = None,
|
| 346 |
+
temporal_context_embeddings: torch.FloatTensor | None = None,
|
| 347 |
+
decoder_attention_mask: torch.Tensor | dict | None = None,
|
| 348 |
+
decoder_position_ids: torch.LongTensor | None = None,
|
| 349 |
+
collect_latent_diagnostics: bool = True,
|
| 350 |
+
memory_bus: Any | None = None,
|
| 351 |
+
memory_slots: torch.Tensor | None = None,
|
| 352 |
+
working_bus: Any | None = None,
|
| 353 |
+
working_state: torch.Tensor | None = None,
|
| 354 |
+
persistent_bus: Any | None = None,
|
| 355 |
+
slot_identity: torch.Tensor | None = None,
|
| 356 |
+
**kwargs: Any,
|
| 357 |
+
) -> ModilifyMk2DecoderOutput:
|
| 358 |
+
if "use_cache" in kwargs:
|
| 359 |
+
raise ValueError("The diffusion decoder always reads the supplied cache.")
|
| 360 |
+
if decoder_token_embeddings is None:
|
| 361 |
+
token_embeddings = self.embed_tokens(decoder_input_ids)
|
| 362 |
+
else:
|
| 363 |
+
token_embeddings = decoder_token_embeddings
|
| 364 |
+
expected = (*decoder_input_ids.shape, self.text_config.hidden_size)
|
| 365 |
+
if token_embeddings.shape != expected:
|
| 366 |
+
raise ValueError("Precomputed decoder embeddings have the wrong shape.")
|
| 367 |
+
context_embeddings = (
|
| 368 |
+
torch.zeros_like(token_embeddings)
|
| 369 |
+
if temporal_context_embeddings is None
|
| 370 |
+
else temporal_context_embeddings.to(token_embeddings)
|
| 371 |
+
)
|
| 372 |
+
inputs_embeds, diagnostics = self.merge_latent_context(
|
| 373 |
+
token_embeddings,
|
| 374 |
+
context_embeddings if temporal_context_embeddings is not None else None,
|
| 375 |
+
collect_diagnostics=collect_latent_diagnostics,
|
| 376 |
+
)
|
| 377 |
+
hidden = self._run_stack(
|
| 378 |
+
inputs_embeds,
|
| 379 |
+
past_key_values=past_key_values,
|
| 380 |
+
decoder_attention_mask=decoder_attention_mask,
|
| 381 |
+
decoder_position_ids=decoder_position_ids,
|
| 382 |
+
memory_bus=memory_bus,
|
| 383 |
+
memory_slots=memory_slots,
|
| 384 |
+
working_bus=working_bus,
|
| 385 |
+
working_state=working_state,
|
| 386 |
+
persistent_bus=persistent_bus,
|
| 387 |
+
slot_identity=slot_identity,
|
| 388 |
+
**kwargs,
|
| 389 |
+
)
|
| 390 |
+
return ModilifyMk2DecoderOutput(
|
| 391 |
+
last_hidden_state=hidden,
|
| 392 |
+
past_key_values=past_key_values,
|
| 393 |
+
token_embeddings=token_embeddings,
|
| 394 |
+
latent_residual_diagnostics=diagnostics,
|
| 395 |
+
)
|
| 396 |
+
|
| 397 |
+
|
| 398 |
+
class ModilifyMk2Model(DiffusionGemmaPreTrainedModel):
|
| 399 |
+
config_class = ModilifyMk2Config
|
| 400 |
+
_tied_weights_keys = {
|
| 401 |
+
"encoder.language_model.norm.weight": "decoder.norm.weight",
|
| 402 |
+
r"encoder.language_model.layers\.(?:[^.]+\.)*weight": r"decoder.layers\.(?:[^.]+\.)*weight",
|
| 403 |
+
r"encoder.language_model.layers\.(?:[^.]+\.)*scale": r"decoder.layers\.(?:[^.]+\.)*scale",
|
| 404 |
+
r"encoder.language_model.layers\.(?:[^.]+\.)*per_expert_scale": r"decoder.layers\.(?:[^.]+\.)*per_expert_scale",
|
| 405 |
+
r"encoder.language_model.layers\.(?:[^.]+\.)*gate_up_proj": r"decoder.layers\.(?:[^.]+\.)*gate_up_proj",
|
| 406 |
+
r"encoder.language_model.layers\.(?:[^.]+\.)*down_proj": r"decoder.layers\.(?:[^.]+\.)*down_proj",
|
| 407 |
+
"encoder.language_model.embed_tokens.weight": "decoder.embed_tokens.weight",
|
| 408 |
+
}
|
| 409 |
+
|
| 410 |
+
def __init__(self, config: ModilifyMk2Config):
|
| 411 |
+
super().__init__(config)
|
| 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):
|
| 418 |
+
return self.encoder
|
| 419 |
+
|
| 420 |
+
def get_decoder(self):
|
| 421 |
+
return self.decoder
|
| 422 |
+
|
| 423 |
+
def get_input_embeddings(self):
|
| 424 |
+
return self.encoder.get_input_embeddings()
|
| 425 |
+
|
| 426 |
+
def set_input_embeddings(self, value):
|
| 427 |
+
self.encoder.set_input_embeddings(value)
|
| 428 |
+
self.decoder.embed_tokens = value
|
| 429 |
+
|
| 430 |
+
def forward(
|
| 431 |
+
self,
|
| 432 |
+
*,
|
| 433 |
+
input_ids: torch.LongTensor | None = None,
|
| 434 |
+
attention_mask: torch.Tensor | dict | None = None,
|
| 435 |
+
past_key_values: Cache | None = None,
|
| 436 |
+
position_ids: torch.LongTensor | None = None,
|
| 437 |
+
decoder_input_ids: torch.LongTensor,
|
| 438 |
+
decoder_token_embeddings: torch.FloatTensor | None = None,
|
| 439 |
+
temporal_context_embeddings: torch.FloatTensor | None = None,
|
| 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,
|
| 448 |
+
persistent_bus: Any | None = None,
|
| 449 |
+
slot_identity: torch.Tensor | None = None,
|
| 450 |
+
**kwargs: Any,
|
| 451 |
+
) -> ModilifyMk2ModelOutput:
|
| 452 |
+
encoder_hidden = None
|
| 453 |
+
encoder_keys = ("pixel_values", "mm_token_type_ids", "image_position_ids", "inputs_embeds")
|
| 454 |
+
encoder_kwargs = {key: kwargs.pop(key) for key in encoder_keys if key in kwargs}
|
| 455 |
+
if input_ids is not None:
|
| 456 |
+
encoded = self.encoder(
|
| 457 |
+
input_ids=input_ids,
|
| 458 |
+
attention_mask=attention_mask,
|
| 459 |
+
past_key_values=past_key_values,
|
| 460 |
+
position_ids=position_ids,
|
| 461 |
+
**encoder_kwargs,
|
| 462 |
+
)
|
| 463 |
+
past_key_values = encoded.past_key_values
|
| 464 |
+
if return_encoder_outputs:
|
| 465 |
+
encoder_hidden = encoded.last_hidden_state
|
| 466 |
+
elif past_key_values is None:
|
| 467 |
+
raise ValueError("Either `input_ids` or `past_key_values` is required.")
|
| 468 |
+
decoded = self.decoder(
|
| 469 |
+
decoder_input_ids=decoder_input_ids,
|
| 470 |
+
decoder_token_embeddings=decoder_token_embeddings,
|
| 471 |
+
past_key_values=past_key_values,
|
| 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,
|
| 480 |
+
persistent_bus=persistent_bus,
|
| 481 |
+
slot_identity=slot_identity,
|
| 482 |
+
**kwargs,
|
| 483 |
+
)
|
| 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 |
+
|
| 493 |
+
class ModilifyMk2ForBlockDiffusion(DiffusionGemmaPreTrainedModel, ModilifyMk2GenerationMixin):
|
| 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)
|
| 508 |
+
self.model = ModilifyMk2Model(config)
|
| 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),
|
| 524 |
+
num_working_readers=(
|
| 525 |
+
sum(layer_type == "full_attention" for layer_type in layer_types)
|
| 526 |
+
if config.working_memory_bus else 0
|
| 527 |
+
),
|
| 528 |
+
num_persistent_readers=(
|
| 529 |
+
sum(layer_type == "full_attention" for layer_type in layer_types)
|
| 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 |
+
|
| 543 |
+
def _finalize_logits(self, hidden: torch.Tensor) -> torch.Tensor:
|
| 544 |
+
logits = self.lm_head(hidden)
|
| 545 |
+
return torch.tanh(logits / self.final_logit_softcapping) * self.final_logit_softcapping
|
| 546 |
+
|
| 547 |
+
def _prepare_latent_context(
|
| 548 |
+
self,
|
| 549 |
+
decoder_input_ids: torch.LongTensor,
|
| 550 |
+
*,
|
| 551 |
+
history: TrajectoryHistory | None,
|
| 552 |
+
tape: TrajectoryTape | None,
|
| 553 |
+
confidence: torch.Tensor | None,
|
| 554 |
+
entropy: torch.Tensor | None,
|
| 555 |
+
age: torch.Tensor | None,
|
| 556 |
+
latent_state: LatentDeliberationState | None,
|
| 557 |
+
) -> tuple[torch.Tensor, LatentDeliberationState, torch.Tensor, torch.Tensor]:
|
| 558 |
+
batch, canvas = decoder_input_ids.shape
|
| 559 |
+
dtype = self.model.decoder.embed_tokens.weight.dtype
|
| 560 |
+
if latent_state is None:
|
| 561 |
+
latent_state = LatentDeliberationState.empty(
|
| 562 |
+
batch_size=batch, canvas_length=canvas,
|
| 563 |
+
latent_dim=self.config.latent_dim, memory_slots=self.config.latent_memory_slots,
|
| 564 |
+
device=decoder_input_ids.device, dtype=dtype,
|
| 565 |
+
)
|
| 566 |
+
if history is None:
|
| 567 |
+
history = TrajectoryHistory.empty(
|
| 568 |
+
batch_size=batch,
|
| 569 |
+
canvas_length=canvas,
|
| 570 |
+
hidden_size=self.config.text_config.hidden_size,
|
| 571 |
+
history_length=self.config.latent_history_length,
|
| 572 |
+
device=decoder_input_ids.device,
|
| 573 |
+
dtype=dtype,
|
| 574 |
+
)
|
| 575 |
+
if tape is None:
|
| 576 |
+
tape = empty_trajectory_tape(
|
| 577 |
+
batch_size=batch,
|
| 578 |
+
config=self.config,
|
| 579 |
+
device=decoder_input_ids.device,
|
| 580 |
+
dtype=dtype,
|
| 581 |
+
)
|
| 582 |
+
confidence = latent_state.confidence if confidence is None else confidence.squeeze(-1).float()
|
| 583 |
+
entropy = latent_state.entropy if entropy is None else entropy.squeeze(-1).float()
|
| 584 |
+
if age is not None:
|
| 585 |
+
latent_state = replace(
|
| 586 |
+
latent_state,
|
| 587 |
+
age=age.to(device=decoder_input_ids.device, dtype=torch.int32),
|
| 588 |
+
)
|
| 589 |
+
token_embeddings = self.model.decoder.embed_tokens(decoder_input_ids)
|
| 590 |
+
processed: LatentProcessorOutput = self.latent_deliberation(
|
| 591 |
+
token_embeddings=token_embeddings,
|
| 592 |
+
confidence=confidence,
|
| 593 |
+
entropy=entropy,
|
| 594 |
+
state=latent_state,
|
| 595 |
+
history=history,
|
| 596 |
+
tape=tape,
|
| 597 |
+
)
|
| 598 |
+
return (
|
| 599 |
+
processed.context,
|
| 600 |
+
processed.state,
|
| 601 |
+
token_embeddings,
|
| 602 |
+
processed.history_projected,
|
| 603 |
+
)
|
| 604 |
+
|
| 605 |
+
def forward(
|
| 606 |
+
self,
|
| 607 |
+
*,
|
| 608 |
+
input_ids: torch.LongTensor | None = None,
|
| 609 |
+
attention_mask: torch.Tensor | dict | None = None,
|
| 610 |
+
past_key_values: Cache | None = None,
|
| 611 |
+
position_ids: torch.LongTensor | None = None,
|
| 612 |
+
decoder_input_ids: torch.LongTensor,
|
| 613 |
+
previous_confidence: torch.FloatTensor | None = None,
|
| 614 |
+
previous_entropy: torch.FloatTensor | None = None,
|
| 615 |
+
token_age: torch.Tensor | None = None,
|
| 616 |
+
latent_state: LatentDeliberationState | None = None,
|
| 617 |
+
history: TrajectoryHistory | None = None,
|
| 618 |
+
tape: TrajectoryTape | None = None,
|
| 619 |
+
history_hidden_state: TrajectoryHistory | torch.FloatTensor | None = None,
|
| 620 |
+
decoder_attention_mask: torch.Tensor | dict | None = None,
|
| 621 |
+
decoder_position_ids: torch.LongTensor | None = None,
|
| 622 |
+
return_encoder_outputs: bool = True,
|
| 623 |
+
compact_vocab: bool = False,
|
| 624 |
+
denoise_temperature: float | None = None,
|
| 625 |
+
repetition_token_mask: torch.BoolTensor | None = None,
|
| 626 |
+
repetition_penalty: float = 1.0,
|
| 627 |
+
sampling_generators: Sequence[torch.Generator] | None = None,
|
| 628 |
+
collect_latent_diagnostics: bool = True,
|
| 629 |
+
**kwargs: Any,
|
| 630 |
+
) -> ModilifyMk2BlockDiffusionOutput:
|
| 631 |
+
if history is None and isinstance(history_hidden_state, TrajectoryHistory):
|
| 632 |
+
history = history_hidden_state
|
| 633 |
+
(
|
| 634 |
+
latent_context,
|
| 635 |
+
next_state,
|
| 636 |
+
decoder_token_embeddings,
|
| 637 |
+
history_projected,
|
| 638 |
+
) = self._prepare_latent_context(
|
| 639 |
+
decoder_input_ids,
|
| 640 |
+
history=history,
|
| 641 |
+
tape=tape,
|
| 642 |
+
confidence=previous_confidence,
|
| 643 |
+
entropy=previous_entropy,
|
| 644 |
+
age=token_age,
|
| 645 |
+
latent_state=latent_state,
|
| 646 |
+
)
|
| 647 |
+
working_bus = self.latent_deliberation.working_memory_bus
|
| 648 |
+
persistent_bus = self.latent_deliberation.persistent_memory_bus
|
| 649 |
+
working_readers_active = working_bus.num_readers > 0
|
| 650 |
+
memory_readers_active = persistent_bus.num_readers > 0
|
| 651 |
+
outputs = self.model(
|
| 652 |
+
input_ids=input_ids, attention_mask=attention_mask,
|
| 653 |
+
past_key_values=past_key_values, position_ids=position_ids,
|
| 654 |
+
decoder_input_ids=decoder_input_ids,
|
| 655 |
+
decoder_token_embeddings=decoder_token_embeddings,
|
| 656 |
+
temporal_context_embeddings=latent_context,
|
| 657 |
+
decoder_attention_mask=decoder_attention_mask,
|
| 658 |
+
decoder_position_ids=decoder_position_ids,
|
| 659 |
+
return_encoder_outputs=return_encoder_outputs,
|
| 660 |
+
collect_latent_diagnostics=collect_latent_diagnostics,
|
| 661 |
+
working_bus=working_bus if working_readers_active else None,
|
| 662 |
+
working_state=latent_context if working_readers_active else None,
|
| 663 |
+
persistent_bus=persistent_bus if memory_readers_active else None,
|
| 664 |
+
memory_slots=next_state.memory_slots if memory_readers_active else None,
|
| 665 |
+
slot_identity=(
|
| 666 |
+
self.latent_deliberation.scaled_memory_slot_identity(
|
| 667 |
+
batch_size=decoder_input_ids.shape[0],
|
| 668 |
+
device=decoder_input_ids.device,
|
| 669 |
+
dtype=latent_context.dtype,
|
| 670 |
+
)
|
| 671 |
+
if memory_readers_active else None
|
| 672 |
+
),
|
| 673 |
+
**kwargs,
|
| 674 |
+
)
|
| 675 |
+
temperature = (
|
| 676 |
+
self.config.denoise_temperature
|
| 677 |
+
if denoise_temperature is None
|
| 678 |
+
else float(denoise_temperature)
|
| 679 |
+
)
|
| 680 |
+
proposal = proposal_confidence = token_entropy = None
|
| 681 |
+
greedy_proposal = greedy_confidence = None
|
| 682 |
+
logits = None
|
| 683 |
+
if compact_vocab:
|
| 684 |
+
(
|
| 685 |
+
proposal,
|
| 686 |
+
proposal_confidence,
|
| 687 |
+
token_entropy,
|
| 688 |
+
greedy_proposal,
|
| 689 |
+
greedy_confidence,
|
| 690 |
+
) = chunked_vocab_statistics(
|
| 691 |
+
outputs.last_hidden_state.detach(),
|
| 692 |
+
self.lm_head.weight.detach(),
|
| 693 |
+
softcap=self.final_logit_softcapping,
|
| 694 |
+
temperature=temperature,
|
| 695 |
+
chunk_size=self.config.vocab_chunk_size,
|
| 696 |
+
repetition_token_mask=repetition_token_mask,
|
| 697 |
+
repetition_penalty=repetition_penalty,
|
| 698 |
+
sampling_generators=sampling_generators,
|
| 699 |
+
)
|
| 700 |
+
else:
|
| 701 |
+
logits = self._finalize_logits(outputs.last_hidden_state)
|
| 702 |
+
return ModilifyMk2BlockDiffusionOutput(
|
| 703 |
+
logits=logits,
|
| 704 |
+
past_key_values=outputs.past_key_values,
|
| 705 |
+
encoder_last_hidden_state=outputs.encoder_last_hidden_state,
|
| 706 |
+
heavy_hidden_state=outputs.last_hidden_state,
|
| 707 |
+
next_latent_state=next_state,
|
| 708 |
+
temporal_context=latent_context,
|
| 709 |
+
history_projected=history_projected,
|
| 710 |
+
working_state=latent_context,
|
| 711 |
+
latent_residual_diagnostics=outputs.latent_residual_diagnostics,
|
| 712 |
+
proposal=proposal,
|
| 713 |
+
proposal_confidence=proposal_confidence,
|
| 714 |
+
token_entropy=token_entropy,
|
| 715 |
+
greedy_proposal=greedy_proposal,
|
| 716 |
+
greedy_confidence=greedy_confidence,
|
| 717 |
+
)
|
| 718 |
+
|
| 719 |
+
|
| 720 |
+
ModilifyMk2Model.register_for_auto_class("AutoModel")
|
| 721 |
+
ModilifyMk2ForBlockDiffusion.register_for_auto_class("AutoModelForCausalLM")
|
| 722 |
+
ModilifyMk2ForBlockDiffusion.register_for_auto_class("AutoModelForMultimodalLM")
|
| 723 |
+
|
| 724 |
+
|
| 725 |
+
__all__ = [
|
| 726 |
+
"ModilifyMk2BlockDiffusionOutput",
|
| 727 |
+
"ModilifyMk2Config",
|
| 728 |
+
"ModilifyMk2DecoderModel",
|
| 729 |
+
"ModilifyMk2EncoderModel",
|
| 730 |
+
"ModilifyMk2ForBlockDiffusion",
|
| 731 |
+
"ModilifyMk2Model",
|
| 732 |
+
"ModilifyMk2RMSNorm",
|
| 733 |
+
"ModilifyMk2TextRouter",
|
| 734 |
+
"install_modilify_mk2_trunk_semantics",
|
| 735 |
+
]
|
mps_ops.py
ADDED
|
@@ -0,0 +1,618 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 6 |
+
|
| 7 |
+
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,
|
| 18 |
+
top_k_index: torch.Tensor,
|
| 19 |
+
top_k_weights: torch.Tensor,
|
| 20 |
+
) -> torch.Tensor:
|
| 21 |
+
final_hidden_states = torch.zeros_like(hidden_states)
|
| 22 |
+
with torch.no_grad():
|
| 23 |
+
expert_mask = torch.nn.functional.one_hot(top_k_index, num_classes=self.num_experts)
|
| 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:
|
| 31 |
+
expert_idx = expert_idx[0]
|
| 32 |
+
if expert_idx == self.num_experts:
|
| 33 |
+
continue
|
| 34 |
+
top_k_pos, token_idx = torch.where(expert_mask[expert_idx])
|
| 35 |
+
current_state = hidden_states[token_idx]
|
| 36 |
+
gate_up = nn.functional.linear(current_state, self.gate_up_proj[expert_idx])
|
| 37 |
+
if hasattr(self, "lora_gate_up_a"):
|
| 38 |
+
update = nn.functional.linear(
|
| 39 |
+
nn.functional.linear(lora_dropout(current_state), self.lora_gate_up_a[expert_idx]),
|
| 40 |
+
self.lora_gate_up_b[expert_idx],
|
| 41 |
+
)
|
| 42 |
+
gate_up = gate_up + lora_scaling * update
|
| 43 |
+
gate, up = gate_up.chunk(2, dim=-1)
|
| 44 |
+
current_hidden_states = self.act_fn(gate) * up
|
| 45 |
+
expert_output = nn.functional.linear(current_hidden_states, self.down_proj[expert_idx])
|
| 46 |
+
if hasattr(self, "lora_down_a"):
|
| 47 |
+
update = nn.functional.linear(
|
| 48 |
+
nn.functional.linear(lora_dropout(current_hidden_states), self.lora_down_a[expert_idx]),
|
| 49 |
+
self.lora_down_b[expert_idx],
|
| 50 |
+
)
|
| 51 |
+
expert_output = expert_output + lora_scaling * update
|
| 52 |
+
current_hidden_states = expert_output * top_k_weights[token_idx, top_k_pos, None]
|
| 53 |
+
final_hidden_states.index_add_(0, token_idx, current_hidden_states.to(final_hidden_states.dtype))
|
| 54 |
+
|
| 55 |
+
return final_hidden_states
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
def grouped_expert_offsets(
|
| 59 |
+
sorted_expert_ids: torch.Tensor,
|
| 60 |
+
num_experts: int,
|
| 61 |
+
) -> torch.Tensor:
|
| 62 |
+
"""Build grouped-mm offsets without MPS ``histc`` host synchronization."""
|
| 63 |
+
counts = torch.bincount(sorted_expert_ids, minlength=num_experts)[:num_experts]
|
| 64 |
+
return torch.cumsum(counts, dim=0, dtype=torch.int32)
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
def mps_grouped_mm_experts_forward(
|
| 68 |
+
self: torch.nn.Module,
|
| 69 |
+
hidden_states: torch.Tensor,
|
| 70 |
+
top_k_index: torch.Tensor,
|
| 71 |
+
top_k_weights: torch.Tensor,
|
| 72 |
+
) -> torch.Tensor:
|
| 73 |
+
"""Transformers grouped-mm MoE with histogram-free MPS routing offsets."""
|
| 74 |
+
num_top_k = top_k_index.size(-1)
|
| 75 |
+
num_tokens = hidden_states.size(0)
|
| 76 |
+
hidden_dim = hidden_states.size(-1)
|
| 77 |
+
sample_weights = top_k_weights.reshape(-1)
|
| 78 |
+
expert_ids = top_k_index.reshape(-1)
|
| 79 |
+
|
| 80 |
+
expert_ids_g, perm = torch.sort(expert_ids)
|
| 81 |
+
selected_hidden_states_g = hidden_states[perm // num_top_k]
|
| 82 |
+
sample_weights_g = sample_weights[perm]
|
| 83 |
+
offsets = grouped_expert_offsets(expert_ids_g, self.num_experts)
|
| 84 |
+
|
| 85 |
+
sentinel_mask = (expert_ids_g >= self.num_experts).unsqueeze(-1)
|
| 86 |
+
expert_ids_g.clamp_(max=self.num_experts - 1)
|
| 87 |
+
selected_hidden_states_g.masked_fill_(sentinel_mask, 0.0)
|
| 88 |
+
|
| 89 |
+
gate_up_weights = self.gate_up_proj if self.has_gate else self.up_proj
|
| 90 |
+
gate_up_biases = (
|
| 91 |
+
self.gate_up_proj_bias[expert_ids_g]
|
| 92 |
+
if self.has_gate and self.has_bias
|
| 93 |
+
else self.up_proj_bias[expert_ids_g]
|
| 94 |
+
if self.has_bias
|
| 95 |
+
else None
|
| 96 |
+
)
|
| 97 |
+
projected = _grouped_linear(
|
| 98 |
+
selected_hidden_states_g,
|
| 99 |
+
gate_up_weights,
|
| 100 |
+
offsets,
|
| 101 |
+
bias=gate_up_biases,
|
| 102 |
+
is_transposed=self.is_transposed,
|
| 103 |
+
)
|
| 104 |
+
projected = self._apply_gate(projected) if self.has_gate else self.act_fn(projected)
|
| 105 |
+
|
| 106 |
+
down_biases = self.down_proj_bias[expert_ids_g] if self.has_bias else None
|
| 107 |
+
projected = _grouped_linear(
|
| 108 |
+
projected,
|
| 109 |
+
self.down_proj,
|
| 110 |
+
offsets,
|
| 111 |
+
bias=down_biases,
|
| 112 |
+
is_transposed=self.is_transposed,
|
| 113 |
+
)
|
| 114 |
+
weighted = projected * sample_weights_g.unsqueeze(-1)
|
| 115 |
+
weighted.masked_fill_(sentinel_mask, 0.0)
|
| 116 |
+
|
| 117 |
+
# Scatter directly back to the original top-k order. Constructing the
|
| 118 |
+
# inverse permutation and then gathering performs the same permutation
|
| 119 |
+
# with an extra index tensor and an extra read pass over `weighted`.
|
| 120 |
+
reordered = torch.empty_like(weighted)
|
| 121 |
+
reordered.index_copy_(0, perm, weighted)
|
| 122 |
+
weighted = reordered
|
| 123 |
+
return weighted.view(num_tokens, num_top_k, hidden_dim).sum(dim=1).to(hidden_states.dtype)
|
| 124 |
+
|
| 125 |
+
|
| 126 |
+
ALL_EXPERTS_FUNCTIONS.register("mps_grouped_mm", mps_grouped_mm_experts_forward)
|
| 127 |
+
|
| 128 |
+
|
| 129 |
+
class _IndependentExpertLoRA(torch.autograd.Function):
|
| 130 |
+
"""Capacity-padded reference path for independent expert LoRA."""
|
| 131 |
+
|
| 132 |
+
@staticmethod
|
| 133 |
+
def forward(
|
| 134 |
+
ctx,
|
| 135 |
+
inputs: torch.Tensor,
|
| 136 |
+
weight_a: torch.Tensor,
|
| 137 |
+
weight_b: torch.Tensor,
|
| 138 |
+
expert_ids: torch.LongTensor,
|
| 139 |
+
counts: torch.LongTensor,
|
| 140 |
+
capacity: int,
|
| 141 |
+
) -> torch.Tensor:
|
| 142 |
+
offsets = counts.cumsum(0)
|
| 143 |
+
starts = offsets - counts
|
| 144 |
+
slots = torch.arange(inputs.shape[0], device=inputs.device) - starts[expert_ids]
|
| 145 |
+
padded = inputs.new_zeros((weight_a.shape[0], capacity, inputs.shape[-1]))
|
| 146 |
+
padded[expert_ids, slots] = inputs
|
| 147 |
+
low_rank = torch.bmm(padded, weight_a.transpose(1, 2))
|
| 148 |
+
updates = torch.bmm(low_rank, weight_b.transpose(1, 2))
|
| 149 |
+
ctx.save_for_backward(
|
| 150 |
+
padded,
|
| 151 |
+
low_rank,
|
| 152 |
+
weight_a,
|
| 153 |
+
weight_b,
|
| 154 |
+
expert_ids,
|
| 155 |
+
slots,
|
| 156 |
+
)
|
| 157 |
+
return updates[expert_ids, slots]
|
| 158 |
+
|
| 159 |
+
@staticmethod
|
| 160 |
+
def backward(ctx, grad_output: torch.Tensor):
|
| 161 |
+
padded, low_rank, weight_a, weight_b, expert_ids, slots = ctx.saved_tensors
|
| 162 |
+
grad_padded = grad_output.new_zeros(
|
| 163 |
+
(weight_a.shape[0], padded.shape[1], grad_output.shape[-1])
|
| 164 |
+
)
|
| 165 |
+
grad_padded[expert_ids, slots] = grad_output
|
| 166 |
+
grad_weight_b = torch.bmm(grad_padded.transpose(1, 2), low_rank)
|
| 167 |
+
grad_low_rank = torch.bmm(grad_padded, weight_b)
|
| 168 |
+
grad_weight_a = torch.bmm(grad_low_rank.transpose(1, 2), padded)
|
| 169 |
+
grad_inputs_padded = torch.bmm(grad_low_rank, weight_a)
|
| 170 |
+
grad_inputs = grad_inputs_padded[expert_ids, slots]
|
| 171 |
+
return grad_inputs, grad_weight_a, grad_weight_b, None, None, None
|
| 172 |
+
|
| 173 |
+
|
| 174 |
+
_INDEPENDENT_EXPERT_LORA_METAL_SOURCE = r"""
|
| 175 |
+
#include <metal_stdlib>
|
| 176 |
+
using namespace metal;
|
| 177 |
+
|
| 178 |
+
constant uint kRank = 8;
|
| 179 |
+
constant uint kThreads = 256;
|
| 180 |
+
|
| 181 |
+
kernel void expert_lora_forward_a(
|
| 182 |
+
device const bfloat* inputs [[buffer(0)]],
|
| 183 |
+
device const bfloat* weight_a [[buffer(1)]],
|
| 184 |
+
device const long* expert_ids [[buffer(2)]],
|
| 185 |
+
device bfloat* low_rank [[buffer(3)]],
|
| 186 |
+
constant uint& row_count [[buffer(4)]],
|
| 187 |
+
constant uint& input_dim [[buffer(5)]],
|
| 188 |
+
uint row [[threadgroup_position_in_grid]],
|
| 189 |
+
uint lane [[thread_index_in_threadgroup]]) {
|
| 190 |
+
if (row >= row_count) return;
|
| 191 |
+
const uint expert = uint(expert_ids[row]);
|
| 192 |
+
threadgroup float partial[kThreads * kRank];
|
| 193 |
+
for (uint rank = 0; rank < kRank; ++rank) {
|
| 194 |
+
float value = 0.0f;
|
| 195 |
+
const uint a_base = (expert * kRank + rank) * input_dim;
|
| 196 |
+
const uint x_base = row * input_dim;
|
| 197 |
+
for (uint column = lane; column < input_dim; column += kThreads) {
|
| 198 |
+
value += float(inputs[x_base + column]) * float(weight_a[a_base + column]);
|
| 199 |
+
}
|
| 200 |
+
partial[lane * kRank + rank] = value;
|
| 201 |
+
}
|
| 202 |
+
threadgroup_barrier(mem_flags::mem_threadgroup);
|
| 203 |
+
for (uint stride = kThreads / 2; stride > 0; stride >>= 1) {
|
| 204 |
+
if (lane < stride) {
|
| 205 |
+
for (uint rank = 0; rank < kRank; ++rank) {
|
| 206 |
+
partial[lane * kRank + rank] +=
|
| 207 |
+
partial[(lane + stride) * kRank + rank];
|
| 208 |
+
}
|
| 209 |
+
}
|
| 210 |
+
threadgroup_barrier(mem_flags::mem_threadgroup);
|
| 211 |
+
}
|
| 212 |
+
if (lane == 0) {
|
| 213 |
+
for (uint rank = 0; rank < kRank; ++rank) {
|
| 214 |
+
low_rank[row * kRank + rank] = bfloat(partial[rank]);
|
| 215 |
+
}
|
| 216 |
+
}
|
| 217 |
+
}
|
| 218 |
+
|
| 219 |
+
kernel void expert_lora_forward_b(
|
| 220 |
+
device const bfloat* low_rank [[buffer(0)]],
|
| 221 |
+
device const bfloat* weight_b [[buffer(1)]],
|
| 222 |
+
device const long* expert_ids [[buffer(2)]],
|
| 223 |
+
device bfloat* output [[buffer(3)]],
|
| 224 |
+
constant uint& row_count [[buffer(4)]],
|
| 225 |
+
constant uint& output_dim [[buffer(5)]],
|
| 226 |
+
uint index [[thread_position_in_grid]]) {
|
| 227 |
+
const uint total = row_count * output_dim;
|
| 228 |
+
if (index >= total) return;
|
| 229 |
+
const uint row = index / output_dim;
|
| 230 |
+
const uint column = index - row * output_dim;
|
| 231 |
+
const uint expert = uint(expert_ids[row]);
|
| 232 |
+
const uint b_base = (expert * output_dim + column) * kRank;
|
| 233 |
+
float value = 0.0f;
|
| 234 |
+
for (uint rank = 0; rank < kRank; ++rank) {
|
| 235 |
+
value += float(low_rank[row * kRank + rank])
|
| 236 |
+
* float(weight_b[b_base + rank]);
|
| 237 |
+
}
|
| 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 |
+
|
| 354 |
+
_independent_expert_lora_metal_library = None
|
| 355 |
+
|
| 356 |
+
|
| 357 |
+
def _get_independent_expert_lora_metal_library():
|
| 358 |
+
global _independent_expert_lora_metal_library
|
| 359 |
+
if _independent_expert_lora_metal_library is None:
|
| 360 |
+
_independent_expert_lora_metal_library = torch.mps.compile_shader(
|
| 361 |
+
_INDEPENDENT_EXPERT_LORA_METAL_SOURCE
|
| 362 |
+
)
|
| 363 |
+
return _independent_expert_lora_metal_library
|
| 364 |
+
|
| 365 |
+
|
| 366 |
+
class _MetalIndependentExpertLoRA(torch.autograd.Function):
|
| 367 |
+
"""Independent rank-8 expert A/B with direct BF16 Metal forward/backward."""
|
| 368 |
+
|
| 369 |
+
@staticmethod
|
| 370 |
+
def forward(
|
| 371 |
+
ctx,
|
| 372 |
+
inputs: torch.Tensor,
|
| 373 |
+
weight_a: torch.Tensor,
|
| 374 |
+
weight_b: torch.Tensor,
|
| 375 |
+
expert_ids: torch.LongTensor,
|
| 376 |
+
counts: torch.LongTensor,
|
| 377 |
+
) -> torch.Tensor:
|
| 378 |
+
inputs = inputs.contiguous()
|
| 379 |
+
weight_a = weight_a.contiguous()
|
| 380 |
+
weight_b = weight_b.contiguous()
|
| 381 |
+
expert_ids = expert_ids.contiguous()
|
| 382 |
+
expert_count, rank, input_dim = weight_a.shape
|
| 383 |
+
row_count = inputs.shape[0]
|
| 384 |
+
output_dim = weight_b.shape[1]
|
| 385 |
+
if rank != 8:
|
| 386 |
+
raise ValueError("The direct Metal expert LoRA kernel requires rank 8.")
|
| 387 |
+
offsets = torch.cat(
|
| 388 |
+
(
|
| 389 |
+
counts.new_zeros(1, dtype=torch.int64),
|
| 390 |
+
counts.cumsum(0, dtype=torch.int64),
|
| 391 |
+
)
|
| 392 |
+
).contiguous()
|
| 393 |
+
low_rank = inputs.new_empty((row_count, rank))
|
| 394 |
+
output = inputs.new_empty((row_count, output_dim))
|
| 395 |
+
library = _get_independent_expert_lora_metal_library()
|
| 396 |
+
library.expert_lora_forward_a(
|
| 397 |
+
inputs,
|
| 398 |
+
weight_a,
|
| 399 |
+
expert_ids,
|
| 400 |
+
low_rank,
|
| 401 |
+
row_count,
|
| 402 |
+
input_dim,
|
| 403 |
+
threads=(row_count * 256, 1, 1),
|
| 404 |
+
group_size=(256, 1, 1),
|
| 405 |
+
)
|
| 406 |
+
library.expert_lora_forward_b(
|
| 407 |
+
low_rank,
|
| 408 |
+
weight_b,
|
| 409 |
+
expert_ids,
|
| 410 |
+
output,
|
| 411 |
+
row_count,
|
| 412 |
+
output_dim,
|
| 413 |
+
threads=(row_count * output_dim, 1, 1),
|
| 414 |
+
group_size=(256, 1, 1),
|
| 415 |
+
)
|
| 416 |
+
ctx.save_for_backward(
|
| 417 |
+
inputs,
|
| 418 |
+
low_rank,
|
| 419 |
+
weight_a,
|
| 420 |
+
weight_b,
|
| 421 |
+
expert_ids,
|
| 422 |
+
offsets,
|
| 423 |
+
)
|
| 424 |
+
ctx.dimensions = (expert_count, row_count, input_dim, output_dim)
|
| 425 |
+
return output
|
| 426 |
+
|
| 427 |
+
@staticmethod
|
| 428 |
+
def backward(ctx, grad_output: torch.Tensor):
|
| 429 |
+
inputs, low_rank, weight_a, weight_b, expert_ids, offsets = ctx.saved_tensors
|
| 430 |
+
expert_count, row_count, input_dim, output_dim = ctx.dimensions
|
| 431 |
+
grad_output = grad_output.contiguous()
|
| 432 |
+
grad_low_rank = low_rank.new_empty(low_rank.shape)
|
| 433 |
+
grad_inputs = inputs.new_empty(inputs.shape)
|
| 434 |
+
grad_weight_a = weight_a.new_empty(weight_a.shape)
|
| 435 |
+
grad_weight_b = weight_b.new_empty(weight_b.shape)
|
| 436 |
+
library = _get_independent_expert_lora_metal_library()
|
| 437 |
+
library.expert_lora_backward_low_rank(
|
| 438 |
+
grad_output,
|
| 439 |
+
weight_b,
|
| 440 |
+
expert_ids,
|
| 441 |
+
grad_low_rank,
|
| 442 |
+
row_count,
|
| 443 |
+
output_dim,
|
| 444 |
+
threads=(row_count * 256, 1, 1),
|
| 445 |
+
group_size=(256, 1, 1),
|
| 446 |
+
)
|
| 447 |
+
library.expert_lora_backward_inputs(
|
| 448 |
+
grad_low_rank,
|
| 449 |
+
weight_a,
|
| 450 |
+
expert_ids,
|
| 451 |
+
grad_inputs,
|
| 452 |
+
row_count,
|
| 453 |
+
input_dim,
|
| 454 |
+
threads=(row_count * input_dim, 1, 1),
|
| 455 |
+
group_size=(256, 1, 1),
|
| 456 |
+
)
|
| 457 |
+
library.expert_lora_backward_a(
|
| 458 |
+
grad_low_rank,
|
| 459 |
+
inputs,
|
| 460 |
+
offsets,
|
| 461 |
+
grad_weight_a,
|
| 462 |
+
expert_count,
|
| 463 |
+
input_dim,
|
| 464 |
+
threads=(expert_count * 8 * input_dim, 1, 1),
|
| 465 |
+
group_size=(256, 1, 1),
|
| 466 |
+
)
|
| 467 |
+
library.expert_lora_backward_b(
|
| 468 |
+
grad_output,
|
| 469 |
+
low_rank,
|
| 470 |
+
offsets,
|
| 471 |
+
grad_weight_b,
|
| 472 |
+
expert_count,
|
| 473 |
+
output_dim,
|
| 474 |
+
threads=(expert_count * output_dim * 8, 1, 1),
|
| 475 |
+
group_size=(256, 1, 1),
|
| 476 |
+
)
|
| 477 |
+
return grad_inputs, grad_weight_a, grad_weight_b, None, None
|
| 478 |
+
|
| 479 |
+
|
| 480 |
+
def _independent_expert_lora(
|
| 481 |
+
inputs: torch.Tensor,
|
| 482 |
+
weight_a: torch.Tensor,
|
| 483 |
+
weight_b: torch.Tensor,
|
| 484 |
+
expert_ids: torch.LongTensor,
|
| 485 |
+
counts: torch.LongTensor,
|
| 486 |
+
) -> torch.Tensor:
|
| 487 |
+
"""Dispatch the compact independent-expert operator without sharing A/B."""
|
| 488 |
+
|
| 489 |
+
if inputs.shape[0] == 0:
|
| 490 |
+
return inputs.new_empty((0, weight_b.shape[1]))
|
| 491 |
+
if (
|
| 492 |
+
inputs.device.type == "mps"
|
| 493 |
+
and inputs.dtype == torch.bfloat16
|
| 494 |
+
and weight_a.dtype == torch.bfloat16
|
| 495 |
+
and weight_b.dtype == torch.bfloat16
|
| 496 |
+
and weight_a.shape[1] == 8
|
| 497 |
+
and hasattr(torch.mps, "compile_shader")
|
| 498 |
+
):
|
| 499 |
+
return _MetalIndependentExpertLoRA.apply(
|
| 500 |
+
inputs,
|
| 501 |
+
weight_a,
|
| 502 |
+
weight_b,
|
| 503 |
+
expert_ids,
|
| 504 |
+
counts,
|
| 505 |
+
)
|
| 506 |
+
return _IndependentExpertLoRA.apply(
|
| 507 |
+
inputs,
|
| 508 |
+
weight_a,
|
| 509 |
+
weight_b,
|
| 510 |
+
expert_ids,
|
| 511 |
+
counts,
|
| 512 |
+
int(counts.max().detach().cpu()),
|
| 513 |
+
)
|
| 514 |
+
|
| 515 |
+
|
| 516 |
+
def mps_segmented_experts_forward(
|
| 517 |
+
self: torch.nn.Module,
|
| 518 |
+
hidden_states: torch.Tensor,
|
| 519 |
+
top_k_index: torch.Tensor,
|
| 520 |
+
top_k_weights: torch.Tensor,
|
| 521 |
+
) -> torch.Tensor:
|
| 522 |
+
"""Run GPU-grouped base experts plus independent Metal LoRA A/B."""
|
| 523 |
+
if hidden_states.device.type != "mps":
|
| 524 |
+
return lora_aware_eager_experts_forward(
|
| 525 |
+
self,
|
| 526 |
+
hidden_states,
|
| 527 |
+
top_k_index,
|
| 528 |
+
top_k_weights,
|
| 529 |
+
)
|
| 530 |
+
|
| 531 |
+
num_top_k = top_k_index.size(-1)
|
| 532 |
+
num_tokens, hidden_dim = hidden_states.shape
|
| 533 |
+
sample_weights = top_k_weights.reshape(-1)
|
| 534 |
+
expert_ids = top_k_index.reshape(-1)
|
| 535 |
+
expert_ids_g, permutation = torch.sort(expert_ids)
|
| 536 |
+
selected_hidden_g = hidden_states[permutation // num_top_k]
|
| 537 |
+
sample_weights_g = sample_weights[permutation]
|
| 538 |
+
|
| 539 |
+
counts = torch.bincount(
|
| 540 |
+
expert_ids_g,
|
| 541 |
+
minlength=self.num_experts,
|
| 542 |
+
)[: self.num_experts]
|
| 543 |
+
offsets = torch.cumsum(counts, dim=0, dtype=torch.int32)
|
| 544 |
+
gate_up_weights = self.gate_up_proj if self.has_gate else self.up_proj
|
| 545 |
+
gate_up_bias = None
|
| 546 |
+
if self.has_bias:
|
| 547 |
+
packed_bias = self.gate_up_proj_bias if self.has_gate else self.up_proj_bias
|
| 548 |
+
gate_up_bias = packed_bias[expert_ids_g]
|
| 549 |
+
gate_up_g = _grouped_linear(
|
| 550 |
+
selected_hidden_g,
|
| 551 |
+
gate_up_weights,
|
| 552 |
+
offsets,
|
| 553 |
+
bias=gate_up_bias,
|
| 554 |
+
is_transposed=self.is_transposed,
|
| 555 |
+
)
|
| 556 |
+
if hasattr(self, "lora_gate_up_a") and self.training:
|
| 557 |
+
gate_up_update = _independent_expert_lora(
|
| 558 |
+
self.lora_dropout(selected_hidden_g),
|
| 559 |
+
self.lora_gate_up_a,
|
| 560 |
+
self.lora_gate_up_b,
|
| 561 |
+
expert_ids_g,
|
| 562 |
+
counts,
|
| 563 |
+
)
|
| 564 |
+
gate_up_g.add_(gate_up_update, alpha=self.lora_scaling)
|
| 565 |
+
elif hasattr(self, "lora_gate_up_a"):
|
| 566 |
+
gate_up_update = _independent_expert_lora(
|
| 567 |
+
selected_hidden_g,
|
| 568 |
+
self.lora_gate_up_a,
|
| 569 |
+
self.lora_gate_up_b,
|
| 570 |
+
expert_ids_g,
|
| 571 |
+
counts,
|
| 572 |
+
)
|
| 573 |
+
gate_up_g.add_(gate_up_update, alpha=self.lora_scaling)
|
| 574 |
+
activated_g = (
|
| 575 |
+
self._apply_gate(gate_up_g)
|
| 576 |
+
if self.has_gate
|
| 577 |
+
else self.act_fn(gate_up_g)
|
| 578 |
+
)
|
| 579 |
+
down_bias = self.down_proj_bias[expert_ids_g] if self.has_bias else None
|
| 580 |
+
projected_g = _grouped_linear(
|
| 581 |
+
activated_g,
|
| 582 |
+
self.down_proj,
|
| 583 |
+
offsets,
|
| 584 |
+
bias=down_bias,
|
| 585 |
+
is_transposed=self.is_transposed,
|
| 586 |
+
)
|
| 587 |
+
if hasattr(self, "lora_down_a") and self.training:
|
| 588 |
+
down_update = _independent_expert_lora(
|
| 589 |
+
self.lora_dropout(activated_g),
|
| 590 |
+
self.lora_down_a,
|
| 591 |
+
self.lora_down_b,
|
| 592 |
+
expert_ids_g,
|
| 593 |
+
counts,
|
| 594 |
+
)
|
| 595 |
+
projected_g.add_(down_update, alpha=self.lora_scaling)
|
| 596 |
+
elif hasattr(self, "lora_down_a"):
|
| 597 |
+
down_update = _independent_expert_lora(
|
| 598 |
+
activated_g,
|
| 599 |
+
self.lora_down_a,
|
| 600 |
+
self.lora_down_b,
|
| 601 |
+
expert_ids_g,
|
| 602 |
+
counts,
|
| 603 |
+
)
|
| 604 |
+
projected_g.add_(down_update, alpha=self.lora_scaling)
|
| 605 |
+
weighted_g = projected_g * sample_weights_g[:, None]
|
| 606 |
+
weighted = torch.empty_like(weighted_g)
|
| 607 |
+
weighted.index_copy_(0, permutation, weighted_g)
|
| 608 |
+
return weighted.view(num_tokens, num_top_k, hidden_dim).sum(dim=1).to(hidden_states.dtype)
|
| 609 |
+
|
| 610 |
+
|
| 611 |
+
ALL_EXPERTS_FUNCTIONS.register("mps_segmented", mps_segmented_experts_forward)
|
| 612 |
+
|
| 613 |
+
|
| 614 |
+
__all__ = [
|
| 615 |
+
"grouped_expert_offsets",
|
| 616 |
+
"mps_grouped_mm_experts_forward",
|
| 617 |
+
"mps_segmented_experts_forward",
|
| 618 |
+
]
|
processor_config.json
ADDED
|
@@ -0,0 +1,75 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"audio_ms_per_token": 40,
|
| 3 |
+
"audio_seq_length": 750,
|
| 4 |
+
"feature_extractor": {
|
| 5 |
+
"dither": 0.0,
|
| 6 |
+
"feature_extractor_type": "Gemma4AudioFeatureExtractor",
|
| 7 |
+
"feature_size": 128,
|
| 8 |
+
"fft_length": 512,
|
| 9 |
+
"fft_overdrive": false,
|
| 10 |
+
"frame_length": 320,
|
| 11 |
+
"hop_length": 160,
|
| 12 |
+
"input_scale_factor": 1.0,
|
| 13 |
+
"max_frequency": 8000.0,
|
| 14 |
+
"mel_floor": 0.001,
|
| 15 |
+
"min_frequency": 0.0,
|
| 16 |
+
"padding_side": "right",
|
| 17 |
+
"padding_value": 0.0,
|
| 18 |
+
"per_bin_mean": null,
|
| 19 |
+
"per_bin_stddev": null,
|
| 20 |
+
"preemphasis": 0.0,
|
| 21 |
+
"preemphasis_htk_flavor": true,
|
| 22 |
+
"return_attention_mask": true,
|
| 23 |
+
"sampling_rate": 16000
|
| 24 |
+
},
|
| 25 |
+
"image_processor": {
|
| 26 |
+
"do_convert_rgb": true,
|
| 27 |
+
"do_normalize": false,
|
| 28 |
+
"do_rescale": true,
|
| 29 |
+
"do_resize": true,
|
| 30 |
+
"image_mean": [
|
| 31 |
+
0.0,
|
| 32 |
+
0.0,
|
| 33 |
+
0.0
|
| 34 |
+
],
|
| 35 |
+
"image_processor_type": "Gemma4ImageProcessor",
|
| 36 |
+
"image_seq_length": 280,
|
| 37 |
+
"image_std": [
|
| 38 |
+
1.0,
|
| 39 |
+
1.0,
|
| 40 |
+
1.0
|
| 41 |
+
],
|
| 42 |
+
"max_soft_tokens": 280,
|
| 43 |
+
"patch_size": 16,
|
| 44 |
+
"pooling_kernel_size": 3,
|
| 45 |
+
"resample": 3,
|
| 46 |
+
"rescale_factor": 0.00392156862745098
|
| 47 |
+
},
|
| 48 |
+
"image_seq_length": 280,
|
| 49 |
+
"processor_class": "Gemma4Processor",
|
| 50 |
+
"video_processor": {
|
| 51 |
+
"do_convert_rgb": true,
|
| 52 |
+
"do_normalize": true,
|
| 53 |
+
"do_rescale": true,
|
| 54 |
+
"do_resize": true,
|
| 55 |
+
"do_sample_frames": true,
|
| 56 |
+
"image_mean": [
|
| 57 |
+
0.0,
|
| 58 |
+
0.0,
|
| 59 |
+
0.0
|
| 60 |
+
],
|
| 61 |
+
"image_std": [
|
| 62 |
+
1.0,
|
| 63 |
+
1.0,
|
| 64 |
+
1.0
|
| 65 |
+
],
|
| 66 |
+
"max_soft_tokens": 70,
|
| 67 |
+
"num_frames": 32,
|
| 68 |
+
"patch_size": 16,
|
| 69 |
+
"pooling_kernel_size": 3,
|
| 70 |
+
"resample": 3,
|
| 71 |
+
"rescale_factor": 0.00392156862745098,
|
| 72 |
+
"return_metadata": false,
|
| 73 |
+
"video_processor_type": "Gemma4VideoProcessor"
|
| 74 |
+
}
|
| 75 |
+
}
|
tokenizer.json
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:cc8d3a0ce36466ccc1278bf987df5f71db1719b9ca6b4118264f45cb627bfe0f
|
| 3 |
+
size 32169626
|
tokenizer_config.json
ADDED
|
@@ -0,0 +1,96 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"audio_token": "<|audio|>",
|
| 3 |
+
"backend": "tokenizers",
|
| 4 |
+
"boa_token": "<|audio>",
|
| 5 |
+
"boi_token": "<|image>",
|
| 6 |
+
"bos_token": "<bos>",
|
| 7 |
+
"eoa_token": "<audio|>",
|
| 8 |
+
"eoc_token": "<channel|>",
|
| 9 |
+
"eoi_token": "<image|>",
|
| 10 |
+
"eos_token": "<eos>",
|
| 11 |
+
"eot_token": "<turn|>",
|
| 12 |
+
"escape_token": "<|\"|>",
|
| 13 |
+
"etc_token": "<tool_call|>",
|
| 14 |
+
"etd_token": "<tool|>",
|
| 15 |
+
"etr_token": "<tool_response|>",
|
| 16 |
+
"extra_special_tokens": [
|
| 17 |
+
"<|video|>"
|
| 18 |
+
],
|
| 19 |
+
"image_token": "<|image|>",
|
| 20 |
+
"is_local": true,
|
| 21 |
+
"local_files_only": false,
|
| 22 |
+
"mask_token": "<mask>",
|
| 23 |
+
"model_max_length": 1000000000000000019884624838656,
|
| 24 |
+
"model_specific_special_tokens": {
|
| 25 |
+
"audio_token": "<|audio|>",
|
| 26 |
+
"boa_token": "<|audio>",
|
| 27 |
+
"boi_token": "<|image>",
|
| 28 |
+
"eoa_token": "<audio|>",
|
| 29 |
+
"eoc_token": "<channel|>",
|
| 30 |
+
"eoi_token": "<image|>",
|
| 31 |
+
"eot_token": "<turn|>",
|
| 32 |
+
"escape_token": "<|\"|>",
|
| 33 |
+
"etc_token": "<tool_call|>",
|
| 34 |
+
"etd_token": "<tool|>",
|
| 35 |
+
"etr_token": "<tool_response|>",
|
| 36 |
+
"image_token": "<|image|>",
|
| 37 |
+
"soc_token": "<|channel>",
|
| 38 |
+
"sot_token": "<|turn>",
|
| 39 |
+
"stc_token": "<|tool_call>",
|
| 40 |
+
"std_token": "<|tool>",
|
| 41 |
+
"str_token": "<|tool_response>",
|
| 42 |
+
"think_token": "<|think|>"
|
| 43 |
+
},
|
| 44 |
+
"pad_token": "<pad>",
|
| 45 |
+
"padding_side": "left",
|
| 46 |
+
"processor_class": "Gemma4Processor",
|
| 47 |
+
"response_schema": {
|
| 48 |
+
"properties": {
|
| 49 |
+
"content": {
|
| 50 |
+
"type": "string"
|
| 51 |
+
},
|
| 52 |
+
"role": {
|
| 53 |
+
"const": "assistant"
|
| 54 |
+
},
|
| 55 |
+
"thinking": {
|
| 56 |
+
"type": "string"
|
| 57 |
+
},
|
| 58 |
+
"tool_calls": {
|
| 59 |
+
"items": {
|
| 60 |
+
"properties": {
|
| 61 |
+
"function": {
|
| 62 |
+
"properties": {
|
| 63 |
+
"arguments": {
|
| 64 |
+
"additionalProperties": {},
|
| 65 |
+
"type": "object",
|
| 66 |
+
"x-parser": "gemma4-tool-call"
|
| 67 |
+
},
|
| 68 |
+
"name": {
|
| 69 |
+
"type": "string"
|
| 70 |
+
}
|
| 71 |
+
},
|
| 72 |
+
"type": "object",
|
| 73 |
+
"x-regex": "call\\:(?P<name>\\w+)(?P<arguments>\\{.*\\})"
|
| 74 |
+
},
|
| 75 |
+
"type": {
|
| 76 |
+
"const": "function"
|
| 77 |
+
}
|
| 78 |
+
},
|
| 79 |
+
"type": "object"
|
| 80 |
+
},
|
| 81 |
+
"type": "array",
|
| 82 |
+
"x-regex-iterator": "<\\|tool_call>(.*?)<tool_call\\|>"
|
| 83 |
+
}
|
| 84 |
+
},
|
| 85 |
+
"type": "object",
|
| 86 |
+
"x-regex": "(\\<\\|channel\\>thought\\n(?P<thinking>.*?)\\<channel\\|\\>)?(?P<tool_calls>\\<\\|tool_call\\>.*\\<tool_call\\|\\>)?(?P<content>(?:(?!\\<turn\\|\\>)(?!\\<\\|tool_response\\>).)+)?(?:\\<turn\\|\\>|\\<\\|tool_response\\>)?"
|
| 87 |
+
},
|
| 88 |
+
"soc_token": "<|channel>",
|
| 89 |
+
"sot_token": "<|turn>",
|
| 90 |
+
"stc_token": "<|tool_call>",
|
| 91 |
+
"std_token": "<|tool>",
|
| 92 |
+
"str_token": "<|tool_response>",
|
| 93 |
+
"think_token": "<|think|>",
|
| 94 |
+
"tokenizer_class": "GemmaTokenizer",
|
| 95 |
+
"unk_token": "<unk>"
|
| 96 |
+
}
|
vocab_ops.py
ADDED
|
@@ -0,0 +1,525 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Copyright 2026 Modilify
|
| 2 |
+
# SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.0
|
| 3 |
+
"""Exact memory-bounded vocabulary projection and sampling."""
|
| 4 |
+
|
| 5 |
+
from __future__ import annotations
|
| 6 |
+
|
| 7 |
+
import math
|
| 8 |
+
from collections.abc import Sequence
|
| 9 |
+
|
| 10 |
+
import torch
|
| 11 |
+
from torch.nn import functional as F
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def _stable_max(values: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
| 15 |
+
"""Return ``max`` indices that stay in-range on MPS NaN/-inf rows."""
|
| 16 |
+
|
| 17 |
+
best, index = values.max(dim=-1)
|
| 18 |
+
width = values.shape[-1]
|
| 19 |
+
if width <= 0:
|
| 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,
|
| 342 |
+
weight: torch.Tensor,
|
| 343 |
+
*,
|
| 344 |
+
softcap: float,
|
| 345 |
+
temperature: float,
|
| 346 |
+
chunk_size: int,
|
| 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,
|
| 353 |
+
torch.Tensor,
|
| 354 |
+
torch.LongTensor,
|
| 355 |
+
torch.Tensor,
|
| 356 |
+
]:
|
| 357 |
+
"""Return exact inference statistics, optionally penalizing seen tokens."""
|
| 358 |
+
|
| 359 |
+
if not math.isfinite(repetition_penalty) or repetition_penalty <= 0:
|
| 360 |
+
raise ValueError("`repetition_penalty` must be a finite positive number.")
|
| 361 |
+
|
| 362 |
+
# Preserve the original custom-autograd path exactly when disabled. In
|
| 363 |
+
# particular, this keeps the same chunk-local RNG draws and default output.
|
| 364 |
+
if repetition_penalty != 1.0 or sampling_generators is not None:
|
| 365 |
+
if hidden.ndim < 2 or weight.ndim != 2:
|
| 366 |
+
raise ValueError("Chunked vocabulary tensors must have matrix features.")
|
| 367 |
+
if hidden.shape[-1] != weight.shape[1]:
|
| 368 |
+
raise ValueError("Hidden and vocabulary projection dimensions differ.")
|
| 369 |
+
if temperature <= 0 or chunk_size <= 0:
|
| 370 |
+
raise ValueError("Temperature and vocabulary chunk size must be positive.")
|
| 371 |
+
expected_mask_shape = (hidden.shape[0], weight.shape[0])
|
| 372 |
+
if repetition_penalty != 1.0:
|
| 373 |
+
if repetition_token_mask is None or repetition_token_mask.shape != expected_mask_shape:
|
| 374 |
+
raise ValueError(
|
| 375 |
+
"`repetition_token_mask` must have shape [batch, vocabulary]."
|
| 376 |
+
)
|
| 377 |
+
if repetition_token_mask.dtype != torch.bool:
|
| 378 |
+
raise ValueError("`repetition_token_mask` must be a boolean tensor.")
|
| 379 |
+
if sampling_generators is not None and len(sampling_generators) != hidden.shape[0]:
|
| 380 |
+
raise ValueError(
|
| 381 |
+
"`sampling_generators` must contain one generator per batch row."
|
| 382 |
+
)
|
| 383 |
+
|
| 384 |
+
output_shape = hidden.shape[:-1]
|
| 385 |
+
flat_hidden = hidden.reshape(-1, hidden.shape[-1])
|
| 386 |
+
rows_per_batch = math.prod(hidden.shape[1:-1])
|
| 387 |
+
batch_indices = torch.arange(
|
| 388 |
+
flat_hidden.shape[0], device=hidden.device
|
| 389 |
+
).div(rows_per_batch, rounding_mode="floor")
|
| 390 |
+
sample_log_z = torch.full(
|
| 391 |
+
(flat_hidden.shape[0],),
|
| 392 |
+
-torch.inf,
|
| 393 |
+
device=hidden.device,
|
| 394 |
+
dtype=torch.float32,
|
| 395 |
+
)
|
| 396 |
+
best_gumbel = torch.full_like(sample_log_z, -torch.inf)
|
| 397 |
+
selected_score = torch.zeros_like(sample_log_z)
|
| 398 |
+
selected = torch.zeros(
|
| 399 |
+
flat_hidden.shape[0], device=hidden.device, dtype=torch.long
|
| 400 |
+
)
|
| 401 |
+
greedy_score = torch.full_like(sample_log_z, -torch.inf)
|
| 402 |
+
greedy = torch.zeros_like(selected)
|
| 403 |
+
moment_max = torch.full_like(sample_log_z, -torch.inf)
|
| 404 |
+
moment_sum = torch.zeros_like(sample_log_z)
|
| 405 |
+
moment_weighted = torch.zeros_like(sample_log_z)
|
| 406 |
+
|
| 407 |
+
for start in range(0, weight.shape[0], int(chunk_size)):
|
| 408 |
+
stop = min(start + int(chunk_size), weight.shape[0])
|
| 409 |
+
raw = F.linear(flat_hidden, weight[start:stop])
|
| 410 |
+
scores = torch.tanh(raw.float() / float(softcap)) * float(softcap)
|
| 411 |
+
if repetition_penalty != 1.0:
|
| 412 |
+
assert repetition_token_mask is not None
|
| 413 |
+
seen = repetition_token_mask[:, start:stop].index_select(
|
| 414 |
+
0, batch_indices
|
| 415 |
+
)
|
| 416 |
+
penalized = torch.where(
|
| 417 |
+
scores < 0,
|
| 418 |
+
scores * float(repetition_penalty),
|
| 419 |
+
scores / float(repetition_penalty),
|
| 420 |
+
)
|
| 421 |
+
scores = torch.where(seen, penalized, scores)
|
| 422 |
+
sample_scores = scores / float(temperature)
|
| 423 |
+
sample_log_z = torch.logaddexp(
|
| 424 |
+
sample_log_z, torch.logsumexp(sample_scores, dim=-1)
|
| 425 |
+
)
|
| 426 |
+
|
| 427 |
+
if sampling_generators is None:
|
| 428 |
+
uniform = torch.rand(
|
| 429 |
+
sample_scores.shape,
|
| 430 |
+
device=sample_scores.device,
|
| 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],
|
| 439 |
+
)
|
| 440 |
+
uniform = torch.cat(
|
| 441 |
+
[
|
| 442 |
+
torch.rand(
|
| 443 |
+
per_request_shape,
|
| 444 |
+
device=sample_scores.device,
|
| 445 |
+
dtype=torch.float32,
|
| 446 |
+
generator=generator,
|
| 447 |
+
)
|
| 448 |
+
for generator in sampling_generators
|
| 449 |
+
],
|
| 450 |
+
dim=0,
|
| 451 |
+
)
|
| 452 |
+
uniform = uniform.clamp_(
|
| 453 |
+
min=torch.finfo(torch.float32).tiny,
|
| 454 |
+
max=1.0 - torch.finfo(torch.float32).eps,
|
| 455 |
+
)
|
| 456 |
+
gumbel_scores = sample_scores - torch.log(-torch.log(uniform))
|
| 457 |
+
chunk_best, chunk_index = _stable_max(gumbel_scores)
|
| 458 |
+
replace_best = chunk_best.gt(best_gumbel)
|
| 459 |
+
candidate_score = sample_scores.gather(
|
| 460 |
+
1, chunk_index[:, None]
|
| 461 |
+
).squeeze(-1)
|
| 462 |
+
best_gumbel = torch.maximum(best_gumbel, chunk_best)
|
| 463 |
+
selected = torch.where(replace_best, chunk_index + start, selected)
|
| 464 |
+
selected_score = torch.where(
|
| 465 |
+
replace_best, candidate_score, selected_score
|
| 466 |
+
)
|
| 467 |
+
|
| 468 |
+
chunk_max, chunk_argmax = _stable_max(sample_scores)
|
| 469 |
+
replace_greedy = chunk_max.gt(greedy_score)
|
| 470 |
+
greedy_score = torch.maximum(greedy_score, chunk_max)
|
| 471 |
+
greedy = torch.where(replace_greedy, chunk_argmax + start, greedy)
|
| 472 |
+
shifted = torch.exp(sample_scores - chunk_max[:, None])
|
| 473 |
+
chunk_sum = shifted.sum(dim=-1)
|
| 474 |
+
chunk_weighted = (shifted * sample_scores).sum(dim=-1)
|
| 475 |
+
merged_max = torch.maximum(moment_max, chunk_max)
|
| 476 |
+
previous_scale = torch.exp(moment_max - merged_max)
|
| 477 |
+
chunk_scale = torch.exp(chunk_max - merged_max)
|
| 478 |
+
moment_sum = moment_sum * previous_scale + chunk_sum * chunk_scale
|
| 479 |
+
moment_weighted = (
|
| 480 |
+
moment_weighted * previous_scale + chunk_weighted * chunk_scale
|
| 481 |
+
)
|
| 482 |
+
moment_max = merged_max
|
| 483 |
+
|
| 484 |
+
confidence = torch.exp(selected_score - sample_log_z).clamp_(0.0, 1.0)
|
| 485 |
+
greedy_confidence = torch.exp(greedy_score - sample_log_z).clamp_(0.0, 1.0)
|
| 486 |
+
entropy = sample_log_z - moment_weighted / moment_sum.clamp_min(
|
| 487 |
+
torch.finfo(torch.float32).tiny
|
| 488 |
+
)
|
| 489 |
+
return (
|
| 490 |
+
selected.view(output_shape),
|
| 491 |
+
confidence.view(output_shape),
|
| 492 |
+
entropy.view(output_shape),
|
| 493 |
+
greedy.view(output_shape),
|
| 494 |
+
greedy_confidence.view(output_shape),
|
| 495 |
+
)
|
| 496 |
+
|
| 497 |
+
labels = torch.full(
|
| 498 |
+
hidden.shape[:-1],
|
| 499 |
+
-100,
|
| 500 |
+
device=hidden.device,
|
| 501 |
+
dtype=torch.long,
|
| 502 |
+
)
|
| 503 |
+
(
|
| 504 |
+
_,
|
| 505 |
+
_,
|
| 506 |
+
proposal,
|
| 507 |
+
confidence,
|
| 508 |
+
entropy,
|
| 509 |
+
greedy,
|
| 510 |
+
greedy_confidence,
|
| 511 |
+
) = chunked_vocab_training_statistics(
|
| 512 |
+
hidden,
|
| 513 |
+
weight,
|
| 514 |
+
labels,
|
| 515 |
+
softcap=softcap,
|
| 516 |
+
temperature=temperature,
|
| 517 |
+
chunk_size=chunk_size,
|
| 518 |
+
)
|
| 519 |
+
return proposal, confidence, entropy, greedy, greedy_confidence
|
| 520 |
+
|
| 521 |
+
|
| 522 |
+
__all__ = [
|
| 523 |
+
"chunked_vocab_statistics",
|
| 524 |
+
"chunked_vocab_training_statistics",
|
| 525 |
+
]
|