Two years ago, I complained to David Peterson that TPUs are difficult to use, and he said that TPU's infra-first design is elegant and natural for a fresh mind, it's just people are too familiar with cuda's SIMT programming model that TPU seems strange.
Now, with the help from strong AI agents, we can catch up with the great minds from David Peterson and @JeffDean to enjoy the beauty of TPU. Single-thread programming model, offload the complexity of scheduling to AI agents.
Put it short, maybe we misunderstood TPU for several years. We are just too dumb to use TPU as efficient as @JeffDean . Now we can glance at what @JeffDean sees with the help of AI agents.
Kudos to the great team from @inferact ๐ We are building the inference software for the incoming AGI era, come to build together with us! ๐๐ป https://t.co/E0CdIetHMP
When I first learned about the results, I was very surprised. Megakernels have been hard on GPUs, barely beating CUDAgraph+PDL.
Then @woosuk_k explained to me. TPU only has 1-2 core compute units, instead of hundreds SMs on NVIDIA GPUs, making megakernels very natural (synchronization across SMs is what kills megakernels on NVIDIA)
Then I wonder why no one has done it earlier. Perhaps cuz they don't have the cracked @woosuk_k and @inferact team ๐ซก
When I first learned about the results, I was very surprised. Megakernels have been hard on GPUs, barely beating CUDAgraph+PDL.
Then @woosuk_k explained to me. TPU only has 1-2 core compute units, instead of hundreds SMs on NVIDIA GPUs, making megakernels very natural (synchronization across SMs is what kills megakernels on NVIDIA)
Then I wonder why no one has done it earlier. Perhaps cuz they don't have the cracked @woosuk_k and @inferact team ๐ซก
Our first TPU megakernel for Kimi K3 reaches 709 tokens/s on low-concurrency decode, against 450 tokens/s for our GB200 baseline, both with DSpark speculative decoding.
To our knowledge, this is the first TPU inference megakernel. The whole model runs in a single Pallas kernel, and without spec decoding it is roughly 1.4 to 2x the GB200 baseline at batch sizes 1 through 8.
We are open sourcing it today.
1/2
SOTA TPU performance that beats GPU by exploiting TPU's large unified VMEM!
Keep in mind that TPU v7 has slightly worse hardware spec (also cheaper) than GB200.
Our first TPU megakernel for Kimi K3 reaches 709 tokens/s on low-concurrency decode, against 450 tokens/s for our GB200 baseline, both with DSpark speculative decoding.
To our knowledge, this is the first TPU inference megakernel. The whole model runs in a single Pallas kernel, and without spec decoding it is roughly 1.4 to 2x the GB200 baseline at batch sizes 1 through 8.
We are open sourcing it today.
1/2
Our first TPU megakernel for Kimi K3 reaches 709 tokens/s on low-concurrency decode, against 450 tokens/s for our GB200 baseline, both with DSpark speculative decoding.
To our knowledge, this is the first TPU inference megakernel. The whole model runs in a single Pallas kernel, and without spec decoding it is roughly 1.4 to 2x the GB200 baseline at batch sizes 1 through 8.
We are open sourcing it today.
1/2