TileMix-comfyui
Bidirectional tile-centric mixed-precision (FP16/INT8) attention for ComfyUI diffusion transformers, with a Krea 2 gated-MMDiT patch.
TileMix-comfyui
Bidirectional tile-centric mixed-precision attention for ComfyUI diffusion transformers. Each tile of the attention score matrix is routed to an FP16 or an INT8 compute path by a static map, and both paths feed one online-softmax state inside a single fused Triton kernel. Text tokens and a configurable diagonal band stay FP16; the rest of the image-image attention runs in INT8. No retraining, no calibration.
Ships a Krea 2 gated-MMDiT patch node. The Triton kernel is bundled — no external attention backend is required.
The tile-centric mixed-precision idea and the routing pattern names come from TileMix (paper, Tile-Centric Mixed-Precision Attention for LLM Inference Acceleration). The upstream release targets causal LLM decoding; this repository is an independent bidirectional kernel for diffusion transformers. See
NOTICE.
Why
Krea 2's SingleStreamDiT runs 28 blocks of dense GQA attention (48 query /
12 KV heads, head_dim 128) over a joint [text, image] sequence that grows
quadratically with resolution. At 1–2 MP the attention is the dominant cost.
TileMix keeps the perceptually important tiles exact and spends INT8 tensor
cores on the rest:
- Exact text. Every tile that reads a text key block or produces a text query row stays FP16. The text conditioning is never quantised.
- FP16 diagonal band. A configurable band of key blocks around each query row stays FP16 (local structure matters most); off-diagonal image tiles go INT8.
- GQA inside the kernel. Query heads are mapped to KV heads in-kernel, so
Krea's
repeat_interleaveof K/V to 48 heads is skipped. - BF16-native. Krea 2 is bf16. bf16→fp16 is lossless in precision (8 vs 11 mantissa bits); score/PV use fp16 tensor cores while the running softmax state stays fp32, so long sequences can't overflow.
- Honest fallback. Anything the kernel can't serve (fp32, an attention
mask, sequences under
min_tokens, a non-Krea model) runs Krea's original attention, logged once per cause.
Install
cd ComfyUI/custom_nodes
git clone https://github.com/chinoll/TileMix-comfyui.git
# Triton + a CUDA PyTorch are already in ComfyUI; nothing else to install.
Restart ComfyUI. Requires an NVIDIA GPU with a working Triton build (developed on SM120 — RTX 5090 / 5060 Ti; the tensor-core paths need SM80+). bf16 or fp16.
Node
Krea 2 TileMix Attention (FP16/INT8 tiles) — category
model_patches/attention. Insert it between the Krea 2 UNETLoader and the
sampler:
UNETLoader → Krea 2 TileMix Attention → KSampler → VAEDecode
It replaces only attn(q, k, v) in the 28 gated blocks. Q/K/V projections, QK
RMSNorm, RoPE, the sigmoid(gate), and the output projection stay on Krea's
stock path.
| input | default | meaning |
|---|---|---|
| pattern | band | Which image-image tiles are INT8. band keeps a diagonal FP16 window; one routes every image tile to INT8; zero is an FP16-only sanity mode (global, row_rand, bigbird also available). |
| percentage | 0.5 | Fraction of image-image tiles routed to INT8 (text tiles are always FP16). |
| min_tokens | 1024 | Below this joint token count, use Krea's original attention. |
| smooth_k | True | Mean-centre K per head before INT8 quantisation (exact for softmax; cuts INT8 error on keys with a shared channel offset ~4×). |
| qk_fp32_accumulate | False | Accumulate FP16-tile scores in fp32 — SDPA-level accuracy for very sharp attention, ~25% slower FP16 tiles. |
| strict | False | Raise kernel errors instead of falling back. |
Guidance: band at percentage 0.5–0.75 is the quality/speed sweet spot;
raise percentage toward one for maximum speed, drop it toward zero (or
enable qk_fp32_accumulate) if you see quality loss.
Results
RTX 5090, torch 2.13 / Triton 3.7, bf16, Krea 2 attention shape (48 query /
12 KV heads, head_dim 128), batch 1. Error is relative L2 vs fp32 PyTorch
SDPA; timing is warm-up + 30 triton.testing.do_bench reps, swept over 19
sequence lengths from 1024 to 16896 (including odd lengths and lengths one
token above/below tile boundaries: 2051, 3333, 4095/4096/4097, 8191/8192/8193,
16383/16384/16385).
| method | INT8 image tiles | speedup vs bf16 SDPA (min / median / max) | rel. L2 error |
|---|---|---|---|
| PyTorch SDPA (ComfyUI default) | – | 1.00 | 1.7e-3 |
| Krea 2 stock block (repeat_interleave + SDPA) | – | 0.85 / 0.96 / 0.99 | 1.7e-3 |
| TileMix zero (all FP16 tiles, GQA in-kernel) | 0% | 1.13 / 1.37 / 1.45 | 3.3e-3 – 4.1e-3 |
| TileMix zero + qk_fp32_accumulate | 0% | 0.95 / 1.12 / 1.17 | 1.8e-3 |
| TileMix band @ 0.5 | 50% | 1.14 / 1.49 / 1.56 | 0.8e-2 – 1.7e-2 |
| TileMix band @ 0.75 | 75% | 1.15 / 1.56 / 1.65 | 1.0e-2 – 2.0e-2 |
| TileMix one (all image tiles INT8) | 100% | 1.16 / 1.61 / 1.73 | 1.1e-2 – 2.3e-2 |
| SageAttention 1.0.6 (Triton INT8) | 100% | 1.16 / 1.70 / 1.82 | 3.0e-2 – 3.3e-2 |
| SageAttention 3 (Blackwell FP4) | – | 1.4 / 2.6 / 3.6 | 0.32 – 0.36 |
Every TileMix layout beats SDPA at all lengths (minimums are at 1024 tokens);
from 2048 tokens on, every layout is ≥ 1.30× SDPA. Because Krea's real path is
repeat_interleave + SDPA (0.96× SDPA), the speedup over what Krea actually
runs is ≈ 1.4–1.7×. Other head layouts (24:24, 32:8) and head_dim 64 follow the
same pattern.
Error notes: the table uses adversarial random N(0, 1) inputs (near-uniform
attention, the worst case for INT8 — every key matters). At the same INT8
budget TileMix is ~2–3× more accurate than SageAttention 1 here because text
tiles and the diagonal band stay FP16. K-smoothing cuts the error on keys with
a shared channel offset from 4.1e-2 to 9.5e-3. The FP16-tile path alone matches
SDPA (≤ 2e-3), and qk_fp32_accumulate holds it at SDPA level even for very
sharp attention.
How it works
- Routing as run pairs. With no causal triangle the routing is fully known
on the host, so each tile row is compressed into
(fp16_start, fp16_end, int8_start, int8_end)runs. The kernel walks the runs with directly-indexed, homogeneous inner loops, so Triton can software-pipeline both precision paths — a data-dependent branch or indirect index in the hot loop costs ~25% on SM120. There is no 64-column bitmask limit, so 16k+ token sequences route at full 128×64 tile granularity. - Numerics. QK^T for FP16 tiles uses fp16-input / fp16-accumulate tensor
cores (2× the fp32-accumulate rate on consumer Ada/Blackwell) with Q
pre-scaled by
sm_scale·log2(e). P·V accumulates in fp16 per tile, but the running accumulator, row max, and row sum are fp32. K and V are materialised once as fp16 (cheap for 12 KV heads) so the hot loop does no conversions. - INT8 tiles. Per-token symmetric quantisation of Q and K; K is
mean-centred per (batch, head) before quantisation, and
q·mean(k)is added back inside the kernel — softmax is invariant to a per-query constant, so the FP16 path is untouched.
Standalone API
The kernel is usable outside ComfyUI for any bidirectional GQA attention:
import torch
from tilemix import (
TileMixBidirRouting,
tilemix_bidirectional_attention,
make_krea2_tile_map,
)
# q: [B, HQ, T, D], k/v: [B, HK, T, D] (bf16/fp16, D in {32, 64, 128})
route = make_krea2_tile_map(T, text_tokens, block_m=128, block_n=64,
pattern_type="band", percentage=0.5, num_head=HK)
routing = TileMixBidirRouting(route, block_m=128, block_n=64) # build once per (T, text)
out = tilemix_bidirectional_attention(q, k, v, routing) # [B, HQ, T, D]
make_bidirectional_tile_map(seq_len, block_m, block_n, pattern_type, percentage, exact_prefix=...) builds maps for the general (non-Krea) case.