Skip to content

Fuse Krea 2 DiT blocks and speed up large-row W4A4 GEMMs (0.10.0) - #15

Merged
iamwavecut merged 1 commit into
mainfrom
perf/dit-fused-runtime
Sep 28, 2026
Merged

iamwavecut merged 1 commit into
mainfrom
perf/dit-fused-runtime

Conversation

@iamwavecut

Copy link
Copy Markdown
Owner

What changes

All W4A4 models on CUDA (no API changes):

  • Triton kernels take row counts as runtime arguments instead of tl.constexpr: a new prompt length or image size no longer recompiles the activation, pack and scale kernels (first denoise step 9.6 → 2.2 s, Qwen3-VL text encoder 3.7 → 0.19 s on Krea 2).
  • Large-row packed W4A4 matmuls use matmul_int8_scaled_with_triton, an INT8 GEMM that applies the scale epilogue in registers. Results are bit-identical to the torch._int_mm + scale path and 20–33% faster per layer.
  • Projections that consume the same activation (Q/K/V, SwiGLU gate/up) reuse one RPBH quantization under torch.no_grad(); INT8 surrogate codes are produced directly.
  • Autotune results persist in the Triton cache (cache_results), and the in-place residual output is restored between benchmark runs (without that the first call of a new shape bucket re-added the residual on every benchmark run).

New orbitquant.runtime.krea2 (Krea 2 Turbo): grouped Q|K|V|gate and gate/up INT8 GEMMs with sigmoid, SwiGLU and gated-residual epilogues, RMSNorm/modulation prologues fused into activation quantization, a Q/K RMSNorm + RoPE kernel, W8A8 down projections, an INT8 Q·Kᵀ / FP16 P·V attention kernel (SageAttention v1 scheme; blocks whose Q/K RMSNorm scales one channel far above the rest stay on BF16 attention), and save_fused / install(fused_path=...) so prebuilt fused weights are memory-mapped instead of copied into host memory.

Measurements (RTX 4060 Ti 16 GB, torch 2.10 + cu128, Triton 3.6)

  • Krea 2 Turbo W4A4, 1024×1024, 8 steps: 25.9 s → 10.8 s per image (1.07 s/step); DiT 10.0 → 7.3 GB.
  • Qwen-Image 2.1 W4A4 pipeline: 15.0 → 10.3 s per image, PNGs pixel-identical on all compared prompts.
  • INT8 attention: 1.9× Flash at 4k–8k tokens; relative error 0.35–1.2% on attention inputs captured from a real denoise.

Checks

  • uv run ruff check ., uv run pytest -ra, scripts/run_paper_methodology_checks.sh, scripts/run_hf_compat_checks.sh --mode all: pass.
  • scripts/run_cuda_kernel_checks.sh on an RTX 4060 Ti: pass. CUDA tests for the new kernels (tests/test_fused_dit_kernels.py, 14 tests) plus test_kernels.py/test_orbit_linear.py: 165 passed.

General CUDA changes for every W4A4 model:
- Triton kernels take row counts as runtime arguments instead of constexpr,
  so new prompt lengths and image sizes no longer trigger recompilation.
- Large-row packed W4A4 matmuls run an INT8 GEMM with the scale epilogue in
  registers (triton_int8_gemm); results are bit-identical to the
  torch._int_mm + scale path and 20-33% faster per layer.
- Projections that consume the same activation reuse one RPBH quantization
  (Q/K/V, SwiGLU gate/up) under no_grad, and INT8 surrogate codes are
  produced directly instead of via packed W4 codes.
- Autotuning results persist in the Triton cache, and in-place outputs are
  restored between benchmark runs.

New orbitquant.runtime.krea2: grouped Q|K|V|gate and gate/up GEMMs with
sigmoid, SwiGLU and gated-residual epilogues, RMSNorm/modulation prologues in
the activation quantizer, a fused Q/K RMSNorm + RoPE kernel, W8A8 down
projections, an INT8 Q.K^T / FP16 P.V attention kernel (triton_attention),
and save_fused/install(fused_path=...) to map prebuilt fused weights.
@iamwavecut
iamwavecut merged commit 3710826 into main Sep 28, 2026
3 checks passed
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