Extensions/TileMix-comfyui
ComfyUI Extension

TileMix-comfyui

Bidirectional tile-centric mixed-precision (FP16/INT8) attention for ComfyUI diffusion transformers, with a Krea 2 gated-MMDiT patch.

By chinoll·Created about a month ago·Updated about a month ago· 1
chinoll/TileMix-comfyui
Nodes—
On cloudLocal install
Stars1
Updatedabout a month ago
Readme

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.

License GPU

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_interleave of 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.

License

Apache-2.0. See LICENSE and NOTICE.