Repository navigation
Fuse Krea 2 DiT blocks and speed up large-row W4A4 GEMMs (0.10.0) - #15
Merged
Merged
Conversation
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.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
What changes
All W4A4 models on CUDA (no API changes):
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).matmul_int8_scaled_with_triton, an INT8 GEMM that applies the scale epilogue in registers. Results are bit-identical to thetorch._int_mm+ scale path and 20–33% faster per layer.torch.no_grad(); INT8 surrogate codes are produced directly.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), andsave_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)
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.shon an RTX 4060 Ti: pass. CUDA tests for the new kernels (tests/test_fused_dit_kernels.py, 14 tests) plustest_kernels.py/test_orbit_linear.py: 165 passed.