Zero-Shot Classification
Transformers
Safetensors
qwen3_5
feature-extraction
decision-model
classification
system-one
multimodal
vision
custom_code
Instructions to use vllm-sr/d3 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use vllm-sr/d3 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("zero-shot-classification", model="vllm-sr/d3", trust_remote_code=True)# pip install -U transformers accelerate # Load model directly from transformers import AutoProcessor, AutoModel processor = AutoProcessor.from_pretrained("vllm-sr/d3", trust_remote_code=True) model = AutoModel.from_pretrained("vllm-sr/d3", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
d3 v3.0.4
Browse files- MODEL_MANIFEST.json +11 -5
- README.md +2 -2
- d3_fast.py +512 -0
- d3_kernels.py +523 -0
- d3_runtime.py +239 -35
MODEL_MANIFEST.json
CHANGED
|
@@ -564,7 +564,9 @@
|
|
| 564 |
},
|
| 565 |
"runtime": {
|
| 566 |
"files_sha256": {
|
| 567 |
-
"d3_runtime.py": "
|
|
|
|
|
|
|
| 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-
|
| 590 |
"files_sha256": {
|
| 591 |
"LICENSE": "c71d239df91726fc519c6eb72d318ec65820627232b2f796219e87dcf35d0ab4",
|
| 592 |
"NOTICE": "61f4ea488e1fed0c3fce97171506a23152c2591116442f78f15827fb9e6af469",
|
| 593 |
-
"README.md": "
|
| 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 |
-
"
|
|
|
|
| 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 |
-
"
|
|
|
|
| 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.
|
| 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
|
| 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``
|
|
|
|
| 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``
|
| 31 |
-
whole prompt while the Gated DeltaNet layers stay causal. The Gated
|
| 32 |
-
transformers binds at import: flash-linear-attention (and causal-conv1d)
|
| 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
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 (
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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 (
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 =
|
| 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 =
|
|
|
|
|
|
|
| 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
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 980 |
-
|
| 981 |
-
prepared.
|
| 982 |
-
|
| 983 |
-
|
| 984 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 985 |
out.update(zip(chunk, probs))
|
| 986 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 997 |
-
|
| 998 |
-
|
| 999 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1000 |
out.update(zip(chunk, probs))
|
| 1001 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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``
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 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()}
|