Inspiration
Transformer inference is dominated by dozens of small, sequential PyTorch dispatches — QKV projections, attention, softmax, output projection, LayerNorm — each launched as its own CUDA kernel with its own overhead. The TikTok TechJam GPU Kernel Challenge asked participants to take a fixed transformer layer and reimplement it as custom GPU kernels that stay numerically close to the PyTorch reference (rel error < 0.02, abs error < 0.002) while running dramatically faster, across shapes ranging from tiny batches to a sequence length of 100,000 tokens that doesn't even fit in memory in its naive form.
What it does
This project — Blackwell Transformer Kernel (sm_120) — is a fused Triton implementation of the challenge's transformer layer, tuned and measured on an NVIDIA RTX 5060 Ti (Blackwell, sm_120, 16 GB). Four Triton kernels replace the reference's ~29 PyTorch dispatches per layer, the forward pass runs as a single CUDA-graph replay wherever that's a net win, and for the one shape whose activations physically exceed VRAM, inputs are streamed in batch chunks sized dynamically from the device's actual free memory.
Headline results:
- Median 9.19x speedup over the 13 test shapes with a runnable PyTorch reference (geometric mean 8.67x, arithmetic mean 11.23x)
- All 14 shapes execute correctly with zero failing elements, worst-case
max_abserror of 1.35e-03 against a 2e-3 gate - Best single-shape speedup: 33.34x (B64 d128 H4 S1024); smallest: 2.87x
- The 14th shape — B=32, d=1024, H=16, S=100,000 — cannot run on any hardware under the naive reference implementation, since materializing the attention score matrix would require 19,073 GB. This implementation streams it to completion in 37.67 s at 36.9 TFLOP/s (73% of the card's measured fp16 peak), validated bit-exact against a memory-safe custom oracle, with per-sequence cost holding steady (1.14–1.18 s) across a 16x batch range.
How we built it
The kernel bodies (QKV+attention fusion, LayerNorm-in-epilogue fusion, fp32-store GEMM, CUDA-graph capture) originate from the upstream COBRA GPU kernel project, tuned on an RTX 3050 (sm_86) under Windows. This branch's contribution — roughly 1,900 lines — re-derives every hardware-specific tuning constant against direct measurement on Blackwell sm_120 rather than assuming portability, and adds capability the original didn't have:
- Shape 14 (the impossible one): host-streamed batch chunking sized from
torch.cuda.mem_get_info(), with halve-on-OOM fallback, plus a memory-safe reference oracle (tools/reference_chunked.py) proventorch.equal-identical to the stock baseline across 16 configurations. - Precision policy fix: discovered that a
float32floor belowd_model=64was silently disabling the Triton GEMM and falling back to cuBLAS TF32 (same 10 mantissa bits as fp16, without the fp32-store accuracy correction). Removing it yielded a 1.34x speedup and better accuracy, confirmed over a 360-trial campaign. - Fusion gate correction: the QKV+attention fusion requires
BM >= S(tile height tied to sequence length). Added a lowerseq_lenbound after measuring the fused path running 1.13x slower at S=32. - Attention tile re-derivation: re-tuned against the card's 36 SMs —
head_dim=64tiling gave a 1.22x gain, and thehead_dimgate was raised from 128 to 256 after measuring the Triton kernel beating FlashAttention-2 by 1.11x in that regime. - Measurement infrastructure: 12 standalone harnesses under
tools/— an analytic shared-memory/occupancy pruner for tile search, in-graph timing, SDPA backend comparison, CUDA-graph gate A/B testing, precision campaigns, and per-kernel timing breakdowns.
Every tuning decision above was reached empirically, not assumed — including nine rejected approaches documented in report/measurements/.
Development tools, APIs, libraries & frameworks
- Language/runtime: Python 3.14.4
- Core libraries: PyTorch 2.11.0 (cu128), Triton 3.6.0 — every kernel is JIT-compiled at first use, no separate compiler toolchain required
- GPU stack: CUDA driver 595.84, targeting sm_120 (Blackwell)
- No external APIs — this is a pure GPU-kernel / systems-performance project; there's no model inference API or web service involved, just PyTorch, Triton, and the hardware
- Custom tooling: an internal suite of 12 Python measurement harnesses (
tools/sweep.py,tools/validate.py,tools/reference_chunked.py,tools/graph_gate.py,tools/sdpa_backend.py,tools/diag_small.py,tools/breakdown.py,tools/time_case14.py, etc.) built specifically for this project to drive tile search, correctness validation, and precision/timing campaigns
Datasets and assets used
No external dataset was needed — the "data" here is the challenge's own synthetic input-shape matrix (14 combinations of batch size, hidden dim, head count, and sequence length, up to B=32, d=1024, H=16, S=100,000), generated at test time, plus a self-built, torch.equal-verified memory-safe reference oracle used to validate correctness on the one shape too large for the stock PyTorch baseline to run at all.
Challenges we ran into
- The reference itself is an unreliable ground truth for speed: its own timings swing up to 2.63x run-to-run on several shapes, while this implementation's timings hold within 1.4%. Every reported speedup is corroborated by at least two independent sweeps to compensate.
- No usable kernel-level profiler was available: CUPTI failed with
CUPTI_ERROR_INVALID_DEVICEand Nsight required root privileges unavailable in the environment, so occupancy had to be modelled via an analytic pruner rather than measured directly — every model-ranked configuration was then re-confirmed with direct timing before being adopted. - Shape 14 is a genuine hardware wall, not just a slow case: the naive reference would need to materialize a 19,073 GB attention score tensor. Solving it required rethinking the problem as a streaming/chunking problem rather than a kernel-fusion problem, sized dynamically off real-time free VRAM.
- A silent precision bug: the float32 fallback path for small
d_modelwas quietly costing both speed and accuracy — an easy thing to miss without a systematic 360-trial precision campaign.
Accomplishments that we're proud of
- A median 9.19x speedup with zero failing elements across all 14 shapes
- Making an "unrunnable" shape (S=100,000) not just execute, but execute at 73% of the card's measured fp16 peak
- Every hardware-specific constant re-derived from first-principles measurement on new hardware (sm_120) rather than inherited assumptions from the sm_86 original
- Transparent reporting of limitations and nine explicitly rejected optimization approaches, rather than only showing what worked
What we learned
Porting a kernel between GPU architectures isn't just "does it run" — every tile size, fusion gate, and precision fallback that was correct on sm_86 needed to be re-measured, and two of them were actively wrong on sm_120. Median-and-mean speedup numbers can hide unstable ground truth; when the reference implementation itself has multi-x run-to-run variance, correctness and reproducibility require independent oracles and repeated sweeps, not a single benchmark run.
What's next for the project
- Reclaim the up-to-2.39x of unclaimed bandwidth headroom on bandwidth-bound shapes once a precision approach is found that doesn't erode the accuracy margin under input-scale stress
- Get real kernel-level profiling working (root/Nsight access) to replace the analytic occupancy model with direct measurement
- Re-measure Case 6, which currently carries a pre-change timestamp due to VRAM contention from an unrelated process
Log in or sign up for Devpost to join the conversation.