Skip to content

Add standalone KDA inference megakernel - #17

Open
fdz-1999 wants to merge 7 commits into
kda-mega-forward-packedfrom
kda-mega-inference-v2
Open

fdz-1999 wants to merge 7 commits into
kda-mega-forward-packedfrom
kda-mega-inference-v2

Conversation

@fdz-1999

@fdz-1999 fdz-1999 commented Sep 22, 2026 •

Copy link
Copy Markdown
Collaborator

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.

pytest -q \
  tokamax/_src/ops/experimental/kda/inference_test.py \
  tokamax/_src/ops/experimental/kda/inference_lowering_test.py

Result: 20 passed in 19.15s.

Findings and conclusion

  • Covers numerical inference behavior and TPU lowering.
  • Covers native segment IDs and different sequence boundaries within an aligned chunk.
  • Full-rank broadcasted_iota predicates avoid the TPU i1 shape-cast failure caused by broadcasting jnp.arange predicates.

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.
Tokens openxla#1103 Mosaic forward Native inference Latency change Estimated peak memory
512 0.322 ms 0.227 ms 29.5% lower 21.0 -> 5.4 MB
8192 1.326 ms 0.541 ms 59.2% lower 305.6 -> 86.2 MB

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 qwix shim is not committed.

2026-09-24 review follow-up (TPU v7x)

  • On the PR HEAD 02ce4fd, inference and lowering suites passed on TPU v7x: 22/22 in 21.99s.
  • This includes the shared rank-1 [T] segment map for B=3, verifying the public API broadcast fix and the independent XLA output comparison.
  • JAX 0.11.2 with libtpu; no additional code change was needed after the review fix.

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.

@fdz-1999
fdz-1999 force-pushed the kda-mega-inference-v2 branch from 32cd27a to f500879 Compare September 22, 2026 16:04
@fdz-1999

Copy link
Copy Markdown
Collaborator Author

TPU validation result

Environment:

  • TPU v7x 4-chip host, topology 2x2x1
  • JAX 0.11.2 with libtpu
  • Python 3.12

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.py

Result: 20 passed in 19.15s.

Validation notes:

  • Covers numerical inference behavior and TPU lowering.
  • Covers native segment-ID handling and different sequence boundaries inside an aligned chunk.
  • Full-rank broadcasted_iota predicates avoid the TPU i1 shape-cast failure previously caused by jnp.arange broadcasting.

Conclusion: the standalone KDA inference megakernel and its lowering path compile and produce the expected results on TPU v7x.

@fdz-1999
fdz-1999 force-pushed the kda-mega-inference-v2 branch 3 times, most recently from 389e0fe to 182c066 Compare September 23, 2026 10:09
fdz-1999 and others added 3 commits September 23, 2026 19:04
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>
@fdz-1999
fdz-1999 force-pushed the kda-mega-inference-v2 branch from 182c066 to 8c11a30 Compare September 23, 2026 11:04
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, :]

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

[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.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.
@fdz-1999
fdz-1999 marked this pull request as ready for review September 28, 2026 08:08
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants