Xunzhuo commited on
Commit
2b02c42
·
verified ·
1 Parent(s): dc6c41c

d3 v3.0.4

Browse files
Files changed (5) hide show
  1. MODEL_MANIFEST.json +11 -5
  2. README.md +2 -2
  3. d3_fast.py +512 -0
  4. d3_kernels.py +523 -0
  5. d3_runtime.py +239 -35
MODEL_MANIFEST.json CHANGED
@@ -564,7 +564,9 @@
564
  },
565
  "runtime": {
566
  "files_sha256": {
567
- "d3_runtime.py": "a4b577dedcc946ad749589ad47e8314a3db7490245c89f896109a32ec4a54f53",
 
 
568
  "modeling_d3.py": "55314f36206484a505e9525db373621d672df7447ddc07b71329a1f2857e8227",
569
  "pipeline_d3.py": "c0975a91316bcda9c09c5f14bf634ceec94e6980ed7e76a29b718fb451fe4091",
570
  "d3_server.py": "434d6aa5e8572b649e1068588bf1811ff970412e409684b80832bd5544080bca",
@@ -586,11 +588,11 @@
586
  }
587
  ]
588
  },
589
- "built_utc": "2026-10-10T20:41:39+00:00",
590
  "files_sha256": {
591
  "LICENSE": "c71d239df91726fc519c6eb72d318ec65820627232b2f796219e87dcf35d0ab4",
592
  "NOTICE": "61f4ea488e1fed0c3fce97171506a23152c2591116442f78f15827fb9e6af469",
593
- "README.md": "089caef8180dfbc3a7725e1e37331f4d71509557d87e5c36cec050b8f56dce1b",
594
  "assets/banner.png": "eb7ecc949232e1de544247899dd55ee629a36cfe276042d5795f65178b9c64c5",
595
  "assets/example-receipt.png": "b9dbb8b103bacd45f1de824d844a3c07b81d57cb1a7466712ef060ad7dcb3926",
596
  "assets/index-areas.png": "49ecc07e5c271a69a42572240cb48afb648b8f9b99cd6bf934c9bf7ff79e4a77",
@@ -598,8 +600,10 @@
598
  "chat_template.jinja": "c3cf9e34abf4f9e36c2d72165aa9c132d3e2a725b6c2586aaa3a8af9d7a81041",
599
  "config.json": "7e16284fafd2d54b73c10073c0391bfd6804886ceaf2cbe37fdb62d4c6813dea",
600
  "d3_engine.py": "ab638d785325d944d3dac891f77448dfc91cad1b8d55a7302bd527c5f1e0e251",
 
601
  "d3_format.py": "e5036154d2e54793f59b320c8726632957bc343c382bac26d825b263a6e9ce62",
602
- "d3_runtime.py": "a4b577dedcc946ad749589ad47e8314a3db7490245c89f896109a32ec4a54f53",
 
603
  "d3_server.py": "434d6aa5e8572b649e1068588bf1811ff970412e409684b80832bd5544080bca",
604
  "decision_config.json": "6b80ca11bd6ba481df4b3c187db8a5d1983786563ca05e2edccb0bdac0d2fc4f",
605
  "merges.txt": "a9d356d7bdf1ef4949e3e748e95b8e10ad9d4e2e838eddc38a0a7b6b94d1db8d",
@@ -636,8 +640,10 @@
636
  "chat_template.jinja": 8952,
637
  "config.json": 3949,
638
  "d3_engine.py": 3726,
 
639
  "d3_format.py": 4882,
640
- "d3_runtime.py": 49044,
 
641
  "d3_server.py": 8295,
642
  "decision_config.json": 5503,
643
  "merges.txt": 3353259,
 
564
  },
565
  "runtime": {
566
  "files_sha256": {
567
+ "d3_runtime.py": "88fcaf3451edcc2139ffcb0db53ea4f9d03152e2736e0ffbb47386b40f618d52",
568
+ "d3_fast.py": "c8f6e74968e2835732e0865814b98a21e4b604c80ecb0416c270219eb8e64c05",
569
+ "d3_kernels.py": "8f7534ca1e1846ebf7c4c73058002f3be4dbf0d5d426f0815150ae0cbe7677bd",
570
  "modeling_d3.py": "55314f36206484a505e9525db373621d672df7447ddc07b71329a1f2857e8227",
571
  "pipeline_d3.py": "c0975a91316bcda9c09c5f14bf634ceec94e6980ed7e76a29b718fb451fe4091",
572
  "d3_server.py": "434d6aa5e8572b649e1068588bf1811ff970412e409684b80832bd5544080bca",
 
588
  }
589
  ]
590
  },
591
+ "built_utc": "2026-10-10T21:44:47+00:00",
592
  "files_sha256": {
593
  "LICENSE": "c71d239df91726fc519c6eb72d318ec65820627232b2f796219e87dcf35d0ab4",
594
  "NOTICE": "61f4ea488e1fed0c3fce97171506a23152c2591116442f78f15827fb9e6af469",
595
+ "README.md": "8cde7822a0840c7c641df439f0f31d64e94216f2bd5df15d6b55eb0ba6aec264",
596
  "assets/banner.png": "eb7ecc949232e1de544247899dd55ee629a36cfe276042d5795f65178b9c64c5",
597
  "assets/example-receipt.png": "b9dbb8b103bacd45f1de824d844a3c07b81d57cb1a7466712ef060ad7dcb3926",
598
  "assets/index-areas.png": "49ecc07e5c271a69a42572240cb48afb648b8f9b99cd6bf934c9bf7ff79e4a77",
 
600
  "chat_template.jinja": "c3cf9e34abf4f9e36c2d72165aa9c132d3e2a725b6c2586aaa3a8af9d7a81041",
601
  "config.json": "7e16284fafd2d54b73c10073c0391bfd6804886ceaf2cbe37fdb62d4c6813dea",
602
  "d3_engine.py": "ab638d785325d944d3dac891f77448dfc91cad1b8d55a7302bd527c5f1e0e251",
603
+ "d3_fast.py": "c8f6e74968e2835732e0865814b98a21e4b604c80ecb0416c270219eb8e64c05",
604
  "d3_format.py": "e5036154d2e54793f59b320c8726632957bc343c382bac26d825b263a6e9ce62",
605
+ "d3_kernels.py": "8f7534ca1e1846ebf7c4c73058002f3be4dbf0d5d426f0815150ae0cbe7677bd",
606
+ "d3_runtime.py": "88fcaf3451edcc2139ffcb0db53ea4f9d03152e2736e0ffbb47386b40f618d52",
607
  "d3_server.py": "434d6aa5e8572b649e1068588bf1811ff970412e409684b80832bd5544080bca",
608
  "decision_config.json": "6b80ca11bd6ba481df4b3c187db8a5d1983786563ca05e2edccb0bdac0d2fc4f",
609
  "merges.txt": "a9d356d7bdf1ef4949e3e748e95b8e10ad9d4e2e838eddc38a0a7b6b94d1db8d",
 
640
  "chat_template.jinja": 8952,
641
  "config.json": 3949,
642
  "d3_engine.py": 3726,
643
+ "d3_fast.py": 21284,
644
  "d3_format.py": 4882,
645
+ "d3_kernels.py": 16129,
646
+ "d3_runtime.py": 57795,
647
  "d3_server.py": 8295,
648
  "decision_config.json": 5503,
649
  "merges.txt": 3353259,
README.md CHANGED
@@ -28,10 +28,10 @@ tags:
28
 
29
  ## Highlights
30
 
31
- - **Jev Decision Index 0.3, public suite: 65.15**, measured with the official 0.3 kit on the released weights: all 140,178 public requests answered, none unsupported.
32
  - **+8.2 on the public suite over Decision 2.0** (its 27B model: 56.97 on the board), ahead in all five areas.
33
  - **Reads images:** multiple images per request (PNG, JPEG or WebP), given as paths, URLs, PIL images or base64 data URLs; every question of the request sees all of them.
34
- - **Speed:** text requests take a median of 84 ms, and requests with an image a median of 342 ms, on one AMD Instinct MI325X, one request at a time.
35
  - **Many questions, one call:** Choice, Yes / No and Score questions about the same input are answered together, each from its own forward pass over the input, with a probability for every option.
36
 
37
  ## Quickstart
 
28
 
29
  ## Highlights
30
 
31
+ - **Jev Decision Index 0.3, public suite: 65.17**, measured with the official 0.3 kit on the released weights: all 140,178 public requests answered, none unsupported.
32
  - **+8.2 on the public suite over Decision 2.0** (its 27B model: 56.97 on the board), ahead in all five areas.
33
  - **Reads images:** multiple images per request (PNG, JPEG or WebP), given as paths, URLs, PIL images or base64 data URLs; every question of the request sees all of them.
34
+ - **Speed:** text requests take a median of 54 ms, and requests with an image a median of 277 ms, on one AMD Instinct MI325X, one request at a time.
35
  - **Many questions, one call:** Choice, Yes / No and Score questions about the same input are answered together, each from its own forward pass over the input, with a probability for every option.
36
 
37
  ## Quickstart
d3_fast.py ADDED
@@ -0,0 +1,512 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """ROCm fast path of the d3 runtime: same function, same weights, less host work.
2
+
3
+ ``d3_runtime.D3`` installs it on AMD GPUs (``torch.version.hip``) for Qwen3.5 backbones with noncausal full
4
+ attention; every other device and model keeps the plain path, and ``D3_FAST=0`` turns it off.
5
+
6
+ - Text forward passes replay HIP graphs captured at warm-up. A pass (one batch of questions, left-padded to
7
+ its longest prompt exactly as in the plain path) is extended on the right with masked padding to a length
8
+ bucket and replays the graph of its (questions, bucket) shape. The real tokens keep their positions,
9
+ chunk boundaries and convolution taps; the right padding is masked out of the full-attention keys and
10
+ comes after every real token in the causal Gated DeltaNet layers; the readout reads each prompt's last
11
+ real token. Passes over the graph budget (``D3_GRAPH_TOKENS`` rows x bucket tokens) run eagerly at their
12
+ own shape.
13
+ - Eager passes build both attention masks from the host-known prompt lengths, so nothing inside the forward
14
+ pass waits for the GPU.
15
+ """
16
+
17
+ from __future__ import annotations
18
+
19
+ import os
20
+ import time
21
+ from typing import Any
22
+
23
+ GRAPH_TOKENS = 0
24
+ BUCKET = 64
25
+ CAPTURE_WARM_RUNS = 2
26
+
27
+
28
+ def enabled(model: Any) -> str | None:
29
+ """None when the fast path applies to this loaded model, else the reason it does not."""
30
+ torch = model.torch
31
+ if os.environ.get("D3_FAST", "").strip().lower() in ("0", "false", "no", "off"):
32
+ return "D3_FAST=0"
33
+ if model.device.type != "cuda" or not getattr(torch.version, "hip", None):
34
+ return "not a ROCm GPU"
35
+ if model.attention_mode != "noncausal_full_attention":
36
+ return f"attention mode {model.attention_mode}"
37
+ if type(model.backbone).__name__ != "Qwen3_5Model":
38
+ return f"backbone {type(model.backbone).__name__}"
39
+ config = model.backbone.language_model.config
40
+ if set(config.layer_types[: config.num_hidden_layers]) - {
41
+ "full_attention",
42
+ "linear_attention",
43
+ }:
44
+ return "layer types other than full / linear attention"
45
+ return None
46
+
47
+
48
+ def fusable(model: Any) -> str | None:
49
+ """None when the fused decoder-layer kernels (``d3_kernels.py``) apply to this model, else the reason."""
50
+ torch = model.torch
51
+ arch = torch.cuda.get_device_properties(model.device).gcnArchName.split(":")[0]
52
+ if arch != "gfx942":
53
+ return f"fused kernels verified on gfx942 only, not {arch}"
54
+ try:
55
+ import triton # noqa: F401
56
+ except ImportError:
57
+ return "triton is not installed"
58
+ config = model.backbone.language_model.config
59
+ chunk = model.kernels.get("torch_chunk_gated_delta_rule", "")
60
+ checks = {
61
+ "hidden size a multiple of 256": config.hidden_size % 256 == 0,
62
+ "128-wide gated-delta heads": config.linear_key_head_dim == 128
63
+ and config.linear_value_head_dim == 128,
64
+ "256-wide attention heads": getattr(config, "head_dim", None) == 256,
65
+ "4-tap convolution": config.linear_conv_kernel_dim == 4,
66
+ "SiLU MLP": config.hidden_act == "silu",
67
+ "FLA chunk kernel": chunk.startswith("fla"),
68
+ "SDPA attention": config._attn_implementation == "sdpa",
69
+ }
70
+ missing = [name for name, ok in checks.items() if not ok]
71
+ return "needs " + ", ".join(missing) if missing else None
72
+
73
+
74
+ class FusedLayers:
75
+ """The decoder layers and final norm of a Qwen3.5 text model through ``d3_kernels`` (eager values)."""
76
+
77
+ def __init__(self, lm: Any, torch: Any):
78
+ import importlib
79
+
80
+ from transformers.models.qwen3_5 import modeling_qwen3_5 as modeling
81
+
82
+ try:
83
+ from . import d3_kernels as kernels
84
+ except ImportError:
85
+ kernels = importlib.import_module("d3_kernels")
86
+ self.k = kernels
87
+ self.torch = torch
88
+ self.chunk = modeling.torch_chunk_gated_delta_rule
89
+ self.repeat_kv = modeling.repeat_kv
90
+ config = lm.config
91
+ self.layers = list(lm.layers[: config.num_hidden_layers])
92
+ self.params = []
93
+ with torch.no_grad():
94
+ for layer in self.layers:
95
+ p = {
96
+ "linear": layer.block_type == "linear_attention",
97
+ "eps": layer.input_layernorm.eps,
98
+ "w1_in": (1.0 + layer.input_layernorm.weight.float()).contiguous(),
99
+ "w1_post": (
100
+ 1.0 + layer.post_attention_layernorm.weight.float()
101
+ ).contiguous(),
102
+ }
103
+ if p["linear"]:
104
+ m = layer.linear_attn
105
+ p.update(
106
+ conv_w=m.conv1d.weight.squeeze(1).contiguous(),
107
+ A_log=m.A_log.float().contiguous(),
108
+ dt_bias=m.dt_bias.float().contiguous(),
109
+ )
110
+ else:
111
+ m = layer.self_attn
112
+ p.update(
113
+ qw1=(1.0 + m.q_norm.weight.float()).contiguous(),
114
+ kw1=(1.0 + m.k_norm.weight.float()).contiguous(),
115
+ )
116
+ self.params.append(p)
117
+ self.final_eps = lm.norm.eps
118
+ self.final_w1 = (1.0 + lm.norm.weight.float()).contiguous()
119
+
120
+ def gated_delta(self, m: Any, p: dict[str, Any], x: Any) -> Any:
121
+ k = self.k
122
+ rows, length, _ = x.shape
123
+ q, key, v, g, beta = k.gdn_prep(
124
+ m.in_proj_qkv(x),
125
+ m.in_proj_b(x),
126
+ m.in_proj_a(x),
127
+ p["conv_w"],
128
+ p["A_log"],
129
+ p["dt_bias"],
130
+ m.num_k_heads,
131
+ m.head_k_dim,
132
+ )
133
+ core, _ = self.chunk(
134
+ q,
135
+ key,
136
+ v,
137
+ g=g,
138
+ beta=beta,
139
+ initial_state=None,
140
+ output_final_state=False,
141
+ use_qk_l2norm_in_kernel=True,
142
+ cu_seqlens=None,
143
+ use_cache=False,
144
+ )
145
+ out = k.gated_rmsnorm(
146
+ core.reshape(-1, m.head_v_dim),
147
+ m.in_proj_z(x).reshape(-1, m.head_v_dim),
148
+ m.norm.weight,
149
+ m.norm.variance_epsilon,
150
+ )
151
+ return m.out_proj(out.reshape(rows, length, -1))
152
+
153
+ def attention(self, m: Any, p: dict[str, Any], x: Any, rotary, full_mask) -> Any:
154
+ rows, length, _ = x.shape
155
+ hd = m.head_dim
156
+ qp = m.q_proj(x)
157
+ kp = m.k_proj(x)
158
+ heads = qp.shape[-1] // (2 * hd)
159
+ kv_heads = kp.shape[-1] // hd
160
+ cos, sin = rotary
161
+ q, key = self.k.attn_prep(
162
+ qp, kp, p["qw1"], p["kw1"], cos, sin, heads, kv_heads, hd, m.q_norm.eps
163
+ )
164
+ v = self.repeat_kv(
165
+ m.v_proj(x).view(rows, length, kv_heads, hd).transpose(1, 2),
166
+ heads // kv_heads,
167
+ )
168
+ attn = self.torch.nn.functional.scaled_dot_product_attention(
169
+ q,
170
+ key,
171
+ v,
172
+ attn_mask=full_mask,
173
+ dropout_p=0.0,
174
+ scale=m.scaling,
175
+ is_causal=False,
176
+ )
177
+ gate = qp.view(rows, length, heads, 2 * hd)[..., hd:]
178
+ return m.o_proj(self.k.sigmoid_gate(attn.transpose(1, 2), gate))
179
+
180
+ def forward(self, embeds: Any, full_mask: Any, rowmask: Any | None, rotary) -> Any:
181
+ """Final-norm hidden states [B, L, D]; ``rowmask`` marks real tokens when the batch is padded."""
182
+ k = self.k
183
+ hidden, delta = embeds, None
184
+ for layer, p in zip(self.layers, self.params):
185
+ hidden, x = k.add_rmsnorm(
186
+ hidden, delta, p["w1_in"], p["eps"], rowmask if p["linear"] else None
187
+ )
188
+ if p["linear"]:
189
+ delta = self.gated_delta(layer.linear_attn, p, x)
190
+ else:
191
+ delta = self.attention(layer.self_attn, p, x, rotary, full_mask)
192
+ hidden, x = k.add_rmsnorm(hidden, delta, p["w1_post"], p["eps"])
193
+ mlp = layer.mlp
194
+ delta = mlp.down_proj(k.silu_mul(mlp.gate_proj(x), mlp.up_proj(x)))
195
+ return k.add_rmsnorm(hidden, delta, self.final_w1, self.final_eps)[1]
196
+
197
+
198
+ class FastText:
199
+ """Text passes of one loaded model: HIP graphs per (rows, bucket) plus a synchronization-free eager pass."""
200
+
201
+ def __init__(
202
+ self,
203
+ model: Any,
204
+ *,
205
+ graph_tokens: int,
206
+ bucket: int,
207
+ max_rows: int,
208
+ fused: FusedLayers | None = None,
209
+ fused_skipped: str | None = None,
210
+ ):
211
+ torch = model.torch
212
+ self.model = model
213
+ self.torch = torch
214
+ self.fused = fused
215
+ self.fused_skipped = fused_skipped
216
+ self.cpu_threads: dict[str, int] | None = None
217
+ self.lm = model.backbone.language_model
218
+ config = self.lm.config
219
+ self.layers = list(self.lm.layers[: config.num_hidden_layers])
220
+ self.kinds = list(config.layer_types[: config.num_hidden_layers])
221
+ self.graph_tokens = graph_tokens
222
+ self.bucket = bucket
223
+ self.max_rows = max_rows
224
+ self.codes = torch.arange(model.readout.shape[0], device=model.device)
225
+ self.graphs: dict[tuple[int, int], dict[str, Any]] = {}
226
+ self.failed: dict[tuple[int, int], str] = {}
227
+ self.pool = None
228
+ self.stats = {"replays": 0, "eager": 0, "captures": 0, "capture_seconds": 0.0}
229
+
230
+ # ------------------------------------------------------------------ forward pieces
231
+
232
+ def hidden(self, ids, mask, linear_mask):
233
+ """Final-norm hidden states [B, L, D]: ``Qwen3_5Model.forward`` for text with the d3 masks given."""
234
+ torch = self.torch
235
+ lm = self.lm
236
+ embeds = lm.embed_tokens(ids)
237
+ rows, length = ids.shape
238
+ positions = torch.arange(length, device=ids.device).view(1, 1, -1)
239
+ positions = positions.expand(4, rows, -1)
240
+ text_positions, positions = positions[0], positions[1:]
241
+ rotary = lm.rotary_emb(embeds, positions)
242
+ full = mask.bool()[:, None, None, :]
243
+ if self.fused is not None:
244
+ rowmask = None if linear_mask is None else linear_mask.reshape(-1)
245
+ return self.fused.forward(embeds, full, rowmask, rotary)
246
+ hidden = embeds
247
+ for layer, kind in zip(self.layers, self.kinds):
248
+ hidden = layer(
249
+ hidden,
250
+ position_embeddings=rotary,
251
+ attention_mask=full if kind == "full_attention" else linear_mask,
252
+ position_ids=text_positions,
253
+ past_key_values=None,
254
+ use_cache=False,
255
+ )
256
+ return lm.norm(hidden)
257
+
258
+ def image_hidden(self, inputs: dict[str, Any], padded: bool):
259
+ """Final-norm hidden states of an image batch: ``Qwen3_5Model.forward`` with the fused text layers."""
260
+ torch = self.torch
261
+ backbone = self.model.backbone
262
+ ids = inputs["input_ids"]
263
+ mask = inputs["attention_mask"]
264
+ embeds = backbone.get_input_embeddings()(ids)
265
+ features = backbone.get_image_features(
266
+ inputs["pixel_values"], inputs["image_grid_thw"], return_dict=True
267
+ ).pooler_output
268
+ features = torch.cat(features, dim=0).to(embeds.device, embeds.dtype)
269
+ image_mask, _ = backbone.get_placeholder_mask(
270
+ ids, inputs_embeds=embeds, image_features=features
271
+ )
272
+ embeds = embeds.masked_scatter(image_mask, features)
273
+ positions = backbone.compute_3d_position_ids(
274
+ input_ids=ids,
275
+ image_grid_thw=inputs["image_grid_thw"],
276
+ inputs_embeds=embeds,
277
+ attention_mask=mask,
278
+ past_key_values=None,
279
+ mm_token_type_ids=inputs["mm_token_type_ids"],
280
+ )
281
+ rotary = self.lm.rotary_emb(embeds, positions)
282
+ full = mask.bool()[:, None, None, :]
283
+ return self.fused.forward(
284
+ embeds, full, mask.reshape(-1) if padded else None, rotary
285
+ )
286
+
287
+ def probabilities_of(self, last, counts):
288
+ """``D3.logits`` -> temperature -> softmax on the last-token hidden states [B, D]."""
289
+ model = self.model
290
+ if model.readout_dtype == "float32":
291
+ logits = last.float() @ model.readout.T
292
+ else:
293
+ logits = self.torch.nn.functional.linear(last, model.readout).float()
294
+ invalid = self.codes[None] >= counts[:, None]
295
+ return (logits.masked_fill(invalid, float("-inf")) / model.temperature).softmax(
296
+ -1
297
+ )
298
+
299
+ # ------------------------------------------------------------------ passes
300
+
301
+ def bucket_of(self, rows: int, width: int) -> int | None:
302
+ size = -(-width // self.bucket) * self.bucket
303
+ if rows > self.max_rows or rows * size > self.graph_tokens:
304
+ return None
305
+ return size
306
+
307
+ def probabilities(self, sequences, counts) -> list[list[float]]:
308
+ torch = self.torch
309
+ rows = len(sequences)
310
+ width = max(len(s) for s in sequences)
311
+ size = self.bucket_of(rows, width)
312
+ entry = self.graphs.get((rows, size)) if size is not None else None
313
+ length = size if entry is not None else width
314
+ ids = torch.full((rows, length), self.model.pad_id, dtype=torch.long)
315
+ mask = torch.zeros((rows, length), dtype=torch.long)
316
+ for i, sequence in enumerate(sequences):
317
+ ids[i, width - len(sequence) : width] = torch.as_tensor(
318
+ sequence, dtype=torch.long
319
+ )
320
+ mask[i, width - len(sequence) : width] = 1
321
+ counts_cpu = torch.as_tensor(list(counts), dtype=torch.long)
322
+ if entry is not None:
323
+ entry["ids"].copy_(ids, non_blocking=True)
324
+ entry["mask"].copy_(mask, non_blocking=True)
325
+ entry["counts"].copy_(counts_cpu, non_blocking=True)
326
+ entry["last"].fill_(width - 1)
327
+ entry["graph"].replay()
328
+ self.stats["replays"] += 1
329
+ probs = entry["probs"].cpu().tolist()
330
+ else:
331
+ device = self.model.device
332
+ ids = ids.to(device, non_blocking=True)
333
+ mask = mask.to(device, non_blocking=True)
334
+ padded = any(len(s) != width for s in sequences)
335
+ with torch.inference_mode():
336
+ last = self.hidden(ids, mask, mask if padded else None)[:, -1]
337
+ probs = (
338
+ self.probabilities_of(last, counts_cpu.to(device)).cpu().tolist()
339
+ )
340
+ self.stats["eager"] += 1
341
+ return [p[:c] for p, c in zip(probs, counts)]
342
+
343
+ # ------------------------------------------------------------------ graphs
344
+
345
+ def shapes(self) -> list[tuple[int, int]]:
346
+ out = []
347
+ for rows in range(1, self.max_rows + 1):
348
+ size = self.bucket
349
+ while rows * size <= self.graph_tokens:
350
+ out.append((rows, size))
351
+ size += self.bucket
352
+ return out
353
+
354
+ def capture(self, rows: int, size: int) -> bool:
355
+ torch = self.torch
356
+ key = (rows, size)
357
+ if key in self.graphs:
358
+ return True
359
+ device = self.model.device
360
+ token = self.model.token_ids[0]
361
+ static = {
362
+ "ids": torch.full((rows, size), token, dtype=torch.long, device=device),
363
+ "mask": torch.ones((rows, size), dtype=torch.long, device=device),
364
+ "counts": torch.full((rows,), 2, dtype=torch.long, device=device),
365
+ "last": torch.full((1,), size - 1, dtype=torch.long, device=device),
366
+ }
367
+
368
+ def body():
369
+ hidden = self.hidden(static["ids"], static["mask"], static["mask"])
370
+ last = hidden.index_select(1, static["last"]).squeeze(1)
371
+ return self.probabilities_of(last, static["counts"])
372
+
373
+ started = time.perf_counter()
374
+ try:
375
+ with torch.inference_mode():
376
+ if self.pool is None:
377
+ self.pool = torch.cuda.graph_pool_handle()
378
+ stream = torch.cuda.Stream(device=device)
379
+ stream.wait_stream(torch.cuda.current_stream(device))
380
+ with torch.cuda.stream(stream):
381
+ for _ in range(CAPTURE_WARM_RUNS):
382
+ body()
383
+ torch.cuda.current_stream(device).wait_stream(stream)
384
+ torch.cuda.synchronize(device)
385
+ graph = torch.cuda.CUDAGraph()
386
+ with torch.cuda.graph(graph, pool=self.pool):
387
+ probs = body()
388
+ torch.cuda.synchronize(device)
389
+ except Exception as exc: # noqa: BLE001 - the shape then runs eagerly
390
+ torch.cuda.synchronize(device)
391
+ self.failed[key] = f"{type(exc).__name__}: {str(exc)[:200]}"
392
+ return False
393
+ self.graphs[key] = {**static, "graph": graph, "probs": probs}
394
+ self.stats["captures"] += 1
395
+ self.stats["capture_seconds"] += time.perf_counter() - started
396
+ return True
397
+
398
+ def capture_all(self) -> float:
399
+ started = time.perf_counter()
400
+ for rows, size in self.shapes():
401
+ self.capture(rows, size)
402
+ return time.perf_counter() - started
403
+
404
+ def report(self) -> dict[str, Any]:
405
+ return {
406
+ "fused_layers": len(self.fused.layers) if self.fused is not None else 0,
407
+ "fused_skipped": self.fused_skipped,
408
+ "cpu_threads": self.cpu_threads,
409
+ "graph_tokens": self.graph_tokens,
410
+ "bucket": self.bucket,
411
+ "max_rows": self.max_rows,
412
+ "graphs": len(self.graphs),
413
+ "failed": dict(list(self.failed.items())[:5]),
414
+ "failed_count": len(self.failed),
415
+ **{
416
+ k: (round(v, 1) if isinstance(v, float) else v)
417
+ for k, v in self.stats.items()
418
+ },
419
+ }
420
+
421
+
422
+ def dedupe_image_processing(processor: Any, torch: Any) -> None:
423
+ """Process each distinct image of one processor call once and repeat its rows for the other copies.
424
+
425
+ The runtime hands the processor one copy of the request's images per question; every copy gives the same
426
+ pixel rows and grid, so the outputs are identical, only the repeated resizing is skipped.
427
+ """
428
+ original = processor._process_images
429
+
430
+ def process_images(images, **kwargs):
431
+ flat = list(images) if isinstance(images, (list, tuple)) else [images]
432
+ index: dict[int, int] = {}
433
+ unique, order = [], []
434
+ for image in flat:
435
+ order.append(index.setdefault(id(image), len(unique)))
436
+ if len(unique) < len(index):
437
+ unique.append(image)
438
+ if len(unique) == len(flat):
439
+ return original(images, **kwargs)
440
+ processed = processor.image_processor(unique, **kwargs)
441
+ if set(processed.keys()) != {"pixel_values", "image_grid_thw"}:
442
+ return original(images, **kwargs)
443
+ grid = processed["image_grid_thw"]
444
+ rows = torch.split(processed["pixel_values"], grid.prod(-1).tolist())
445
+ processed["pixel_values"] = torch.cat([rows[i] for i in order])
446
+ processed["image_grid_thw"] = grid[order]
447
+ replacements = [
448
+ processor.replace_image_token(processed, image_idx=i, **kwargs)
449
+ for i in range(len(flat))
450
+ ]
451
+ return processed, replacements
452
+
453
+ processor._process_images = process_images
454
+
455
+
456
+ def cpu_quota() -> int | None:
457
+ """CPUs this process may use: the cgroup CPU quota (containers), else the affinity mask."""
458
+ try:
459
+ with open("/sys/fs/cgroup/cpu.max", encoding="utf-8") as stream:
460
+ quota, period = stream.read().split()[:2]
461
+ if quota != "max":
462
+ return max(1, int(quota) // int(period))
463
+ except (OSError, ValueError):
464
+ pass
465
+ try:
466
+ return len(os.sched_getaffinity(0))
467
+ except (AttributeError, OSError):
468
+ return None
469
+
470
+
471
+ def cap_cpu_threads(torch: Any) -> dict[str, int] | None:
472
+ """Keep torch's CPU threads within the container's CPU quota when OMP_NUM_THREADS is not set.
473
+
474
+ torch sizes its pool by the host's CPU count; in a container with a smaller quota the image
475
+ preprocessing then oversubscribes it and the idle-spinning workers starve the thread that feeds the GPU.
476
+ """
477
+ if "OMP_NUM_THREADS" in os.environ:
478
+ return None
479
+ limit, current = cpu_quota(), torch.get_num_threads()
480
+ if limit is None or current <= limit:
481
+ return None
482
+ torch.set_num_threads(limit)
483
+ return {"from": current, "to": limit}
484
+
485
+
486
+ def install(model: Any) -> FastText | None:
487
+ """The fast path for a loaded ``D3`` model, or None (the reason is in ``model.fast_skipped``)."""
488
+ reason = enabled(model)
489
+ if reason is not None:
490
+ model.fast_skipped = reason
491
+ return None
492
+ graph_tokens = int(os.environ.get("D3_GRAPH_TOKENS", GRAPH_TOKENS))
493
+ bucket = int(os.environ.get("D3_GRAPH_BUCKET", BUCKET))
494
+ if model.processor is not None and hasattr(model.processor, "_process_images"):
495
+ dedupe_image_processing(model.processor, model.torch)
496
+ fused, skipped = None, None
497
+ if os.environ.get("D3_FUSED", "").strip().lower() in ("0", "false", "no", "off"):
498
+ skipped = "D3_FUSED=0"
499
+ else:
500
+ skipped = fusable(model)
501
+ if skipped is None:
502
+ fused = FusedLayers(model.backbone.language_model, model.torch)
503
+ fast = FastText(
504
+ model,
505
+ graph_tokens=graph_tokens,
506
+ bucket=bucket,
507
+ max_rows=model.batch_size,
508
+ fused=fused,
509
+ fused_skipped=skipped,
510
+ )
511
+ fast.cpu_threads = cap_cpu_threads(model.torch)
512
+ return fast
d3_kernels.py ADDED
@@ -0,0 +1,523 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Fused Triton kernels for the element-wise ops of the d3 Qwen3.5 decoder layers (ROCm gfx942).
2
+
3
+ Each kernel computes what the BF16 eager forward computes, rounding where the eager ops round: the BF16
4
+ residual adds, the norms' FP32 math rounded to BF16 at the end, SiLU / sigmoid outputs in BF16, the depthwise
5
+ convolution summed in FP64 and rounded to BF16 like MIOpen's naive kernel, RoPE products and sums each rounded
6
+ to BF16. Transcendentals call the same OCML functions as the ATen kernels, floating-point contraction is off,
7
+ the RMSNorm means are summed in the order of ATen's ROCm row reduction (one 64-lane wavefront per row, four
8
+ accumulators per lane, then a lane tree; rows of 128 use 32 lanes) and ``torch.rsqrt``, which is correctly
9
+ rounded on ROCm, is reproduced through FP64. GEMMs, attention and the gated-delta chunk kernel are unchanged.
10
+
11
+ Kernels:
12
+ add_rmsnorm BF16 residual add + zero-centred RMSNorm (optionally zeroing padding rows)
13
+ silu_mul SiLU(gate) * up
14
+ gdn_prep causal conv + SiLU + q / k (repeated to the value heads) / v split + beta + g
15
+ gated_rmsnorm Gated DeltaNet output norm with the SiLU(z) gate
16
+ attn_prep q / gate split, q / k RMSNorm, partial RoPE, [B, H, T, D] layout (k repeated)
17
+ sigmoid_gate attention output * sigmoid(gate)
18
+ """
19
+
20
+ from __future__ import annotations
21
+
22
+ from typing import Any
23
+
24
+ import numpy as np
25
+ import torch
26
+ import triton
27
+ import triton.language as tl
28
+ from triton.language.extra import libdevice
29
+
30
+ EXACT = {"enable_fp_fusion": False}
31
+
32
+
33
+ @triton.jit
34
+ def rsqrt_rn(x):
35
+ """``torch.rsqrt`` on ROCm is correctly rounded; OCML's FP32 rsqrt is not."""
36
+ return libdevice.rsqrt(x.to(tl.float64)).to(tl.float32)
37
+
38
+
39
+ @triton.jit
40
+ def bf16(x):
41
+ return x.to(tl.bfloat16).to(tl.float32)
42
+
43
+
44
+ @triton.jit
45
+ def lane_tree64(v, R: tl.constexpr):
46
+ """[R, 64] -> [R]: the shfl_down tree of offsets 1, 2, ..., 32."""
47
+ a, b = tl.split(tl.reshape(v, [R, 32, 2]))
48
+ v = a + b
49
+ a, b = tl.split(tl.reshape(v, [R, 16, 2]))
50
+ v = a + b
51
+ a, b = tl.split(tl.reshape(v, [R, 8, 2]))
52
+ v = a + b
53
+ a, b = tl.split(tl.reshape(v, [R, 4, 2]))
54
+ v = a + b
55
+ a, b = tl.split(tl.reshape(v, [R, 2, 2]))
56
+ v = a + b
57
+ a, b = tl.split(tl.reshape(v, [R, 1, 2]))
58
+ v = a + b
59
+ return tl.reshape(v, [R])
60
+
61
+
62
+ @triton.jit
63
+ def combine_vec4(acc, R: tl.constexpr):
64
+ """[R, 64, 4] accumulators -> [R, 64]: ((a0 + a1) + a2) + a3."""
65
+ even, odd = tl.split(tl.reshape(acc, [R, 64, 2, 2]))
66
+ a0, a2 = tl.split(even)
67
+ a1, a3 = tl.split(odd)
68
+ return ((a0 + a1) + a2) + a3
69
+
70
+
71
+ @triton.jit
72
+ def sumsq_128(x, R: tl.constexpr):
73
+ """[R, 128] -> [R]: 32 lanes of four accumulators, then a five-level lane tree."""
74
+ even, odd = tl.split(tl.reshape(x * x, [R, 32, 2, 2]))
75
+ a0, a2 = tl.split(even)
76
+ a1, a3 = tl.split(odd)
77
+ v = ((a0 + a1) + a2) + a3
78
+ a, b = tl.split(tl.reshape(v, [R, 16, 2]))
79
+ v = a + b
80
+ a, b = tl.split(tl.reshape(v, [R, 8, 2]))
81
+ v = a + b
82
+ a, b = tl.split(tl.reshape(v, [R, 4, 2]))
83
+ v = a + b
84
+ a, b = tl.split(tl.reshape(v, [R, 2, 2]))
85
+ v = a + b
86
+ a, b = tl.split(tl.reshape(v, [R, 1, 2]))
87
+ return tl.reshape(a + b, [R])
88
+
89
+
90
+ @triton.jit
91
+ def sumsq_256(x, R: tl.constexpr):
92
+ """[R, 256] -> [R]: one four-wide load per lane, then the lane tree."""
93
+ return lane_tree64(combine_vec4(tl.reshape(x * x, [R, 64, 4]), R), R)
94
+
95
+
96
+ @triton.jit
97
+ def _add_rmsnorm_kernel(
98
+ res_ptr,
99
+ delta_ptr,
100
+ w1_ptr,
101
+ rowmask_ptr,
102
+ hidden_ptr,
103
+ out_ptr,
104
+ M,
105
+ H,
106
+ inv_h,
107
+ eps,
108
+ HAS_DELTA: tl.constexpr,
109
+ HAS_MASK: tl.constexpr,
110
+ R: tl.constexpr,
111
+ ):
112
+ rows = tl.program_id(0) * R + tl.arange(0, R)
113
+ rmask = (rows < M)[:, None]
114
+ base = rows[:, None].to(tl.int64) * H
115
+ cols = tl.arange(0, 256)[None, :]
116
+ acc = tl.zeros([R, 64, 4], dtype=tl.float32)
117
+ for c in range(0, H // 256):
118
+ offs = base + c * 256 + cols
119
+ x = tl.load(res_ptr + offs, mask=rmask, other=0.0).to(tl.float32)
120
+ if HAS_DELTA:
121
+ x = bf16(
122
+ x + tl.load(delta_ptr + offs, mask=rmask, other=0.0).to(tl.float32)
123
+ )
124
+ tl.store(hidden_ptr + offs, x.to(tl.bfloat16), mask=rmask)
125
+ acc = acc + tl.reshape(x * x, [R, 64, 4])
126
+ var = lane_tree64(combine_vec4(acc, R), R) * inv_h
127
+ rstd = rsqrt_rn(var + eps)[:, None]
128
+ if HAS_MASK:
129
+ keep = tl.load(rowmask_ptr + rows, mask=rows < M, other=0).to(tl.float32)[
130
+ :, None
131
+ ]
132
+ for c in range(0, H // 256):
133
+ offs = base + c * 256 + cols
134
+ x = tl.load(res_ptr + offs, mask=rmask, other=0.0).to(tl.float32)
135
+ if HAS_DELTA:
136
+ x = bf16(
137
+ x + tl.load(delta_ptr + offs, mask=rmask, other=0.0).to(tl.float32)
138
+ )
139
+ y = bf16((x * rstd) * tl.load(w1_ptr + c * 256 + cols))
140
+ if HAS_MASK:
141
+ y = y * keep
142
+ tl.store(out_ptr + offs, y.to(tl.bfloat16), mask=rmask)
143
+
144
+
145
+ def add_rmsnorm(
146
+ residual: Any,
147
+ delta: Any | None,
148
+ weight_plus_one: Any,
149
+ eps: float,
150
+ rowmask: Any | None = None,
151
+ ) -> tuple[Any, Any]:
152
+ """(hidden, normed) BF16: ``hidden = residual + delta`` (or ``residual``), then ``Qwen3_5RMSNorm``.
153
+
154
+ ``weight_plus_one`` is the FP32 ``1 + w``; ``rowmask`` (one integer per row) zeroes the normed rows of
155
+ padding as the Gated DeltaNet layer's padding multiply does. Rows of a multiple of 256.
156
+ """
157
+ H = residual.shape[-1]
158
+ rows = residual.numel() // H
159
+ hidden = residual if delta is None else torch.empty_like(residual)
160
+ out = torch.empty_like(residual)
161
+ _add_rmsnorm_kernel[(triton.cdiv(rows, 2),)](
162
+ residual,
163
+ delta if delta is not None else residual,
164
+ weight_plus_one,
165
+ rowmask if rowmask is not None else residual,
166
+ hidden,
167
+ out,
168
+ rows,
169
+ H,
170
+ float(np.float32(1.0) / np.float32(H)),
171
+ eps,
172
+ HAS_DELTA=delta is not None,
173
+ HAS_MASK=rowmask is not None,
174
+ R=2,
175
+ num_warps=4,
176
+ **EXACT,
177
+ )
178
+ return hidden, out
179
+
180
+
181
+ @triton.jit
182
+ def _silu_mul_kernel(g_ptr, u_ptr, out_ptr, n_cols, BLOCK: tl.constexpr):
183
+ row = tl.program_id(0).to(tl.int64)
184
+ cols = tl.program_id(1) * BLOCK + tl.arange(0, BLOCK)
185
+ mask = cols < n_cols
186
+ g = tl.load(g_ptr + row * n_cols + cols, mask=mask, other=0.0).to(tl.float32)
187
+ u = tl.load(u_ptr + row * n_cols + cols, mask=mask, other=0.0).to(tl.float32)
188
+ s = bf16(g / (1.0 + libdevice.exp(-g)))
189
+ tl.store(out_ptr + row * n_cols + cols, (s * u).to(tl.bfloat16), mask=mask)
190
+
191
+
192
+ def silu_mul(gate: Any, up: Any) -> Any:
193
+ """``silu(gate) * up`` of two contiguous BF16 tensors of one shape."""
194
+ n = gate.shape[-1]
195
+ rows = gate.numel() // n
196
+ out = torch.empty_like(gate)
197
+ _silu_mul_kernel[(rows, triton.cdiv(n, 1024))](
198
+ gate, up, out, n, BLOCK=1024, num_warps=4, **EXACT
199
+ )
200
+ return out
201
+
202
+
203
+ @triton.jit
204
+ def _gdn_prep_kernel(
205
+ x_ptr,
206
+ w_ptr,
207
+ b_ptr,
208
+ a_ptr,
209
+ alog_ptr,
210
+ dtb_ptr,
211
+ q_ptr,
212
+ k_ptr,
213
+ v_ptr,
214
+ g_ptr,
215
+ beta_ptr,
216
+ T,
217
+ C,
218
+ NK,
219
+ NV,
220
+ DK: tl.constexpr,
221
+ REP: tl.constexpr,
222
+ BT: tl.constexpr,
223
+ KW: tl.constexpr,
224
+ ):
225
+ pid_t = tl.program_id(0)
226
+ j = tl.program_id(1)
227
+ nblk = tl.cdiv(T, BT)
228
+ bidx = (pid_t // nblk).to(tl.int64)
229
+ t = (pid_t % nblk) * BT + tl.arange(0, BT)
230
+ tmask = t < T
231
+ d = tl.arange(0, DK)
232
+ c = j * DK + d
233
+ row = bidx * T + t
234
+ acc = tl.zeros([BT, DK], dtype=tl.float64)
235
+ for w in tl.static_range(KW):
236
+ tt = t - (KW - 1) + w
237
+ m = (tt >= 0) & tmask
238
+ xv = tl.load(
239
+ x_ptr + (bidx * T + tt)[:, None] * C + c[None, :],
240
+ mask=m[:, None],
241
+ other=0.0,
242
+ )
243
+ wv = tl.load(w_ptr + c * KW + w)
244
+ acc = acc + wv.to(tl.float64)[None, :] * xv.to(tl.float64)
245
+ conv = bf16(acc.to(tl.float32))
246
+ y = (conv / (1.0 + libdevice.exp(-conv))).to(tl.bfloat16)
247
+ if j < NK:
248
+ for r in tl.static_range(REP):
249
+ tl.store(
250
+ q_ptr + row[:, None] * (NV * DK) + ((j * REP + r) * DK + d)[None, :],
251
+ y,
252
+ mask=tmask[:, None],
253
+ )
254
+ elif j < 2 * NK:
255
+ for r in tl.static_range(REP):
256
+ tl.store(
257
+ k_ptr
258
+ + row[:, None] * (NV * DK)
259
+ + (((j - NK) * REP + r) * DK + d)[None, :],
260
+ y,
261
+ mask=tmask[:, None],
262
+ )
263
+ else:
264
+ hv = j - 2 * NK
265
+ tl.store(
266
+ v_ptr + row[:, None] * (NV * DK) + (hv * DK + d)[None, :],
267
+ y,
268
+ mask=tmask[:, None],
269
+ )
270
+ bb = tl.load(b_ptr + row * NV + hv, mask=tmask, other=0.0).to(tl.float32)
271
+ tl.store(
272
+ beta_ptr + row * NV + hv,
273
+ (1.0 / (1.0 + libdevice.exp(-bb))).to(tl.bfloat16),
274
+ mask=tmask,
275
+ )
276
+ s = tl.load(a_ptr + row * NV + hv, mask=tmask, other=0.0).to(
277
+ tl.float32
278
+ ) + tl.load(dtb_ptr + hv)
279
+ sp = tl.where(s > 20.0, s, libdevice.log1p(libdevice.exp(s)))
280
+ g = -libdevice.exp(tl.load(alog_ptr + hv)) * sp
281
+ tl.store(g_ptr + row * NV + hv, g, mask=tmask)
282
+
283
+
284
+ def gdn_prep(
285
+ mixed_qkv: Any,
286
+ b: Any,
287
+ a: Any,
288
+ conv_weight: Any,
289
+ A_log: Any,
290
+ dt_bias: Any,
291
+ k_heads: int,
292
+ head_dim: int,
293
+ ) -> tuple[Any, Any, Any, Any, Any]:
294
+ """(q, k, v, g, beta) as the eager layer hands them to the chunk kernel (q / k repeated to the value heads).
295
+
296
+ ``mixed_qkv`` is the contiguous [B, T, C] BF16 ``in_proj_qkv`` output, ``conv_weight`` the [C, KW] BF16
297
+ depthwise filter, ``A_log`` / ``dt_bias`` FP32 copies of the parameters; q / k / v are SiLU(conv) in BF16,
298
+ beta = sigmoid(b) in BF16 and g = -exp(A_log) * softplus(a + dt_bias) in FP32.
299
+ """
300
+ B, T, C = mixed_qkv.shape
301
+ nv = (C - 2 * k_heads * head_dim) // head_dim
302
+ dev = mixed_qkv.device
303
+ q = torch.empty(B, T, nv, head_dim, dtype=torch.bfloat16, device=dev)
304
+ k = torch.empty_like(q)
305
+ v = torch.empty_like(q)
306
+ g = torch.empty(B, T, nv, dtype=torch.float32, device=dev)
307
+ beta = torch.empty(B, T, nv, dtype=torch.bfloat16, device=dev)
308
+ _gdn_prep_kernel[(B * triton.cdiv(T, 16), 2 * k_heads + nv)](
309
+ mixed_qkv,
310
+ conv_weight,
311
+ b,
312
+ a,
313
+ A_log,
314
+ dt_bias,
315
+ q,
316
+ k,
317
+ v,
318
+ g,
319
+ beta,
320
+ T,
321
+ C,
322
+ k_heads,
323
+ nv,
324
+ DK=head_dim,
325
+ REP=nv // k_heads,
326
+ BT=16,
327
+ KW=conv_weight.shape[-1],
328
+ num_warps=4,
329
+ **EXACT,
330
+ )
331
+ return q, k, v, g, beta
332
+
333
+
334
+ @triton.jit
335
+ def _gated_rmsnorm_kernel(
336
+ x_ptr, z_ptr, w_ptr, out_ptr, N, eps, inv_d, D: tl.constexpr, BR: tl.constexpr
337
+ ):
338
+ r = (tl.program_id(0) * BR + tl.arange(0, BR)).to(tl.int64)
339
+ d = tl.arange(0, D)
340
+ m = (r < N)[:, None]
341
+ offs = r[:, None] * D + d[None, :]
342
+ x = tl.load(x_ptr + offs, mask=m, other=0.0).to(tl.float32)
343
+ rstd = rsqrt_rn(sumsq_128(x, BR) * inv_d + eps)
344
+ xn = bf16(x * rstd[:, None])
345
+ y = bf16(tl.load(w_ptr + d).to(tl.float32)[None, :] * xn)
346
+ z = tl.load(z_ptr + offs, mask=m, other=0.0).to(tl.float32)
347
+ y = y * (z / (1.0 + libdevice.exp(-z)))
348
+ tl.store(out_ptr + offs, y.to(tl.bfloat16), mask=m)
349
+
350
+
351
+ def gated_rmsnorm(core: Any, z: Any, weight: Any, eps: float) -> Any:
352
+ """``Qwen3_5RMSNormGated`` (BF16 weight) on contiguous [N, 128] BF16 rows ``core`` and gates ``z``."""
353
+ D = core.shape[-1]
354
+ n = core.numel() // D
355
+ out = torch.empty_like(core)
356
+ _gated_rmsnorm_kernel[(triton.cdiv(n, 16),)](
357
+ core,
358
+ z,
359
+ weight,
360
+ out,
361
+ n,
362
+ eps,
363
+ float(np.float32(1.0) / np.float32(D)),
364
+ D=D,
365
+ BR=16,
366
+ num_warps=4,
367
+ **EXACT,
368
+ )
369
+ return out
370
+
371
+
372
+ @triton.jit
373
+ def _attn_prep_kernel(
374
+ x_ptr,
375
+ w1_ptr,
376
+ cos_ptr,
377
+ sin_ptr,
378
+ o_ptr,
379
+ T,
380
+ NH,
381
+ eps,
382
+ inv_d,
383
+ cs_bstride,
384
+ D: tl.constexpr,
385
+ ROT: tl.constexpr,
386
+ QW: tl.constexpr,
387
+ REP: tl.constexpr,
388
+ ):
389
+ bt = tl.program_id(0).to(tl.int64)
390
+ hid = tl.program_id(1)
391
+ b = bt // T
392
+ t = bt % T
393
+ d = tl.arange(0, D)
394
+ half: tl.constexpr = ROT // 2
395
+ partner = tl.where(d < half, d + half, tl.where(d < ROT, d - half, d))
396
+ src = x_ptr + bt * (NH * QW) + hid * QW
397
+ x = tl.load(src + d).to(tl.float32)
398
+ sumsq = sumsq_256(tl.reshape(x, [1, D]), 1)
399
+ rstd = rsqrt_rn(tl.reshape(sumsq, []) * inv_d + eps)
400
+ xp = tl.load(src + partner).to(tl.float32)
401
+ y = bf16((x * rstd) * tl.load(w1_ptr + d))
402
+ yp = bf16((xp * rstd) * tl.load(w1_ptr + partner))
403
+ yp = tl.where(d < half, -yp, yp)
404
+ rot = d < ROT
405
+ cs = b * cs_bstride + t * ROT + d
406
+ c = tl.load(cos_ptr + cs, mask=rot, other=1.0).to(tl.float32)
407
+ s = tl.load(sin_ptr + cs, mask=rot, other=0.0).to(tl.float32)
408
+ out = tl.where(rot, bf16(bf16(y * c) + bf16(yp * s)), y).to(tl.bfloat16)
409
+ for r in tl.static_range(REP):
410
+ tl.store(o_ptr + ((b * NH * REP + hid * REP + r) * T + t) * D + d, out)
411
+
412
+
413
+ def attn_prep(
414
+ q_proj_out: Any,
415
+ k_proj_out: Any,
416
+ q_norm_w1: Any,
417
+ k_norm_w1: Any,
418
+ cos: Any,
419
+ sin: Any,
420
+ heads: int,
421
+ kv_heads: int,
422
+ head_dim: int,
423
+ eps: float,
424
+ ) -> tuple[Any, Any]:
425
+ """(q [B, H, T, D], k [B, H, T, D]) in BF16: head RMSNorms, then the partial RoPE, k repeated to H heads.
426
+
427
+ ``q_proj_out`` [B, T, H * 2D] holds query and gate per head, ``k_proj_out`` [B, T, Hkv * D] (both
428
+ contiguous BF16); the norms multiply by ``1 + w`` (FP32) and round to BF16 before the RoPE, whose products
429
+ and sum round to BF16 like the eager BF16 ops; ``cos`` / ``sin`` are the contiguous [B or 1, T, rotary dim]
430
+ BF16 rotary tables; the head dim is 256.
431
+ """
432
+ if head_dim != 256:
433
+ raise ValueError("attn_prep reduces rows of 256")
434
+ B, T = q_proj_out.shape[:2]
435
+ dev = q_proj_out.device
436
+ q = torch.empty(B, heads, T, head_dim, dtype=torch.bfloat16, device=dev)
437
+ k = torch.empty_like(q)
438
+ rot = cos.shape[-1]
439
+ # One launch per tensor: Triton's AMD pointer canonicalization fails on a runtime branch between two
440
+ # pointers when only one of their tensors fits the 2 GiB buffer range.
441
+ for x, w1, out, n, width, rep in (
442
+ (q_proj_out, q_norm_w1, q, heads, 2 * head_dim, 1),
443
+ (k_proj_out, k_norm_w1, k, kv_heads, head_dim, heads // kv_heads),
444
+ ):
445
+ _attn_prep_kernel[(B * T, n)](
446
+ x,
447
+ w1,
448
+ cos,
449
+ sin,
450
+ out,
451
+ T,
452
+ n,
453
+ eps,
454
+ float(np.float32(1.0) / np.float32(head_dim)),
455
+ 0 if cos.shape[0] == 1 else T * rot,
456
+ D=head_dim,
457
+ ROT=rot,
458
+ QW=width,
459
+ REP=rep,
460
+ num_warps=2,
461
+ **EXACT,
462
+ )
463
+ return q, k
464
+
465
+
466
+ @triton.jit
467
+ def _sigmoid_gate_kernel(
468
+ a_ptr,
469
+ g_ptr,
470
+ out_ptr,
471
+ T,
472
+ H,
473
+ sab,
474
+ sat,
475
+ sah,
476
+ sgb,
477
+ sgt,
478
+ sgh,
479
+ D: tl.constexpr,
480
+ HB: tl.constexpr,
481
+ ):
482
+ bt = tl.program_id(0).to(tl.int64)
483
+ h = tl.program_id(1) * HB + tl.arange(0, HB)
484
+ b = bt // T
485
+ t = bt % T
486
+ d = tl.arange(0, D)
487
+ m = (h < H)[:, None]
488
+ a = tl.load(
489
+ a_ptr + b * sab + t * sat + h[:, None] * sah + d[None, :], mask=m, other=0.0
490
+ ).to(tl.float32)
491
+ gt = tl.load(
492
+ g_ptr + b * sgb + t * sgt + h[:, None] * sgh + d[None, :], mask=m, other=0.0
493
+ ).to(tl.float32)
494
+ s = bf16(1.0 / (1.0 + libdevice.exp(-gt)))
495
+ tl.store(
496
+ out_ptr + bt * (H * D) + h[:, None] * D + d[None, :],
497
+ (a * s).to(tl.bfloat16),
498
+ mask=m,
499
+ )
500
+
501
+
502
+ def sigmoid_gate(attn_out: Any, gate: Any) -> Any:
503
+ """``attn_out * sigmoid(gate)`` -> [B, T, H * D] BF16 from [B, T, H, D] views with unit last stride."""
504
+ B, T, H, D = attn_out.shape
505
+ out = torch.empty(B, T, H * D, dtype=torch.bfloat16, device=attn_out.device)
506
+ _sigmoid_gate_kernel[(B * T, triton.cdiv(H, 4))](
507
+ attn_out,
508
+ gate,
509
+ out,
510
+ T,
511
+ H,
512
+ attn_out.stride(0),
513
+ attn_out.stride(1),
514
+ attn_out.stride(2),
515
+ gate.stride(0),
516
+ gate.stride(1),
517
+ gate.stride(2),
518
+ D=D,
519
+ HB=4,
520
+ num_warps=4,
521
+ **EXACT,
522
+ )
523
+ return out
d3_runtime.py CHANGED
@@ -13,7 +13,8 @@ probability for every option. No text is generated and no input is truncated.
13
  model.system_one(state="...", questions={...}, images=["photo.png", "label.jpg"])
14
 
15
  The checkpoint directory holds ``config.json`` + ``model*.safetensors`` (a transformers
16
- ``Qwen3_5Model``), ``readout.safetensors`` (``{"weight": [255, hidden]}``), ``decision_config.json``
 
17
  (prompt family, answer codes, attention mode, pooling, temperature, input limit) and the tokenizer.
18
  ``d3_format.py`` next to this file is the prompt and answer-code contract of the model.
19
 
@@ -27,10 +28,17 @@ with the vision tower (``visual.*`` weights).
27
 
28
  Numerics: BF16 backbone with SDPA attention, FP32 readout and softmax (unless the checkpoint says
29
  otherwise). A request's questions run in request order, ``batch_size`` per forward pass, each batch
30
- left-padded to its longest prompt. ``noncausal_full_attention`` lets the full-attention layers see the
31
- whole prompt while the Gated DeltaNet layers stay causal. The Gated DeltaNet kernels are the ones
32
- transformers binds at import: flash-linear-attention (and causal-conv1d) when installed, its PyTorch
33
- reference implementation otherwise.
 
 
 
 
 
 
 
34
 
35
  The noncausal attention mask hook is adapted from perplexity-ai/pplx-decider-v1.1-27b, Copyright
36
  Perplexity AI, Apache License 2.0.
@@ -73,11 +81,19 @@ except ImportError:
73
  to_answer,
74
  user_prompt,
75
  )
 
 
 
 
 
 
 
76
 
77
  RUNTIME = "d3-runtime/1"
78
  FORMAT_VERSION = 1
79
  PROMPTS = ("d3",)
80
  ATTENTION_MODES = ("causal", "noncausal_full_attention")
 
81
  DEFAULT_BATCH_SIZE = 8
82
  SCORE_LEVELS = (2, 10)
83
  MANIFEST = "MODEL_MANIFEST.json"
@@ -97,6 +113,15 @@ MAX_IMAGE_SOURCE_PIXELS = 16_000_000
97
  IMAGE_FORMATS = ("PNG", "JPEG", "WEBP")
98
  DOWNLOAD_TIMEOUT_SECONDS = 30
99
  MAX_DOWNLOAD_BYTES = 64 << 20
 
 
 
 
 
 
 
 
 
100
 
101
 
102
  class MaxLengthExceeded(ValueError):
@@ -114,9 +139,7 @@ class Question:
114
 
115
  kind: str # choice | noul | score
116
  original: Mapping[str, Any]
117
- rendered: dict[
118
- str, Any
119
- ] # a choice or noul question in the d3_format contract
120
  keys: list[str]
121
  descriptions: list[Any]
122
 
@@ -216,6 +239,21 @@ def image_messages(
216
  raise ValueError(f"unknown prompt family {prompt!r}")
217
 
218
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
219
  def _canonical(value: Any) -> str:
220
  return (
221
  value
@@ -272,10 +310,19 @@ def product_answer(
272
  def _data_url_payload(value: str, strict: bool) -> bytes:
273
  header, separator, encoded = value.partition(",")
274
  kind = header.strip().lower()
275
- if not separator or not kind.startswith("data:image/") or not kind.endswith(";base64"):
 
 
 
 
276
  raise ValueError("a data URL image is data:image/<format>;base64,<data>")
277
  if strict:
278
- if kind[len("data:image/") : -len(";base64")] not in ("png", "jpeg", "jpg", "webp"):
 
 
 
 
 
279
  raise ValueError("images must be base64 PNG, JPEG or WebP data URLs")
280
  if len(encoded) > 4 * -(-MAX_IMAGE_BYTES // 3):
281
  raise ValueError(f"each image must be at most {MAX_IMAGE_BYTES:,} bytes")
@@ -292,7 +339,9 @@ def _download(url: str) -> bytes:
292
  with urllib.request.urlopen(request, timeout=DOWNLOAD_TIMEOUT_SECONDS) as response:
293
  payload = response.read(MAX_DOWNLOAD_BYTES + 1)
294
  if len(payload) > MAX_DOWNLOAD_BYTES:
295
- raise ValueError(f"the image at {url} is larger than {MAX_DOWNLOAD_BYTES:,} bytes")
 
 
296
  return payload
297
 
298
 
@@ -339,7 +388,12 @@ def load_image(value: Any, *, strict: bool = False):
339
  image.verify()
340
  with Image.open(io.BytesIO(payload)) as image:
341
  return image.convert("RGB")
342
- except (OSError, SyntaxError, UnidentifiedImageError, Image.DecompressionBombError) as exc:
 
 
 
 
 
343
  raise ValueError(f"invalid image data ({type(exc).__name__})") from exc
344
 
345
 
@@ -473,6 +527,27 @@ def vision_weights_present(root: Path) -> bool:
473
  # ---------------------------------------------------------------------------------------------
474
 
475
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
476
  def enable_noncausal_full_attention(text_model) -> None:
477
  """Let the softmax-attention layers see future tokens; keep padding and the causal recurrence.
478
 
@@ -586,6 +661,10 @@ class Prepared:
586
  images: list[Any] = field(default_factory=list)
587
  texts: dict[str, str] = field(default_factory=dict)
588
  lengths: dict[str, int] = field(default_factory=dict)
 
 
 
 
589
 
590
  @property
591
  def runnable(self) -> list[str]:
@@ -605,11 +684,11 @@ class D3:
605
  max_length: int | None = None,
606
  readout_dtype: str | None = None,
607
  model_name: str | None = None,
 
608
  ):
609
  import torch
610
  from safetensors.torch import load_file
611
  from transformers import AutoTokenizer
612
- from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5Model
613
 
614
  self.torch = torch
615
  self.root = Path(root)
@@ -621,12 +700,20 @@ class D3:
621
  if config.get("format_version") != FORMAT_VERSION:
622
  raise ValueError("unsupported decision_config.json format_version")
623
  self.config = config
 
624
  self.prompt = config.get("prompt", "d3")
625
  if self.prompt not in PROMPTS:
626
  raise ValueError(f"unknown prompt family {self.prompt!r}")
627
  self.attention_mode = config.get("attention_mode", "causal")
628
  if self.attention_mode not in ATTENTION_MODES:
629
  raise ValueError(f"unknown attention mode {self.attention_mode!r}")
 
 
 
 
 
 
 
630
  if config.get("pooling", "last") != "last":
631
  raise ValueError(f"unsupported pooling {config.get('pooling')!r}")
632
  self.temperature = float(config.get("temperature", 1.0))
@@ -643,6 +730,7 @@ class D3:
643
  self.model_name = (
644
  model_name or (manifest or {}).get("model_name") or self.root.name
645
  )
 
646
 
647
  self.tokenizer = AutoTokenizer.from_pretrained(str(self.root))
648
  self.tokenizer.padding_side = "left"
@@ -667,7 +755,7 @@ class D3:
667
  self.device = torch.device(device)
668
  if self.device.type == "cuda":
669
  torch.cuda.set_device(self.device)
670
- self.kernels = kernel_report()
671
  if self.device.type == "cpu" and any(
672
  v.startswith(("fla", "causal_conv1d"))
673
  for k, v in self.kernels.items()
@@ -678,7 +766,7 @@ class D3:
678
  "use an environment without them for CPU inference"
679
  )
680
  torch.manual_seed(20260919)
681
- self.backbone = Qwen3_5Model.from_pretrained(
682
  str(self.root),
683
  dtype=torch.bfloat16,
684
  attn_implementation="sdpa",
@@ -697,16 +785,22 @@ class D3:
697
  self.processor = None
698
  self.image_unavailable: str | None = None
699
  if not vision_weights_present(self.root):
700
- self.image_unavailable = "the checkpoint has no vision tower (visual.* weights)"
 
 
701
  elif not linearize_patch_embed(self.backbone):
702
  self.image_unavailable = "the vision patch embedding was not found"
703
  else:
704
  try:
705
  self.processor = load_processor(self.root)
706
- except Exception as exc: # noqa: BLE001 - text requests do not use the processor
 
 
707
  self.image_unavailable = (
708
  f"the image processor failed to load ({type(exc).__name__}: {exc})"
709
  )
 
 
710
  self.loaded_seconds = time.perf_counter() - started
711
 
712
  @classmethod
@@ -725,12 +819,14 @@ class D3:
725
  max_length: int | None = None,
726
  readout_dtype: str | None = None,
727
  model_name: str | None = None,
 
728
  ) -> D3:
729
  """Load a package directory or Hub repository.
730
 
731
  ``verify``: ``fast`` (default) hashes every file of ``MODEL_MANIFEST.json`` up to 64 MiB and checks
732
  the size of the weight shards; ``full`` hashes every file; ``none`` skips the check. A checkpoint
733
- without a manifest (a plain code-readout export) loads unverified.
 
734
  """
735
  root = resolve_dir(
736
  name_or_path,
@@ -749,6 +845,7 @@ class D3:
749
  max_length=max_length,
750
  readout_dtype=readout_dtype,
751
  model_name=model_name,
 
752
  )
753
 
754
  # ------------------------------------------------------------------ requests
@@ -769,7 +866,9 @@ class D3:
769
  if not images:
770
  return []
771
  if self.image_unavailable is not None:
772
- raise ValueError(f"image inputs are not available: {self.image_unavailable}")
 
 
773
  decoded = []
774
  for number, value in enumerate(images):
775
  try:
@@ -829,8 +928,53 @@ class D3:
829
  }
830
  continue
831
  prepared.sequences[key] = sequence
 
 
832
  return prepared
833
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
834
  def _prepare_images(
835
  self, state: Any, questions: Mapping[str, Any], images: list[Any]
836
  ) -> Prepared:
@@ -882,6 +1026,8 @@ class D3:
882
  continue
883
  prepared.texts[key] = text
884
  prepared.lengths[key] = length
 
 
885
  return prepared
886
 
887
  def logits(self, sequences: Sequence[Sequence[int]], counts: Sequence[int]):
@@ -914,6 +1060,8 @@ class D3:
914
  self, sequences: Sequence[Sequence[int]], counts: Sequence[int]
915
  ) -> list[list[float]]:
916
  """Softmax over each prompt's own codes, in option order."""
 
 
917
  probs = (
918
  (self.logits(sequences, counts) / self.temperature)
919
  .softmax(-1)
@@ -946,8 +1094,15 @@ class D3:
946
  f"planned {width} input tokens, the processor produced {encoded['input_ids'].shape[1]}"
947
  )
948
  inputs = {name: value.to(self.device) for name, value in encoded.items()}
 
949
  with torch.inference_mode(), sdpa_backends(self.device):
950
- hidden = self.backbone(**inputs, use_cache=False).last_hidden_state[:, -1]
 
 
 
 
 
 
951
  if self.readout_dtype == "float32":
952
  logits = hidden.float() @ self.readout.T
953
  else:
@@ -976,14 +1131,31 @@ class D3:
976
  out: dict[str, list[float]] = {}
977
  for start in range(0, len(keys), self.batch_size):
978
  chunk = keys[start : start + self.batch_size]
979
- probs = self.image_probabilities(
980
- [prepared.texts[k] for k in chunk],
981
- prepared.images,
982
- [len(prepared.questions[k].keys) for k in chunk],
983
- max(prepared.lengths[k] for k in chunk),
984
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
985
  out.update(zip(chunk, probs))
986
- return out, sum(prepared.lengths[k] for k in keys)
 
 
 
 
987
 
988
  def run(self, prepared: Prepared) -> tuple[dict[str, list[float]], int]:
989
  """Probabilities per runnable question (request order, ``batch_size`` per pass) and the input tokens."""
@@ -993,12 +1165,24 @@ class D3:
993
  out: dict[str, list[float]] = {}
994
  for start in range(0, len(keys), self.batch_size):
995
  chunk = keys[start : start + self.batch_size]
996
- probs = self.probabilities(
997
- [prepared.sequences[k] for k in chunk],
998
- [len(prepared.questions[k].keys) for k in chunk],
999
- )
 
 
 
 
 
 
 
 
1000
  out.update(zip(chunk, probs))
1001
- return out, sum(len(prepared.sequences[k]) for k in keys)
 
 
 
 
1002
 
1003
  def respond(
1004
  self,
@@ -1049,7 +1233,8 @@ class D3:
1049
  lengths: Sequence[int] = (37, 64, 320, 333, 1000, 1024),
1050
  images: bool = True,
1051
  ) -> float:
1052
- """Compile and autotune the kernels for every batch size up to ``batch_size`` before serving.
 
1053
 
1054
  The Gated DeltaNet kernels take the batch size as a compile-time constant, and Triton specializes their
1055
  length and chunk-count arguments on being 1 or a multiple of 16; these lengths cover every combination
@@ -1058,7 +1243,10 @@ class D3:
1058
  four 1.6 MP images) also warm the vision tower. Answers are unchanged.
1059
  """
1060
  started = time.perf_counter()
1061
- for size in range(1, self.batch_size + 1):
 
 
 
1062
  for length in lengths:
1063
  sequences = [
1064
  [
@@ -1091,6 +1279,8 @@ class D3:
1091
  self.backbone.to(target)
1092
  self.readout = self.readout.to(target)
1093
  self.device = target
 
 
1094
  return self
1095
 
1096
  def parameter_count(self) -> int:
@@ -1108,7 +1298,7 @@ class D3:
1108
  if (self.root / name).is_file()
1109
  }
1110
  identity = (self.manifest or {}).get("identity", {})
1111
- return {
1112
  "kind": "d3-code-readout",
1113
  "runtime": RUNTIME,
1114
  "model_name": self.model_name,
@@ -1131,6 +1321,13 @@ class D3:
1131
  "refused); no option filtering; one fixed prompt for every request.",
1132
  "images": self.image_contract(),
1133
  }
 
 
 
 
 
 
 
1134
 
1135
  def image_contract(self) -> dict[str, Any]:
1136
  """How image inputs are read (or why they are not available)."""
@@ -1167,4 +1364,11 @@ class D3:
1167
  }
1168
  if self.device.type == "cuda":
1169
  info["gpu"] = torch.cuda.get_device_name(self.device)
 
1170
  return info
 
 
 
 
 
 
 
13
  model.system_one(state="...", questions={...}, images=["photo.png", "label.jpg"])
14
 
15
  The checkpoint directory holds ``config.json`` + ``model*.safetensors`` (a transformers
16
+ ``Qwen3_5Model``, or a ``Qwen3VLModel`` when ``config.json`` says ``model_type: qwen3_vl``),
17
+ ``readout.safetensors`` (``{"weight": [255, hidden]}``), ``decision_config.json``
18
  (prompt family, answer codes, attention mode, pooling, temperature, input limit) and the tokenizer.
19
  ``d3_format.py`` next to this file is the prompt and answer-code contract of the model.
20
 
 
28
 
29
  Numerics: BF16 backbone with SDPA attention, FP32 readout and softmax (unless the checkpoint says
30
  otherwise). A request's questions run in request order, ``batch_size`` per forward pass, each batch
31
+ left-padded to its longest prompt. ``noncausal_full_attention`` (Qwen3.5 backbones only) lets the
32
+ full-attention layers see the whole prompt while the Gated DeltaNet layers stay causal. The Gated
33
+ DeltaNet kernels are the ones transformers binds at import: flash-linear-attention (and causal-conv1d)
34
+ when installed, its PyTorch reference implementation otherwise. Qwen3-VL backbones have no linear-attention
35
+ layers and run causal attention only.
36
+
37
+ ``permutation_average=True`` (off by default) also scores every choice question with two or more options
38
+ with its options in reversed order, in the same forward passes as the original order, and answers with
39
+ the per-option mean of the two distributions. Noul and score questions are scored once. A pass carrying both
40
+ orders that would exceed ``MERGE_TOKENS`` padded tokens runs the reversed prompts in a pass of their own, so
41
+ peak memory stays that of the original order.
42
 
43
  The noncausal attention mask hook is adapted from perplexity-ai/pplx-decider-v1.1-27b, Copyright
44
  Perplexity AI, Apache License 2.0.
 
81
  to_answer,
82
  user_prompt,
83
  )
84
+ try:
85
+ from . import d3_fast
86
+ except ImportError:
87
+ try:
88
+ import d3_fast
89
+ except ImportError:
90
+ d3_fast = None
91
 
92
  RUNTIME = "d3-runtime/1"
93
  FORMAT_VERSION = 1
94
  PROMPTS = ("d3",)
95
  ATTENTION_MODES = ("causal", "noncausal_full_attention")
96
+ BACKBONES = ("qwen3_5", "qwen3_vl")
97
  DEFAULT_BATCH_SIZE = 8
98
  SCORE_LEVELS = (2, 10)
99
  MANIFEST = "MODEL_MANIFEST.json"
 
113
  IMAGE_FORMATS = ("PNG", "JPEG", "WEBP")
114
  DOWNLOAD_TIMEOUT_SECONDS = 30
115
  MAX_DOWNLOAD_BYTES = 64 << 20
116
+ MERGE_TOKENS = 8192
117
+ # What the processor returns for an image batch; the fast path runs exactly these through the fused layers.
118
+ IMAGE_INPUTS = (
119
+ "input_ids",
120
+ "attention_mask",
121
+ "mm_token_type_ids",
122
+ "pixel_values",
123
+ "image_grid_thw",
124
+ )
125
 
126
 
127
  class MaxLengthExceeded(ValueError):
 
139
 
140
  kind: str # choice | noul | score
141
  original: Mapping[str, Any]
142
+ rendered: dict[str, Any] # a choice or noul question in the d3_format contract
 
 
143
  keys: list[str]
144
  descriptions: list[Any]
145
 
 
239
  raise ValueError(f"unknown prompt family {prompt!r}")
240
 
241
 
242
+ def reversed_question(question: Question) -> dict[str, Any] | None:
243
+ """The rendered choice question with its options in reverse order (None when there is nothing to permute)."""
244
+ if question.kind != "choice" or len(question.keys) < 2:
245
+ return None
246
+ criteria = question.rendered["criteria"]
247
+ return dict(
248
+ question.rendered, criteria={key: criteria[key] for key in reversed(criteria)}
249
+ )
250
+
251
+
252
+ def average_orders(forward: Sequence[float], backward: Sequence[float]) -> list[float]:
253
+ """Per-option mean of the original-order and reversed-order distributions, in original option order."""
254
+ return [(a + b) / 2 for a, b in zip(forward, reversed(backward))]
255
+
256
+
257
  def _canonical(value: Any) -> str:
258
  return (
259
  value
 
310
  def _data_url_payload(value: str, strict: bool) -> bytes:
311
  header, separator, encoded = value.partition(",")
312
  kind = header.strip().lower()
313
+ if (
314
+ not separator
315
+ or not kind.startswith("data:image/")
316
+ or not kind.endswith(";base64")
317
+ ):
318
  raise ValueError("a data URL image is data:image/<format>;base64,<data>")
319
  if strict:
320
+ if kind[len("data:image/") : -len(";base64")] not in (
321
+ "png",
322
+ "jpeg",
323
+ "jpg",
324
+ "webp",
325
+ ):
326
  raise ValueError("images must be base64 PNG, JPEG or WebP data URLs")
327
  if len(encoded) > 4 * -(-MAX_IMAGE_BYTES // 3):
328
  raise ValueError(f"each image must be at most {MAX_IMAGE_BYTES:,} bytes")
 
339
  with urllib.request.urlopen(request, timeout=DOWNLOAD_TIMEOUT_SECONDS) as response:
340
  payload = response.read(MAX_DOWNLOAD_BYTES + 1)
341
  if len(payload) > MAX_DOWNLOAD_BYTES:
342
+ raise ValueError(
343
+ f"the image at {url} is larger than {MAX_DOWNLOAD_BYTES:,} bytes"
344
+ )
345
  return payload
346
 
347
 
 
388
  image.verify()
389
  with Image.open(io.BytesIO(payload)) as image:
390
  return image.convert("RGB")
391
+ except (
392
+ OSError,
393
+ SyntaxError,
394
+ UnidentifiedImageError,
395
+ Image.DecompressionBombError,
396
+ ) as exc:
397
  raise ValueError(f"invalid image data ({type(exc).__name__})") from exc
398
 
399
 
 
527
  # ---------------------------------------------------------------------------------------------
528
 
529
 
530
+ def backbone_type(root: Path) -> str:
531
+ """``model_type`` of the checkpoint's ``config.json``: ``qwen3_5`` or ``qwen3_vl``."""
532
+ kind = json.loads((Path(root) / "config.json").read_text(encoding="utf-8")).get(
533
+ "model_type"
534
+ )
535
+ if kind not in BACKBONES:
536
+ raise ValueError(f"unsupported backbone model_type {kind!r}")
537
+ return kind
538
+
539
+
540
+ def backbone_class(kind: str):
541
+ """The transformers backbone class of a ``model_type``."""
542
+ if kind == "qwen3_vl":
543
+ from transformers.models.qwen3_vl.modeling_qwen3_vl import Qwen3VLModel
544
+
545
+ return Qwen3VLModel
546
+ from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5Model
547
+
548
+ return Qwen3_5Model
549
+
550
+
551
  def enable_noncausal_full_attention(text_model) -> None:
552
  """Let the softmax-attention layers see future tokens; keep padding and the causal recurrence.
553
 
 
661
  images: list[Any] = field(default_factory=list)
662
  texts: dict[str, str] = field(default_factory=dict)
663
  lengths: dict[str, int] = field(default_factory=dict)
664
+ # Reversed-option prompts of the choice questions (permutation_average only).
665
+ reversed_sequences: dict[str, list[int]] = field(default_factory=dict)
666
+ reversed_texts: dict[str, str] = field(default_factory=dict)
667
+ reversed_lengths: dict[str, int] = field(default_factory=dict)
668
 
669
  @property
670
  def runnable(self) -> list[str]:
 
684
  max_length: int | None = None,
685
  readout_dtype: str | None = None,
686
  model_name: str | None = None,
687
+ permutation_average: bool = False,
688
  ):
689
  import torch
690
  from safetensors.torch import load_file
691
  from transformers import AutoTokenizer
 
692
 
693
  self.torch = torch
694
  self.root = Path(root)
 
700
  if config.get("format_version") != FORMAT_VERSION:
701
  raise ValueError("unsupported decision_config.json format_version")
702
  self.config = config
703
+ self.backbone_type = backbone_type(self.root)
704
  self.prompt = config.get("prompt", "d3")
705
  if self.prompt not in PROMPTS:
706
  raise ValueError(f"unknown prompt family {self.prompt!r}")
707
  self.attention_mode = config.get("attention_mode", "causal")
708
  if self.attention_mode not in ATTENTION_MODES:
709
  raise ValueError(f"unknown attention mode {self.attention_mode!r}")
710
+ if (
711
+ self.attention_mode == "noncausal_full_attention"
712
+ and self.backbone_type != "qwen3_5"
713
+ ):
714
+ raise ValueError(
715
+ "noncausal_full_attention is defined for Qwen3.5 backbones only"
716
+ )
717
  if config.get("pooling", "last") != "last":
718
  raise ValueError(f"unsupported pooling {config.get('pooling')!r}")
719
  self.temperature = float(config.get("temperature", 1.0))
 
730
  self.model_name = (
731
  model_name or (manifest or {}).get("model_name") or self.root.name
732
  )
733
+ self.permutation_average = bool(permutation_average)
734
 
735
  self.tokenizer = AutoTokenizer.from_pretrained(str(self.root))
736
  self.tokenizer.padding_side = "left"
 
755
  self.device = torch.device(device)
756
  if self.device.type == "cuda":
757
  torch.cuda.set_device(self.device)
758
+ self.kernels = kernel_report() if self.backbone_type == "qwen3_5" else {}
759
  if self.device.type == "cpu" and any(
760
  v.startswith(("fla", "causal_conv1d"))
761
  for k, v in self.kernels.items()
 
766
  "use an environment without them for CPU inference"
767
  )
768
  torch.manual_seed(20260919)
769
+ self.backbone = backbone_class(self.backbone_type).from_pretrained(
770
  str(self.root),
771
  dtype=torch.bfloat16,
772
  attn_implementation="sdpa",
 
785
  self.processor = None
786
  self.image_unavailable: str | None = None
787
  if not vision_weights_present(self.root):
788
+ self.image_unavailable = (
789
+ "the checkpoint has no vision tower (visual.* weights)"
790
+ )
791
  elif not linearize_patch_embed(self.backbone):
792
  self.image_unavailable = "the vision patch embedding was not found"
793
  else:
794
  try:
795
  self.processor = load_processor(self.root)
796
+ except (
797
+ Exception
798
+ ) as exc: # noqa: BLE001 - text requests do not use the processor
799
  self.image_unavailable = (
800
  f"the image processor failed to load ({type(exc).__name__}: {exc})"
801
  )
802
+ self.fast_skipped = None if d3_fast else "d3_fast.py is not present"
803
+ self.fast = d3_fast.install(self) if d3_fast else None
804
  self.loaded_seconds = time.perf_counter() - started
805
 
806
  @classmethod
 
819
  max_length: int | None = None,
820
  readout_dtype: str | None = None,
821
  model_name: str | None = None,
822
+ permutation_average: bool = False,
823
  ) -> D3:
824
  """Load a package directory or Hub repository.
825
 
826
  ``verify``: ``fast`` (default) hashes every file of ``MODEL_MANIFEST.json`` up to 64 MiB and checks
827
  the size of the weight shards; ``full`` hashes every file; ``none`` skips the check. A checkpoint
828
+ without a manifest (a plain code-readout export) loads unverified. ``permutation_average``: see the
829
+ module docstring.
830
  """
831
  root = resolve_dir(
832
  name_or_path,
 
845
  max_length=max_length,
846
  readout_dtype=readout_dtype,
847
  model_name=model_name,
848
+ permutation_average=permutation_average,
849
  )
850
 
851
  # ------------------------------------------------------------------ requests
 
866
  if not images:
867
  return []
868
  if self.image_unavailable is not None:
869
+ raise ValueError(
870
+ f"image inputs are not available: {self.image_unavailable}"
871
+ )
872
  decoded = []
873
  for number, value in enumerate(images):
874
  try:
 
928
  }
929
  continue
930
  prepared.sequences[key] = sequence
931
+ if self.permutation_average:
932
+ self._prepare_reversed(state, prepared)
933
  return prepared
934
 
935
+ def _prepare_reversed(
936
+ self, state: Any, prepared: Prepared, visual: int = 0
937
+ ) -> None:
938
+ """Render and tokenize the reversed-option prompt of every runnable choice question.
939
+
940
+ ``visual``: the image tokens of the request (image path), counted like ``_prepare_images`` counts them.
941
+ """
942
+ n_images = len(prepared.images)
943
+ texts = []
944
+ for key in prepared.runnable:
945
+ flipped = reversed_question(prepared.questions[key])
946
+ if flipped is not None:
947
+ texts.append(
948
+ (
949
+ key,
950
+ (
951
+ self.image_text(state, flipped, n_images)
952
+ if n_images
953
+ else self.text(state, flipped)
954
+ ),
955
+ )
956
+ )
957
+ if not texts:
958
+ return
959
+ tokenizer = self.processor.tokenizer if n_images else self.tokenizer
960
+ ids = tokenizer([t for _, t in texts], add_special_tokens=False)["input_ids"]
961
+ for (key, text), sequence in zip(texts, ids):
962
+ length = len(sequence) - n_images + visual
963
+ if self.max_length is not None and length > self.max_length:
964
+ for planned in (prepared.sequences, prepared.texts, prepared.lengths):
965
+ planned.pop(key, None)
966
+ prepared.errors[key] = {
967
+ "type": prepared.questions[key].kind,
968
+ "error": "max_length_exceeded",
969
+ "message": f"the reversed-option prompt has {length} tokens, over the maximum context "
970
+ f"length of {self.max_length} tokens; nothing was truncated",
971
+ }
972
+ elif n_images:
973
+ prepared.reversed_texts[key] = text
974
+ prepared.reversed_lengths[key] = length
975
+ else:
976
+ prepared.reversed_sequences[key] = sequence
977
+
978
  def _prepare_images(
979
  self, state: Any, questions: Mapping[str, Any], images: list[Any]
980
  ) -> Prepared:
 
1026
  continue
1027
  prepared.texts[key] = text
1028
  prepared.lengths[key] = length
1029
+ if self.permutation_average:
1030
+ self._prepare_reversed(state, prepared, sum(visual))
1031
  return prepared
1032
 
1033
  def logits(self, sequences: Sequence[Sequence[int]], counts: Sequence[int]):
 
1060
  self, sequences: Sequence[Sequence[int]], counts: Sequence[int]
1061
  ) -> list[list[float]]:
1062
  """Softmax over each prompt's own codes, in option order."""
1063
+ if self.fast is not None:
1064
+ return self.fast.probabilities(sequences, counts)
1065
  probs = (
1066
  (self.logits(sequences, counts) / self.temperature)
1067
  .softmax(-1)
 
1094
  f"planned {width} input tokens, the processor produced {encoded['input_ids'].shape[1]}"
1095
  )
1096
  inputs = {name: value.to(self.device) for name, value in encoded.items()}
1097
+ fused = self.fast is not None and self.fast.fused is not None
1098
  with torch.inference_mode(), sdpa_backends(self.device):
1099
+ if fused and set(inputs) == set(IMAGE_INPUTS):
1100
+ padded = not bool(encoded["attention_mask"].all())
1101
+ hidden = self.fast.image_hidden(inputs, padded)[:, -1]
1102
+ else:
1103
+ hidden = self.backbone(**inputs, use_cache=False).last_hidden_state[
1104
+ :, -1
1105
+ ]
1106
  if self.readout_dtype == "float32":
1107
  logits = hidden.float() @ self.readout.T
1108
  else:
 
1131
  out: dict[str, list[float]] = {}
1132
  for start in range(0, len(keys), self.batch_size):
1133
  chunk = keys[start : start + self.batch_size]
1134
+ extra = [k for k in chunk if k in prepared.reversed_texts]
1135
+ texts = [prepared.texts[k] for k in chunk] + [
1136
+ prepared.reversed_texts[k] for k in extra
1137
+ ]
1138
+ counts = [len(prepared.questions[k].keys) for k in chunk + extra]
1139
+ widths = [prepared.lengths[k] for k in chunk] + [
1140
+ prepared.reversed_lengths[k] for k in extra
1141
+ ]
1142
+ n = len(chunk)
1143
+ if extra and len(widths) * max(widths) > MERGE_TOKENS:
1144
+ probs = self.image_probabilities(
1145
+ texts[:n], prepared.images, counts[:n], max(widths[:n])
1146
+ ) + self.image_probabilities(
1147
+ texts[n:], prepared.images, counts[n:], max(widths[n:])
1148
+ )
1149
+ else:
1150
+ probs = self.image_probabilities(
1151
+ texts, prepared.images, counts, max(widths)
1152
+ )
1153
  out.update(zip(chunk, probs))
1154
+ for key, backward in zip(extra, probs[n:]):
1155
+ out[key] = average_orders(out[key], backward)
1156
+ return out, sum(prepared.lengths[k] for k in keys) + sum(
1157
+ prepared.reversed_lengths.values()
1158
+ )
1159
 
1160
  def run(self, prepared: Prepared) -> tuple[dict[str, list[float]], int]:
1161
  """Probabilities per runnable question (request order, ``batch_size`` per pass) and the input tokens."""
 
1165
  out: dict[str, list[float]] = {}
1166
  for start in range(0, len(keys), self.batch_size):
1167
  chunk = keys[start : start + self.batch_size]
1168
+ extra = [k for k in chunk if k in prepared.reversed_sequences]
1169
+ sequences = [prepared.sequences[k] for k in chunk] + [
1170
+ prepared.reversed_sequences[k] for k in extra
1171
+ ]
1172
+ counts = [len(prepared.questions[k].keys) for k in chunk + extra]
1173
+ n = len(chunk)
1174
+ if extra and len(sequences) * max(map(len, sequences)) > MERGE_TOKENS:
1175
+ probs = self.probabilities(
1176
+ sequences[:n], counts[:n]
1177
+ ) + self.probabilities(sequences[n:], counts[n:])
1178
+ else:
1179
+ probs = self.probabilities(sequences, counts)
1180
  out.update(zip(chunk, probs))
1181
+ for key, backward in zip(extra, probs[n:]):
1182
+ out[key] = average_orders(out[key], backward)
1183
+ return out, sum(len(prepared.sequences[k]) for k in keys) + sum(
1184
+ len(s) for s in prepared.reversed_sequences.values()
1185
+ )
1186
 
1187
  def respond(
1188
  self,
 
1233
  lengths: Sequence[int] = (37, 64, 320, 333, 1000, 1024),
1234
  images: bool = True,
1235
  ) -> float:
1236
+ """Compile and autotune the kernels for every batch size up to ``batch_size`` (twice that with
1237
+ ``permutation_average``) before serving.
1238
 
1239
  The Gated DeltaNet kernels take the batch size as a compile-time constant, and Triton specializes their
1240
  length and chunk-count arguments on being 1 or a multiple of 16; these lengths cover every combination
 
1243
  four 1.6 MP images) also warm the vision tower. Answers are unchanged.
1244
  """
1245
  started = time.perf_counter()
1246
+ if self.fast is not None:
1247
+ self.fast.capture_all()
1248
+ widest = self.batch_size * (2 if self.permutation_average else 1)
1249
+ for size in range(1, widest + 1):
1250
  for length in lengths:
1251
  sequences = [
1252
  [
 
1279
  self.backbone.to(target)
1280
  self.readout = self.readout.to(target)
1281
  self.device = target
1282
+ if self.fast is not None:
1283
+ self.fast, self.fast_skipped = None, "moved after loading"
1284
  return self
1285
 
1286
  def parameter_count(self) -> int:
 
1298
  if (self.root / name).is_file()
1299
  }
1300
  identity = (self.manifest or {}).get("identity", {})
1301
+ record = {
1302
  "kind": "d3-code-readout",
1303
  "runtime": RUNTIME,
1304
  "model_name": self.model_name,
 
1321
  "refused); no option filtering; one fixed prompt for every request.",
1322
  "images": self.image_contract(),
1323
  }
1324
+ if self.permutation_average:
1325
+ record["permutation_average"] = True
1326
+ record["policy"] += (
1327
+ " Choice questions with two or more options are also scored with their options in reversed "
1328
+ "order, in the same forward passes, and answered with the per-option mean of both distributions."
1329
+ )
1330
+ return record
1331
 
1332
  def image_contract(self) -> dict[str, Any]:
1333
  """How image inputs are read (or why they are not available)."""
 
1364
  }
1365
  if self.device.type == "cuda":
1366
  info["gpu"] = torch.cuda.get_device_name(self.device)
1367
+ info["fast_path"] = self.fast_report()
1368
  return info
1369
+
1370
+ def fast_report(self) -> dict[str, Any]:
1371
+ """What the ROCm fast path (``d3_fast.py``) does in this process, or why it is off."""
1372
+ if self.fast is None:
1373
+ return {"active": False, "reason": self.fast_skipped}
1374
+ return {"active": True, **self.fast.report()}