Independent Researcher Makes TPU Pallas top-k Bitwise Correct and 1.67x Faster
Francis_YAO_ · x · 2026-10-11
Independent researcher Muhan Zhong published a preprint showing that a known Pallas TPU bug — kept on purpose because a fix was deemed too expensive — is both cheap to fix and beatable on speed.
- The bug: calling jax.lax.topk inside a Pallas kernel on TPU returns duplicate indices when inputs contain -inf. Errors go beyond duplicates: 70,706 of 197,592 inputs with special values produce wrong results.
- Cost of correctness: defining correctness via IEEE 754 totalOrder, a bitwise-correct formulation takes only 4.7% more cycles for top-8 of f32[8,128] — the "too expensive" trade-off doesn't hold.
- Going faster: modeling runtime via XLU round-trip latency, issue interval, and elementwise vector instruction count, the author picks formulations per shape — faster than official Pallas on 20 of 25 shapes (geometric-mean 1.67x speedup) and faster than native XLA on all tested shapes.
- For remaining shapes, two algorithms (fold-and-rank and loser-tree merge after full transpose) push coverage to 24 shapes, using instructions not currently reachable from Pallas.
More from Infra
- Anthropic engineer hits 1 billion tokens per day for a week — AaronBergman18 · 2026-10-11
- NVIDIA inference expert on when to keep, repurpose or replace aging GPUs — kimmonismus · 2026-10-11
- AI agent tunes Triton kernels on AMD MI210, flipping grid order yields 1.29x speedup — zmkzmkz · 2026-10-11
- Engineer with $120k of GPUs: bought 75% before the price surge, don't follow me — TheZachMueller · 2026-10-11
- Migrating AI workloads often means 70%+ redevelopment, warns cloud analyst — DavidLinthicum · 2026-10-11
- Qwen3.8-Flash-Next on a 5090 beats Claude Code on airbench: 11 min vs 14 min — dh7net · 2026-10-11