KerneLab
Case studies

The work so far.

Three ports, described qualitatively. Each names its reference, its target and how it was verified. Figures will follow, published with reproductions.

01 · Kernel language to kernel language

A mixture-of-experts kernel, from CuTe DSL to Pallas Mosaic GPU

FlashInfer MegaMoE·CuTe DSL→Pallas Mosaic GPU·Blackwell B200
Port complete · tuning ongoing
Reference
The masked mixture-of-experts path in FlashInfer, written in NVIDIA's CuTe DSL.
Target
Pallas Mosaic GPU, JAX's own kernel layer for NVIDIA GPUs, so the kernel can live inside a JAX stack.
Number format
NVFP4 kept end to end: 4-bit values with shared block scales, consumed natively by Blackwell's tensor cores.
Correctness
The grouped matrix multiply is bit-exact against a float64 oracle. The full path agrees with the reference to bfloat16 precision.
Speed
Faster than the dense bfloat16 baseline the same computation would otherwise use.
Figures
to be published with reproductions

What was done

  • Four kernels rebuilt, plus a model of the numerics that every test is checked against.
  • Three of the reference's four stages fused into one kernel, removing a large intermediate tensor. This was one of the biggest end-to-end gains.
  • Six structural variants taken from the reference's own design were each implemented, verified bit-exact and measured.
  • The reference's launch configuration was reconstructed from its capture and executed, not inferred from its source.

What we found

  • Fusion pays end to end. The intermediate tensor between stages never shows up when kernels are timed one at a time.
  • The launch record matters. How the reference is scheduled across the processor is not visible in its instruction listing; capturing a real launch revealed it.
  • Producing the scales is half the work. A block-scaled format needs its scale factors computed and laid out the way the hardware expects, and that shapes the whole kernel.
  • The manual search and the automated loop agreed. Both arrived at the same operating point, the loop in a handful of steps.
Why it matters. A kernel that exists in one kernel language is out of reach for a stack built on another. This one now runs inside JAX with its number format and its numerics intact.
02 · Accelerator to accelerator

Splash attention, from TPU to NVIDIA GPU

JAX splash attention·TPU→Blackwell B200·Pallas Mosaic GPU
Correct on hardware · tuning ongoing
Reference
JAX's TPU splash attention: block-sparse flash attention that never loads or computes fully masked blocks.
Target
Blackwell, using its bulk memory transfers, tensor-core instructions and tensor memory, with warp specialization.
Coverage
Every upstream mask type; multi-head, grouped-query and multi-query attention; segment ids; logit soft-capping; bf16 and f16.
Gradients
Fully differentiable. For the common head size, a fused single-kernel backward pass replaces the two-kernel split.
Speed
Tuned per problem: the kernel variant and tile sizes are chosen automatically for each mask and head size.
Figures
to be published with reproductions

How it was verified

  • Interpreter with race detection. Numerics against a dense reference and gradients against automatic differentiation, with the barrier protocol checked for data races.
  • Lowering for the target on machines with no GPU, including shared-memory and tensor-memory budgets for every configuration.
  • Native tests on Blackwell, on a multi-GPU node.
  • Randomized fuzzing across masks, shapes, attention variants and types, for outputs and gradients, on whatever GPUs were idle.

What we found

  • Two query tiles per block, interleaved, keep the tensor core busy while the other tile's softmax runs. The variant is chosen per problem.
  • Ablation shows where the time goes. Switching off one part of the computation at a time tells the loop what to optimize next.
  • Every optimization was A/B tested on hardware across the full problem grid, and fuzzed, before it was kept.
  • Four hardware behaviours surfaced that the interpreter does not model. Each is documented and handled.
Why it matters. A team whose attention kernel exists only for one accelerator family is tied to it. This is that kernel, running and differentiable on another.
03 · The other direction

A mixed-precision grouped multiply, from CUTLASS C++ to CuTe DSL

CUTLASS C++→CuTe DSL·Hopper H100
At parity
Reference
A CUTLASS example: a grouped matrix multiply with fp8 and bf16 inputs.
Correctness
Bit-exact against the C++ original.
Speed
At parity with the C++ on the example's own benchmark configurations.
Why it matters. The loop is not specific to one target. Here the destination is NVIDIA's own Python kernel language, on the previous hardware generation.
The standard

What every case study here has to state

A kernel result without its conditions is not a result. Anything we publish, now in words and later in figures, carries the same six things.

  1. The referenceExactly which kernel, at which version, is being matched or ported.
  2. The hardware and softwareDevice, driver, toolchain and library versions, pinned.
  3. The correctness checkWhat the oracle is, what tolerance applies, and why that tolerance.
  4. The measurementDevice time, the sample count and the measured noise floor.
  5. The baselineWhat the same computation costs without the work, and against the best available alternative.
  6. The reproductionScripts that regenerate the result from pinned inputs.

The method behind these.

How a candidate is compared with its reference, and the rule that decides what is kept.