Skip to content

Fuse the transformer blocks of six diffusion families into stored checkpoints (0.11.0) - #16

Merged
iamwavecut merged 4 commits into
mainfrom
perf/fused-transformers
Oct 2, 2026
Merged

iamwavecut merged 4 commits into
mainfrom
perf/fused-transformers

Conversation

@iamwavecut

@iamwavecut iamwavecut commented Oct 2, 2026 •

Copy link
Copy Markdown
Owner

What changes

New orbitquant.fused: fused transformer blocks stored as checkpoints. Per block, the projections that read one input become one INT8-surrogate GEMM (attention Q|K|V, SwiGLU gate|up, both streams of a joint attention); RMSNorm/LayerNorm and AdaLN modulation (per tensor, or per row from a modulation table) run in the activation-quantization prologue; SwiGLU, sigmoid gates, GELU-tanh and gated residual updates (per column or per row from a table) run in the GEMM epilogue; Q/K normalization and RoPE (interleaved or rotate-half, optionally on a leading slice of the head) are one kernel; attention uses the INT8 Q·Kᵀ / FP16 P·V kernel where a block's Q/K RMSNorm weights stay within a spread limit and BF16 SDPA elsewhere.

Families: Krea 2, FLUX.2 (double/single stream, reference-image KV cache), Ideogram 4, Qwen-Image 2.1 (block-causal prefill, KV-cached decode), MiniMax-H3 (per-row AdaLN tables, padding documents) and Boogu-Image (refiner, single/double stream).

  • fuse(model) converts a loaded per-projection model in place; save_pretrained writes a checkpoint holding only the fused groups and records the layout in quantization_config.fused_layout. from_pretrained and load_orbitquant_artifact rebuild the empty groups before the weights load and finish the blocks after (prepare_skeleton / finalize). fuse_component_artifact converts an orbitquant-v1 component artifact.
  • Groups can mix weight widths (low-bit boundary/interior protection) as row segments of one GEMM with per-row scale factors, and can take 8-bit activations: per-token absmax INT8 of the RPBH-rotated input, no codebook (activation_bits in the layout). Boogu-Image uses it for every group: its 3360 channels only allow a 32-wide rotation block, and 4-bit Q/K/V inputs produce visible texture artifacts.
  • The activation quantizer handles RPBH blocks up to 16384 (FLUX.2 single-stream output projection) by rotating 4096-wide chunks and pairing whole chunks for the last butterfly stages.
  • The Triton prologue/post-norm kernels compile without FP contraction so the bf16 round trips that mirror eager rounding are kept.
  • INT8 attention converts V to FP16 once before the kernel instead of per tile (about 10% faster at MiniMax-H3 sizes) and accepts a different key/value length than the query length.
  • Importing orbitquant.fused does not import Triton; CPU-only code can read layouts.
  • Fused weights are frozen parameters rather than buffers: Diffusers' streamed group offloading (use_stream=True) moves only parameters back to the host, so buffers piled up on the GPU until it ran out of memory.
  • Loading keeps the stored dtypes of fused tensors (FP32 row scales, BF16 dense rows). Diffusers casts floating checkpoint tensors to the model dtype before handing them to the quantizer, which rounded the INT8 group scales and made a reloaded checkpoint differ from the in-memory fused model.
  • Krea 2: the text-fusion stack runs once per prompt instead of once per step (cached on the identity and version of its inputs), and the blocks skip padded prompt rows.
  • fuse_component_artifact accepts artifacts whose manifest lists files that a later repository edit removed; the fused manifest lists the files the fused copy ships.

orbitquant.runtime.krea2 from 0.10 is unchanged.

Measurements

RTX 4060 Ti 16 GB (torch 2.10 + cu128, Triton 3.6, diffusers 80c7ed26 unless noted), hot, per-projection W4A4 → fused W4A4 of the same published checkpoint:

Model Setting Transformer time Peak
FLUX.2-klein-9B 1024², 4 steps 5.44 → 3.15 s 12.93 → 12.60 GiB
FLUX.2-klein-9B reference edit, 1024², 4 steps (image wall) 13.4 → 8.3 s
Ideogram 4 Instant 1024², 8 steps (diffusers 0.39.0) 27.15 → 19.25 s 13.96 → 14.12 GiB
Boogu-Image Turbo 1024², 4 steps vs the SDNQ UINT4 deployment 8.75 → 5.8 s 6.7 → 6.0 GB NVML
Boogu-Image Turbo two reference images, 1024² 35.8 → 16.6 s 9.0 → 7.4 GB NVML
Krea 2 Turbo 1024², 8 steps, model CPU offload, per step 2.58 → 1.11 s 10.7 → 8.1 GiB
Qwen-Image 2.1 (Turbo-Image) 1024², steps 2–6, fp16, model CPU offload 6.24 → 3.87 s 8.19 → 7.91 GiB
Qwen-Image 2.1 (Turbo-Image) edit with one reference, steps 2–6 7.39 → 4.56 s 10.87 → 10.59 GiB

RTX 3090, MiniMax-H3 608×480×124 frames, 24 sigma points / 23 forwards, release runner: 6.06 → 4.0 s per forward resident; generation 175 → 116 s resident, and with block-level streamed group offloading (no record_stream) 174 → 114 s at 5.1 GiB of GPU memory.

Replays of real block inputs (fused vs per-projection, rel. L2): FLUX.2 0.002–0.027, Ideogram 4 0.001–0.017. A fused checkpoint loaded with from_pretrained renders bit-identical images to the in-memory fused model (FLUX.2 10/10, Krea 2 6/6, Qwen-Image 2.1 6/6). Same-seed images move with the different rounding but keep their detail and legibility (paired contact sheets); reference edits stay close (FLUX.2: PSNR 32.7 dB, SSIM 0.985).

Checks

  • uv run ruff check ., uv run pytest (CPU): pass.
  • scripts/run_paper_methodology_checks.sh, scripts/run_hf_compat_checks.sh --mode current|release|dev: pass.
  • CUDA tests on an RTX 4060 Ti and an RTX 3090: tests/test_fused_family_kernels.py, tests/test_fused_dit_kernels.py, tests/test_krea2_runtime.py, tests/test_fused_layout.py: pass.

…ckpoints (0.11.0)

orbitquant.fused groups the projections that read one input into one INT8-surrogate GEMM per
block (Q|K|V, SwiGLU gate|up, both streams of a joint attention), runs norm and AdaLN modulation
in the activation-quantization prologue and SwiGLU, sigmoid gates and gated residual updates in
the GEMM epilogue, fuses Q/K normalization with RoPE and uses the INT8 Q.K^T / FP16 P.V attention
kernel where a block's Q/K norm weights allow it. Families: Krea 2, FLUX.2, Ideogram 4,
Qwen-Image 2.1, MiniMax-H3 and Boogu-Image.

A fused checkpoint stores only the fused groups and records the layout in its
quantization_config; from_pretrained and load_orbitquant_artifact rebuild the groups before the
weights load. Groups may mix weight widths (low-bit boundary/interior protection) as row
segments of one GEMM and may take 8-bit activations (per-token absmax INT8 of the rotated
input). The activation quantizer covers RPBH blocks up to 16384.
…locks on valid rows

Fused groups keep their weights as frozen parameters instead of buffers: diffusers' streamed
group offloading moves only parameters back to the host, so block-level streaming kept every
fused block on the GPU and ran out of memory under a VRAM cap.

Krea 2 blocks run on the rows of valid tokens only (padded text rows are never attended to
and are dropped at the output), find those rows once per forward, and the text fusion stack
runs once per prompt instead of in every denoising step.
Diffusers casts every floating tensor of a checkpoint to the model dtype before the quantizer
places it, which rounded the FP32 per-row scales of INT8 groups (and would round BF16 dense
rows of an FP16 model). The quantizer now places fused group tensors from the stored values,
so a fused checkpoint loads bit-identical to the model it was saved from.
@iamwavecut
iamwavecut merged commit 74f0bc8 into main Oct 2, 2026
3 checks passed
@iamwavecut
iamwavecut deleted the perf/fused-transformers branch October 2, 2026 02:34
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant