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_abs error 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) proven torch.equal-identical to the stock baseline across 16 configurations.
  • Precision policy fix: discovered that a float32 floor below d_model=64 was 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 lower seq_len bound 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=64 tiling gave a 1.22x gain, and the head_dim gate 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_DEVICE and 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_model was 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

Built With

Share this project:

Updates

Submission history