Inspiration
Transformer architectures are heavily bottlenecked by memory bandwidth, kernel launch overhead, and tensor core underutilization. While fully-fused kernels (like FlashAttention) address memory issues, they frequently suffer from numerical instability (precision drift) when errors accumulate across deep residual blocks in mixed-precision or FP32 workloads. Furthermore, an "always-on" fusion approach often degrades performance on smaller shapes due to launch overhead. We were inspired to learn this interesting topic and build a solution that doesn't just fuse operations blindly, but adapts to the hardware and shape dynamically to guarantee absolute bit-level accuracy without sacrificing extreme speed.
What it does
Our solution is a dynamically dispatched, shape-aware UserOptimizedTransformer adapter backed by custom Triton and Gluon kernels. Instead of a monolithic kernel, it evaluates the sequence length, dimension, and precision type, and routes the workload to the most optimal execution path. It achieves massive latency reductions (up to 11.97x) and safely processes extreme sequence lengths (up to 100,000 tokens) that typically crash standard baselines with Out-Of-Memory (OOM) errors.
How we built it
We built the suite using PyTorch 2.13.0, Triton 3.7.1, and CUDA 13.0, developing heavily in VSCode with assistance from OpenAI Codex/GPT-4o via a Test-Driven Development workflow. Key engineering highlights include:
- Blackwell-Optimized Gluon Core (cc >= 12 / sm_120): For short sequences, we implemented a true full-row fusion kernel targeting NVIDIA's Blackwell architecture via mma_v2. It keeps QK, masked softmax, and PV entirely in registers/shared memory, bypassing global memory writes.
- Warp and Lane Exploitation: We hard-coded specific sequence dimension cases to perfectly allocate parallel threads (warps and lanes). By matching exact tile sizes to specific warp counts, we maximize GPU occupancy.
- Fast Inverse Square Root: We hyper-optimized our custom Triton LayerNorm using _tl_software_rsqrt, an FP32 bit-level hack utilizing a Quake-style magic exponent seed (0x5F3759DF) followed by Newton-Raphson steps.
- Thermal-Drift-Resistant Benchmarking: We engineered a custom 5-factor modular ablation runner (ablation.py) to systematically measure 448 technique combinations. The suite uses an alternating measurement order to cancel out GPU clock ramping and thermal throttling drift during latency comparisons. ## Challenges we ran into
- Numerical Instability: We faced significant precision drift when accumulating errors across deep residual blocks. We solved this by inventing a Hybrid Exact/Fused Dispatching policy (e.g., EEFF for Case 6 and EFFF for Cases 1-5). By executing a mix of exact native PyTorch blocks and fused Triton blocks, we bounded error accumulation and passed strict accuracy tolerances (atol=0.002, rtol=0.02).
- FP16 / FP32 Precision Inconsistencies: We discovered structural inconsistencies between data types in our kernels. We resolved this by updating our Gluon MMA adapter to dynamically derive operand k_width from input primitive bit widths, ensuring bit-level exactness.
- Hardware Limits: We initially explored TMEM tcgen05 instructions for Blackwell but ran into LLVM lowering issues on the target environment, requiring us to pivot to the highly stable mma_v2 path. ## Accomplishments that we're proud of
100% Accuracy Pass Rate: Achieved across all viable official test shapes.
11.97x Peak Latency Reduction: Our Tiled FP32 TF32 MMA implementation on Case 13 dropped median latency from 99.35ms down to 8.30ms (an 11.971x speedup).
Extreme Long-Sequence Execution: Where standard dense baselines crash with OOM errors, our bounded-memory kernel processed 100,000 tokens in FP16 (21.4 seconds) and Bfloat16 (25.6 seconds), tightly bounding peak memory to just 14.12 GiB on our 16GB GPU.
2.02x Average Ablation Speedup: In our 448-measurement full-factorial ablation sweep, our 11000 (QKV + SDPA) configuration achieved the best overall mean speedup of 2.02x across the 13 viable benchmark shapes.
What we learned
We learned how to implement a GPU Kernel for a transformer layer and that "always fuse everything" is a flawed approach. Launch and layout overhead can easily outweigh the saved work for small or unfavorable GEMM sizes. We also learned how critical reduction-order-sensitive operations are to precision. Matching PyTorch's persistent softmax lane association was key to preventing one-ULP FP16 differences from amplifying across the full Transformer stack.
What's next for Implement a GPU Kernel for a Transformer Layer
Our profiling shows that attention is no longer the dominant cost in our optimized workloads. The next optimization targets are repeated native masked/pointwise kernels, followed by projection/FFN GEMMs and LayerNorms. Additionally, we plan to fully port our kernels to TMEM tcgen05 instructions once the backend LLVM lowering issues are resolved for the Blackwell architecture.

Log in or sign up for Devpost to join the conversation.