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

Publish Modilify Mk2 Preview

Browse files
README.md CHANGED
@@ -14,49 +14,91 @@ tags:
14
 
15
  ![LOGO](assets/01-LOGO.jpg)
16
 
17
- # Modilify Mk2 Preview
18
 
19
- A 26B-A4B multimodal block-diffusion model with dual-timescale latent deliberation.
20
 
21
- Mk2 does not dump chain-of-thought into extra visible tokens. Each heavy denoise runs a latent Transformer over a packed trajectory history and a denoise-time tape, then writes persistent memory only when the canvas actually commits. Working state is recomputed every step. Persistent slots survive the rolling window. The exclusive excess-entropy commit formula still decides how many tokens lock in; temperature and failure budget are first-class inference knobs.
22
 
23
- This repository is the first public Mk2 preview checkpoint: merged BF16 weights, remote code, processor, and tokenizer. The text trunk and dual-timescale latent stack come from schema23 training step 900 (~12.4 million adaptation tokens). The Gemma 4 vision tower is restored from DiffusionGemma so text, image, and video share one decoder. This is not a full benchmark release.
24
 
25
- ## Architecture
26
 
27
- The heavy trunk is DiffusionGemma 26B-A4B. Inside every denoise, a 4-layer latent Transformer reads the noisy 256-token canvas plus:
28
 
29
- 1. **Packed history** (T=16). Four views of each canvas position's recent deliberation, projected into rank-1024 space. New canvas positions start empty; they do not inherit the previous token's thought.
30
- 2. **Denoise tape**. Row-level probes written on a time ring. They do not shift when tokens commit.
31
- 3. **Persistent slots** (256 × 2816). Updated only at commit by a Transformer writer. Full-attention decoder layers read working and persistent buses.
 
 
 
 
 
32
 
33
- Visible tokens are the product of that loop, not the workspace. Easy prompts commit a long prefix. Hard prompts keep pondering.
34
 
35
- The encoder is the official Gemma 4 multimodal encoder. Image and video tokens condition prefix KV the same way text does; the rolling canvas and latent stack stay text-side.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
36
 
37
  ## Model Summary
38
 
39
  | | |
40
  | --- | ---: |
41
  | Architecture | Mixture-of-Experts block diffusion + dual-timescale latent Transformer |
42
- | Total Parameters | 26.139B text trunk + 569.550M vision encoder + latent stack |
43
- | Text Heavy-Denoise Activated Parameters | 4.159B |
44
- | Vision Encoder | Gemma 4 Vision, 569.550M |
45
  | Layers | 30 |
46
  | Number of Experts | 128 |
47
  | Selected Experts per Token | 8 |
48
  | Vocabulary Size | 262,144 |
49
- | Context Length | 262,144 tokens |
50
  | Sliding Window | 1024 |
51
  | Canvas Length | 256 |
52
  | Latent Width | 2,816 |
53
- | Latent Memory | 256 slots × 2,816-d, 4 layers |
 
 
54
  | Trajectory History | 16 frames, 4 views, rank 1,024 |
55
- | Denoise Tape | 16 probes |
56
  | Modality | Text, Image, Video |
57
- | Preview checkpoint | schema23 step 900 |
58
  | Adaptation tokens | ~12.4 million |
59
 
 
 
 
 
 
 
60
  ## Getting Started
61
 
62
  Transformers 5.14.1 is the minimum supported version.
@@ -157,12 +199,12 @@ output = model.generate(**inputs, max_new_tokens=256)
157
 
158
  ## Thinking mode
159
 
160
- The chat template controls the prompt, not the model's first generated tokens.
161
 
162
  - `enable_thinking=True` inserts a system turn that contains `<|think|>` and still ends the prompt at `<|turn>model`.
163
  - `enable_thinking=False` does **not** inject an empty thought channel. The prompt ends at `<|turn>model`.
164
 
165
- The model may still open `<|channel>thought` on its own. That is generation, not a template artifact. Applications should not assume hidden reasoning is complete, correct, or appropriate to expose to end users.
166
 
167
  ## Configurable inference
168
 
@@ -208,13 +250,13 @@ model = AutoModelForMultimodalLM.from_pretrained(
208
  )
209
  ```
210
 
211
- Lower temperature and a tighter budget make the model more cautious and usually slower. Higher temperature and a looser budget commit more tokens per denoise. These are compute-control decisions, not guarantees of correctness.
212
 
213
  Generation supports left-padded batches with independent stopping. Batch prompts of similar lengths together for the best throughput. Streaming and caller-supplied KV caches remain limited to batch size 1.
214
 
215
  ## Evaluation status, limitations, and risks
216
 
217
- This is a preview. It does not include a complete accuracy, robustness, calibration, fairness, or safety evaluation. Structural export checks and a small graduate-level qualitative probe, if present, do not establish fitness for use.
218
 
219
  The model can hallucinate facts, citations, visual details, or temporal relationships; reproduce bias, unsafe content, personal information, or copyrighted material; and consume substantial time and memory during long iterative generation. Confidence-based commits control compute. They do not certify that a prefix is true. Visual performance can degrade with poor resolution, motion, occlusion, unusual aspect ratios, or domain shift.
220
 
@@ -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
  ![LOGO](assets/01-LOGO.jpg)
16
 
17
+ # Modilify Mk2 Preview · 26B-A5B
18
 
19
+ **Refine a whole canvas. Think in latent space. Carry memory forward.**
20
 
21
+ Modilify Mk2 is a **26B-A5B multimodal block-diffusion model** that brings parallel token refinement, latent deliberation, and persistent trajectory memory into one generation loop. A rolling 256-token canvas gives the model room to revise upcoming text together. A dedicated latent Transformer turns the history of those revisions into context for the next step. As output advances, a learned memory writer carries information across canvas boundaries.
22
 
23
+ The central idea is simple: **give each answer a workspace, a memory, and a variable compute budget.** Mk2 can refine its internal state over multiple passes before committing text, and release multiple tokens together when its confidence-and-entropy policy permits. Deliberation happens in continuous hidden states, without requiring every internal update to become a visible reasoning token.
24
 
25
+ Built on DiffusionGemma, Mk2 combines sparse expert routing with a dual-timescale latent architecture. Text, images, and sampled video frames feed the same generation path.
26
 
27
+ ## What makes Mk2 different
28
 
29
+ | Architecture choice | What it enables |
30
+ | --- | --- |
31
+ | **Parallel block diffusion** | Revise a 256-token canvas jointly and commit a variable-length prefix, allowing multiple output tokens per denoising pass. |
32
+ | **Trajectory-aware latent deliberation** | Condition the next revision on how hidden states have evolved, including their changes, acceleration, and residuals relative to the latest state. |
33
+ | **Memory at two timescales** | Rebuild working state every pass while retaining persistent slots across rolling-window shifts within a generation. |
34
+ | **Memory inside the decoder** | Feed working and persistent memory directly into full-attention decoder layers so both can influence token refinement. |
35
+ | **Adaptive commitment** | Use confidence and entropy to decide how much text to release, with explicit budgets for continued refinement and forced progress. |
36
+ | **Sparse multimodal foundation** | Select 8 of 128 experts per token and bring text, image, and video context into a shared decoder. |
37
 
38
+ ## Inside the generation loop
39
 
40
+ ### A canvas built for revision
41
+
42
+ Mk2 maintains a rolling canvas of 256 candidate tokens. Each denoising pass updates the candidates using the encoded prompt, the current canvas, and latent context. Attention lets positions within the canvas inform one another before the output prefix is finalized. After a commit, the canvas shifts forward and opens space for new candidates.
43
+
44
+ This creates two useful degrees of freedom: **how many times to refine** and **how many tokens to release**. A pass can commit a longer prefix when the policy permits, or spend additional computation refining an uncertain frontier.
45
+
46
+ ### Deliberation that reads its own trajectory
47
+
48
+ A 4-layer latent Transformer, 2,816 dimensions wide, builds working context before each decoder denoising pass. It reads the current canvas alongside three complementary sources of state:
49
+
50
+ - **Per-token trajectory history:** 16 recent frames represented through four views—hidden state, first difference, second difference, and residual from the latest state—projected to rank 1,024. These give the processor access to the direction and stability of recent revisions. History follows surviving canvas tokens; new positions start empty.
51
+ - **Denoising tape:** 16 pooled probes per frame summarize canvas activity in a time-indexed ring. The tape preserves recent step-level context as token positions move through the window.
52
+ - **Persistent memory:** 256 slots of 2,816 dimensions carry learned summaries across commits within the generation.
53
+
54
+ The resulting working context enters the decoder through its self-conditioning bridge and working-memory bus. Full-attention layers also read a separate persistent-memory bus. **The refinement history becomes an input to the next refinement.**
55
+
56
+ ### Fast working state, lasting commit memory
57
+
58
+ Working state is recomputed on every denoising pass. Persistent slots update only when tokens commit, using a Transformer writer with a separate gate for each slot. That writer draws on the committed region's trajectory, working state, and final decoder representation.
59
+
60
+ This separates rapid revision from memory consolidation: the canvas can keep changing while persistent memory stays stable between commits. When the window advances, the memory slots remain available to later tokens.
61
+
62
+ ### Compute that follows the commit frontier
63
+
64
+ Mk2 combines proposal confidence with an excess-entropy penalty and selects the longest prefix whose cumulative failure score stays below the configured budget. A tighter budget requires stronger evidence before normal commitment; a looser budget admits longer prefixes for the same scores. Stagnation handling and a pondering watchdog bound continued refinement.
65
+
66
+ Temperature and commit budget are exposed directly at inference time, making the generation policy adjustable per request. These scores govern commitment; they are not calibrated guarantees of factual correctness.
67
+
68
+ ### One generation path for text, images, and video
69
+
70
+ The Gemma 4 vision tower supplies visual features through the DiffusionGemma multimodal encoder. Text, image, and sampled video-frame inputs are encoded into the prefix KV cache that conditions the decoder. The rolling text canvas then uses the same latent deliberation and memory loop across all three input modalities.
71
 
72
  ## Model Summary
73
 
74
  | | |
75
  | --- | ---: |
76
  | Architecture | Mixture-of-Experts block diffusion + dual-timescale latent Transformer |
77
+ | Model Size | 26B-A5B |
78
+ | Vision Encoder | Gemma 4 Vision |
 
79
  | Layers | 30 |
80
  | Number of Experts | 128 |
81
  | Selected Experts per Token | 8 |
82
  | Vocabulary Size | 262,144 |
83
+ | Configured Context Length | 262,144 tokens |
84
  | Sliding Window | 1024 |
85
  | Canvas Length | 256 |
86
  | Latent Width | 2,816 |
87
+ | Latent Transformer | 4 layers, 16 attention heads |
88
+ | Persistent Memory | 256 slots × 2,816 dimensions |
89
+ | Memory Writer | 2-layer commit-sequence Transformer + per-slot gated writer |
90
  | Trajectory History | 16 frames, 4 views, rank 1,024 |
91
+ | Denoise Tape | 16 probes per frame |
92
  | Modality | Text, Image, Video |
93
+ | Preview checkpoint | Training step 900 |
94
  | Adaptation tokens | ~12.4 million |
95
 
96
+ ## Preview release
97
+
98
+ This first public preview includes **merged BF16 weights, inference code, processor, and tokenizer**. The text backbone and latent stack are exported from training step 900, after approximately 12.4 million adaptation tokens. The Gemma 4 vision tower is restored from DiffusionGemma.
99
+
100
+ The release makes the architecture available for hands-on exploration and evaluation. Comprehensive benchmark results are not included; measured speed, reasoning quality, and multimodal reliability remain to be established for specific workloads.
101
+
102
  ## Getting Started
103
 
104
  Transformers 5.14.1 is the minimum supported version.
 
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

  • SHA256: 4d2316525989c32b9028374439f97af3b90be87f0a7ff0666a6c35d283f83358
  • Pointer size: 131 Bytes
  • Size of remote file: 431 kB

Git LFS Details

  • SHA256: dbbe5819f47d5a367c9bb6bb448d82afe5bdb881ef8676f4d4c22b2a6b20e261
  • Pointer size: 131 Bytes
  • Size of remote file: 335 kB
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
+ ]