Skip to content

TokenSpeed MLA 2 for NVIDIA Vera Rubin

Oct 12, 2026
by TokenSpeed Team

TokenSpeed MLA 2 optimizes our multi-head latent attention (MLA) decode kernel for NVIDIA Vera Rubin. It succeeds the original TokenSpeed MLA implementation for Blackwell, referred to here as MLA 1. In our long-context benchmarks, MLA 2 on NVIDIA Vera Rubin delivers kernel speedups of up to 3.44× over MLA 1 on NVIDIA B200 and up to 2.84× over MLA 1 on NVIDIA GB200. These cross-hardware peaks include both hardware and software changes.

We released TokenSpeed MLA in May 2026 for long-context agentic inference on Blackwell. The TokenSpeed launch benchmarks outperformed the then state-of-the-art TensorRT-LLM MLA baseline on representative prefill and speculative-decoding workloads. TokenSpeed served Kimi K2.5 at launch and later added Kimi K3 support.

The tokenspeed-mla package has reached more than 2 million monthly downloads on PyPI.[1]

Weekly PyPI downloads, June 15–October 10, 2026. Labels use w (10,000), rounded to one decimal; the last point covers 6 days.

Figure 1. Weekly PyPI downloads from June 15 to October 10, 2026; the final point covers six days.

NVIDIA Vera Rubin capabilities ​

On NVIDIA Vera Rubin, the K=64-byte matrix multiply-accumulate (MMA) instruction for 8-bit floating-point (FP8) data provides twice the arithmetic throughput per streaming multiprocessor (SM) per clock of the K=32-byte form. Larger shared memory and tensor memory (TMEM) support wider tiles and deeper buffers, but padding, pipeline overhead, and reduction costs can limit the benefit.

MLA 2 builds on a FlashInfer-derived pipeline and adapts it to Kimi K3 workloads with 96 query heads (H96) and 4 or 8 query tokens (Q4 or Q8).

Attention pipeline overlap ​

MLA computes query-key scores (QK), applies softmax, then multiplies probabilities by values (PV). MLA 1 processes 128 query-head rows per tile (M128). One cooperative thread array (CTA) pair alternates QK and PV using K=32-byte MMA. Scores and output accumulators share TMEM, and probabilities stay in local shared memory.

MLA 2 separates QK and softmax from PV into two CTA pairs, so stages can overlap across tiles. Its M256 tiles reuse key-value (KV) data across 256 query-head rows and use K=64-byte MMA. Queries stay in TMEM; probabilities pass between pairs through distributed shared memory (DSMEM). Deeper key and value buffers and two interleaved softmax groups support the overlap. Figure 2 compares the two pipelines.

Architecture comparison of the Blackwell M128 two-CTA pipeline and Vera Rubin M256 four-CTA pipeline. Vera Rubin separates QK/softmax and PV into two CTA pairs, transfers probabilities through DSMEM, and overlaps the stages. Orange marks the Blackwell baseline; blue marks the Vera Rubin pipeline.

Figure 2. MLA 1’s Blackwell data flow and the four-CTA pipeline for NVIDIA Vera Rubin shared by FlashInfer and MLA 2. The Blackwell configuration uses 3 key and 2 value stages; MLA 1 recompiled for NVIDIA Vera Rubin uses 4 of each.

Kimi K3 workload optimizations ​

U.S. AI companies are building specialized agents on Kimi K3’s open weights. Cognition post-trained SWE-2 for coding, while Harvey built Tenet for long-horizon legal work, both using Kimi K3 as their foundation model.

To improve Kimi K3’s long-context decoding on NVIDIA Vera Rubin, MLA 2 uses native query-row packing, tunes KV split selection for better scheduling, and optimizes reduction to lower overhead.

Native row packing ​

The benchmarked FlashInfer multi-token prediction (MTP) kernel supports 128 query heads (H128) and 2 or 4 query tokens (Q2 or Q4). Its adapter pads H96 to H128 and splits Q8 into 2 Q4 requests, executing 4 M256 tiles. We call this kernel plus adapter the FlashInfer baseline. MLA 2 packs each token’s 96 heads consecutively, fitting all 768 useful rows into 3 tiles while preserving causal masking.

Q4 requires 2 M256 tiles for 384 rows. The final tile executes a full MMA, leaving 128 of 512 row slots unused. Figure 3 compares this 75% row occupancy with the fully occupied Q8 tiles.

Native H96 row packing: Q4 uses two M256 tiles with 75% useful rows; Q8 uses three fully occupied tiles. FlashInfer H128 adaptation uses four tiles with 75% useful rows. Blue marks TokenSpeed rows, orange marks FlashInfer rows, and gray marks padding.

Figure 3. Native H96 row packing compared with H128 adaptation. Native Q8 uses 3 M256 tiles instead of 4. Q4 uses 2 tiles with 75% useful rows in both implementations; gray marks padding.

Scheduling and reduction ​

Splitting the KV context adds parallelism, but each split incurs pipeline and reduction overhead. MLA 2 selects splits for 4 CTAs per work item and up to 2 scheduling waves. It leaves full execution grids unsplit.

MLA 2 also tunes the reducer to lower the cost of combining partial attention results. The benchmarks measure these optimizations together, not their individual contributions.

Benchmark results ​

Across 16 H96 configurations, we compare MLA 2 with two separate baselines on the same NVIDIA Vera Rubin GPU: recompiled MLA 1 and adapted FlashInfer. The final comparison changes both hardware and implementation. Speedup is baseline latency divided by target latency for the same configuration: 2× means half the latency. Summary values are geometric means; peaks are the best individual results.[2]

MLA 2 versus recompiled MLA 1 ​

For the same-GPU comparison, we recompiled MLA 1’s Blackwell M128 kernel for NVIDIA Vera Rubin with four key and four value stages. Both kernels use FP8 inputs and outputs. MLA 2 has lower median latency in 13 of 16 configurations (Figure 4), with geometric mean speedups of 1.142× overall, 1.271× for Q8, and 1.027× for Q4.

MLA 2 versus recompiled MLA 1 on the same Vera Rubin GPU across 16 H96 FP8 workloads. MLA 2 geometric-mean speedup is 1.142× overall, 1.027× for Q4, and 1.271× for Q8. Orange bars mark the 1× MLA 1 baseline; blue bars mark MLA 2.

Figure 4. MLA 2 versus recompiled MLA 1 on NVIDIA Vera Rubin. MLA 1 is normalized to 1×; MLA 2 bars show MLA 1 latency divided by MLA 2 latency. Values greater than 1× indicate lower MLA 2 latency.

Q8 uses 3 M256 tiles instead of 6 M128 tiles and improves in all 8 configurations. Q4’s partially filled final tile adds 33% more row slots than MLA 1, consistent with its smaller gain.

The small-batch behavior is consistent with pipeline fill/drain, split reduction, and fixed launch costs taking a larger fraction of total latency.

Together, these results show that MLA 2’s NVIDIA Vera Rubin pipeline benefits most when workloads fill its M256 tiles.

MLA 2 versus adapted FlashInfer ​

On the same NVIDIA Vera Rubin GPU, MLA 2 achieves geometric mean speedups of 1.327× overall, 1.227× for Q4, and 1.434× for Q8 over adapted FlashInfer. The two share the same attention pipeline (Figure 2). This comparison measures native Kimi K3 support, tuning, and adapter removal against the benchmarked implementation, not native H128 performance or current FlashInfer releases.

MLA 2 versus the H96-adapted FlashInfer PR 6177 baseline on the same Vera Rubin GPU. Across 16 shapes, MLA 2 geometric-mean speedup is 1.327× overall, 1.227× for Q4, and 1.434× for Q8. Adapter timing includes query packing and output unpacking. Orange marks FlashInfer; blue marks MLA 2.

Figure 5. MLA 2 versus adapted FlashInfer on NVIDIA Vera Rubin. FlashInfer is normalized to 1×; MLA 2 bars show baseline latency divided by MLA 2 latency. Baseline timing includes query packing and output unpacking.

At large-batch Q8, native packing removes one of four M256 tiles. This yields speedups of 1.325× at B64/64K and 1.271× at B64/128K, close to the tile-count model’s estimated 4/3 = 1.33× benefit.

At large-batch Q4, both paths execute two M256 tiles at 75% utilization, so native packing does not reduce the tile count and performance remains nearly identical.

At small batch, however, adapter work, split-KV reduction, and fixed launch overhead account for a larger share of latency; removing the adaptation path and reducing split-KV overhead therefore has a larger effect.

NVIDIA Vera Rubin versus Blackwell ​

Across the same 16 configurations, MLA 2 on NVIDIA Vera Rubin achieves geometric mean speedups of 2.62× over MLA 1 on B200 and 2.05× over MLA 1 on GB200. Each run uses one GPU; the hardware differs between runs. The peaks cited in the introduction, 3.44× over B200 and 2.84× over GB200, both occur at batch size 16 with Q8 and KV64K.

TokenSpeed MLA speedup across B200, B300, GB200, GB300, and Vera Rubin for Q4/Q8, KV64K/KV128K, and batch sizes 1, 4, 16, and 64. B200 is fixed at 1×. Orange marks the B200 baseline; three gray shades distinguish the other Blackwell GPUs; blue highlights Vera Rubin. Vera Rubin bars are directly labeled and tables provide all series values.

Figure 6. TokenSpeed MLA across GPU configurations. Every value is normalized to B200 at 1×: B200 latency divided by target latency. To compare with GB200, divide the NVIDIA Vera Rubin ratio by the GB200 ratio for the same configuration. Blackwell GPUs run MLA 1’s M128 kernel; NVIDIA Vera Rubin runs MLA 2.

Hardware, power, clocks, and implementations differ, so these results do not isolate architectural gains. Interpret small differences within the reported run variability.

Results and next steps ​

MLA 2 pairs NVIDIA Vera Rubin’s wider tiles with Kimi K3’s native layout. Q8 benefits most because it fills those tiles without padding. Q4 gains are smaller, and its small-batch regressions need further profiling.

These results measure kernel latency, not end-to-end serving performance on NVIDIA Vera Rubin.

Acknowledgments ​

We thank the FlashInfer and Fast Kernel teams for the kernel work that MLA 2 builds on. Special thanks to NVIDIA for Vera Rubin early access and to the NVIDIA DevTech team for their brilliant work and continued collaboration from MLA 1 through MLA 2.


  1. Source: PyPI Stats, retrieved October 11, 2026. Monthly downloads refer to the latest 30-day total; counts exclude known mirrors. ↩︎

  2. The benchmark covers 16 combinations of batch sizes 1, 4, 16, and 64; 96 query heads; 4 or 8 query tokens; and KV context lengths of 65,536 or 131,072 tokens. Figures use B for batch size, H for query heads, Q for query tokens, and KV64K or KV128K for context length.

    The same-GPU tests on NVIDIA Vera Rubin use FP8 queries, KV data, and outputs; 64-token pages; a cold L2 cache through KV rotation; and disabled programmatic dependent launch (PDL). We report median latency from 15 repeats of 20 CUDA graph iterations. Timing includes split-KV reduction and, for adapted FlashInfer, query packing and output unpacking.

    Overall geometric means give equal weight to all 16 configurations; the Q4 and Q8 means each cover 8 configurations. Cross-hardware ratios use the named Blackwell baseline and MLA 2 on NVIDIA Vera Rubin. They are separate comparisons, not additional factors to multiply by the same-GPU speedups. ↩︎

© 2026 LightSeek Foundation. CC BY 4.0.