Repository navigation
Conversation
32cd27a to
f500879
Compare
TPU validation resultEnvironment:
Executed in the Tokamax checkout for this PR: pytest -q \
tokamax/_src/ops/experimental/kda/inference_test.py \
tokamax/_src/ops/experimental/kda/inference_lowering_test.pyResult: 20 passed in 19.15s. Validation notes:
Conclusion: the standalone KDA inference megakernel and its lowering path compile and produce the expected results on TPU v7x. |
389e0fe to
182c066
Compare
Port the inference-only source without a training residual tape. Handle arbitrary in-chunk boundaries and preserve empty final-state slots. Co-authored-by: Fred33146 <liudonglai.ldl@antgroup.com>
Use mask/reduce/select instead of dynamic slices and add abstract TPU7x export coverage independent of interpretation. Co-authored-by: Fred33146 <liudonglai.ldl@antgroup.com>
182c066 to
8c11a30
Compare
| if q.shape[1] % 64: | ||
| raise ValueError("T must be padded to a multiple of 64") | ||
| if segment_ids.ndim == 1: | ||
| segment_ids = segment_ids[None, :] |
There was a problem hiding this comment.
[P2] The public signature documents segment_ids with shape [T], but for B>1 this only becomes [1,T] and the following shape check rejects it. The underlying native kernel already broadcasts a 1D segment map across the batch. Please broadcast to q.shape[:2] here, or explicitly limit the documented [T] form to B=1, and add a B>1 test. This issue is shared with #15.
There was a problem hiding this comment.
Fixed in 6e2b372 (broadcast a rank-1 [T] segment map to q.shape[:2]) and 02ce4fd (broadcast the same shared map for the independent XLA reference). Added a B=3 shared-[T] regression case. On PR HEAD 02ce4fd, the inference and lowering suites passed on TPU v7x: 22/22 in 21.99s. The PR description records the command and environment.
The wrapper expanded a [T] segment_ids to [1, T], so B>1 inputs failed the shape check even though the native kernel broadcasts the map. Broadcast 1D maps to [B, T] instead and add a regression test.
The XLA reference only accepts the documented [B, T] layout; the inference wrapper under test keeps the raw 1D form.
Summary
PR D of the KDA megakernel migration, independently stacked on PR #19. Adds standalone native segment-ID inference, TPU-safe token selection, lowering regression coverage, Tokamax numerical tests, documentation, and benchmark scaffolding.
Stack:
main-> PR #19 -> this PR.Incremental size: 2,017 additions across 7 files (including review fixes and benchmark-shape scaffolding inherited through the stack).
TPU validation
Environment: TPU v7x 4-chip host (2x2x1), JAX 0.11.2 with libtpu, Python 3.12.
Result: 20 passed in 19.15s.
Findings and conclusion
broadcasted_iotapredicates avoid the TPU i1 shape-cast failure caused by broadcastingjnp.arangepredicates.The standalone KDA inference megakernel and its lowering path compile and produce expected results on TPU v7x.
TPU performance versus openxla/tokamax openxla#1103 (2026-09-23)
Baseline: openxla/tokamax #1103 pinned at
939da5c24b90dca4c73ad50fba9b9021d9f35a27, using its existing Mosaic forward path. This PR adds standalone native inference; openxla#1103 has no equivalent standalone inference kernel. Both checkouts used the same TPU v7x Pod, JAX 0.11.2, BF16, B=1, H=8, K=V=128, and two packed segments with a non-64-aligned boundary. Variants were compiled independently and timed for 10 synchronized wall-clock iterations; values are medians. This replaces the earlier within-branch Mosaic comparison.The shared benchmark script is inherited from PR #19. Copy it into the pinned openxla#1103 checkout and run with
--baseline; run without that flag on this PR:python -m tokamax.benchmarks.kda_megakernel_compare --suite inference --tokens 512 --baseline --iterations 10 python -m tokamax.benchmarks.kda_megakernel_compare --suite inference --tokens 8192 --baseline --iterations 10 # Repeat on this PR without --baseline.Conclusion: native inference improves both measured latency and estimated peak memory for these two packed shapes versus openxla#1103. This compares complete forward calls, not autoregressive single-token decoding or serving throughput. Compilation and host input preparation are excluded; the Pod's temporary
qwixshim is not committed.2026-09-24 review follow-up (TPU v7x)
02ce4fd, inference and lowering suites passed on TPU v7x: 22/22 in 21.99s.[T]segment map for B=3, verifying the public API broadcast fix and the independent XLA output comparison.2026-09-29 stacked-HEAD TPU verification
On HEAD
98dab14, the Tokamax inference and inference-lowering suites passed 22/22 in 22.04s on the TPU v7x Pod. The merged #19 base adds forward review tests and the benchmark-shape script; the inference kernel itself is unchanged.