Skip to content

Add packed KDA forward megakernel - #19

Open
fdz-1999 wants to merge 13 commits into
mainfrom
kda-mega-forward-packed
Open

fdz-1999 wants to merge 13 commits into
mainfrom
kda-mega-forward-packed

Conversation

@fdz-1999

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

Copy link
Copy Markdown
Collaborator

Summary

PR A of the KDA megakernel migration. Adds the no-CP fused forward path, native packed input staging, packed output compaction, singleton-tile handling, Tokamax numerical tests, and CI shard coverage.

Stack: main -> this PR -> PR #20. PR #17 is independently stacked on this PR.

Incremental size: 1,698 additions, 85 deletions across 12 files, including the representative-shape benchmark parameters.

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/pallas_mosaic_tpu_fwd_fused_test.py \
  tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_packed_test.py \
  tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_output_test.py

Result: 29 passed in 78.81s.

Findings and conclusion

  • Forward tests are scoped to the forward contract; packed-backward/VJP coverage is in PR Add packed KDA backward megakernel #20.
  • Replaced a rank-broadcasted i1 predicate that aborted TPU lowering.
  • Zeroed each copied sequence tail before the VMEM store, preventing aligned source tails from leaking into compact output.
  • Singleton tiled layouts compile without tpu.memref_squeeze failures.

The fused no-CP forward path, native packed input path, and packed output compaction are numerically correct for the covered TPU v7x cases.

TPU performance versus openxla/tokamax openxla#1103 (2026-09-23)

Baseline: openxla/tokamax #1103 at pinned head 939da5c24b90dca4c73ad50fba9b9021d9f35a27 (open PR), using its default Pallas/Mosaic KDA path. The baseline and this PR were checked out separately on the same TPU v7x Pod with JAX 0.11.2. Inputs and timing code are identical: BF16, H=8, B=1, T=1024, K=V=128; the packed case uses four non-64-aligned segments. Each variant is compiled independently and then timed for 10 synchronized wall-clock iterations; values are medians. This replaces the earlier within-branch staged comparison, which was not the upstream baseline.

The reproducible script is tokamax/benchmarks/kda_megakernel_compare.py. Copy it into the pinned openxla#1103 checkout and run the same suite there with --baseline; run without that flag on this branch:

python -m tokamax.benchmarks.kda_megakernel_compare --suite forward --baseline --iterations 10
python -m tokamax.benchmarks.kda_megakernel_compare --suite forward-packed --baseline --iterations 10
# Repeat on this PR without --baseline.
Workload openxla#1103 baseline This PR Latency change Estimated peak memory
Fixed forward 0.299 ms 0.208 ms 30.3% lower 27.3 -> 25.2 MB
Packed forward 0.465 ms 0.337 ms 27.6% lower 32.6 -> 27.0 MB

Conclusion: both forward paths improve on these shapes versus the pinned upstream PR. Measurements exclude compilation and host input preparation and do not establish end-to-end training throughput. The TPU Pod's temporary qwix compatibility shim is not committed.

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

  • On the PR HEAD 4085015, forward fused, packed, and output suites passed on TPU v7x: 33/33 in 110.19s.
  • The suite includes H=2/4, B=2 forward and VJP regression cases, including packed output and the corrected [B,T] segment-map test input.
  • JAX 0.11.2 with libtpu; no additional code change was needed after the review fix.

2026-09-29 review follow-ups: performance scope and x_block reuse

  • The measured forward latency numbers above apply only to BF16, B=1, H=8, T=1024, K=V=128. The packed case uses four non-64-aligned segments. They are not measurements of the representative B=1/H=32/T=8192 fixed or packed N=25 training specs in arg_specs.py; no speedup is claimed for those shapes yet. A like-for-like run against the same pinned Implement KDA Pallas Kernel openxla/tokamax#1103 baseline on those two specs remains a performance follow-up.
  • The fixed fused variant uses Config(fuse_forward=True). The native-packed variant uses Config(fuse_forward=True, packed_forward=True, packed_output=True); Implement KDA Pallas Kernel openxla/tokamax#1103 uses its default Mosaic config. Both native packed flags are opt-in in this PR, so the reported packed result does not describe default dispatch.
  • The forward solver still uses jnp.stack(rows[:j], axis=1) and jnp.stack(rows, axis=1). An alternative mutable-scratch/buffer-reuse formulation has not been A/B evaluated. We retain the validated formulation for this migration and defer any change until numerical equivalence, compiler memory behavior, and latency are compared on the representative fixed and packed shapes. No buffer-reuse performance benefit is claimed here.

2026-09-29 representative-shape forward A/B and x_block evaluation

The benchmark script now accepts --heads, --training-tokens, and --num-segments (commit 7ad795a). On the same TPU v7x Pod and JAX 0.11.2, we compared this PR's 7ad795a kernel tree with pinned openxla/tokamax openxla#1103 939da5c. Inputs match the B=1/H=32/T=8192/K=V=128 BF16 training shapes in arg_specs.py; packed N=25 uses its balanced segment lengths. Two runs were made in alternating checkout order, with 10 and 20 synchronized wall-clock iterations respectively; entries below are run medians. This is forward-only, not training step throughput.

Workload openxla#1103 run 1 / run 2 #19 run 1 / run 2 #19 configuration Estimated peak memory (openxla#1103 -> #19)
Fixed forward 1.990 / 1.970 ms 1.906 / 1.870 ms fuse_forward=True 1074.3 -> 671.6 MB
Packed forward, N=25 4.052 / 4.035 ms 2.891 / 2.902 ms fuse_forward=True, packed_forward=True, packed_output=True 1459.2 -> 835.1 MB

The fixed result is roughly 4-5% lower latency across the two runs; native packed is roughly 28% lower. The packed flags remain opt-in, and these results must not be read as default-dispatch performance. The script was copied unchanged into the pinned openxla#1103 checkout and run there with --baseline; on this PR, omit that flag:

python -m tokamax.benchmarks.kda_megakernel_compare --suite forward --heads 32 --training-tokens 8192 --iterations 20 --baseline
python -m tokamax.benchmarks.kda_megakernel_compare --suite forward-packed --heads 32 --training-tokens 8192 --num-segments 25 --iterations 20 --baseline
# Repeat both commands on this PR without --baseline.

For the x_block follow-up, an immutable fixed-size block with per-row concatenate updates passed the forward/packed/output TPU suite (33/33), but one representative-shape A/B showed no stable benefit: fixed 1.885 vs 1.906 ms and packed 2.903 vs 2.891 ms (candidate vs original); estimated peak memory was unchanged. Direct scatter-add and dynamic_update_slice update variants did not lower in current Pallas TPU. The candidate kernel change was reverted; no buffer-reuse speedup is claimed. A different Pallas-compatible scratch formulation would need a fresh correctness and paired performance evaluation.

2026-09-29 stacked-HEAD TPU verification

On HEAD 7ad795a, the forward-fused, packed-input, and compact-output Tokamax suites passed 33/33 in 111.81s on the TPU v7x Pod (pytest -q on pallas_mosaic_tpu_fwd_fused_test.py, pallas_mosaic_tpu_packed_test.py, and pallas_mosaic_tpu_output_test.py). This is after the benchmark-shape script commit; no forward kernel change was landed in this follow-up.

fdz-1999 and others added 6 commits September 22, 2026 23:24
Connect the existing intra-chunk and state/output kernels through VMEM
scratch while preserving aligned preprocessing and backward residuals.
Add staged/fused comparisons, focused correctness tests and design notes.

Co-authored-by: Fred33146 <liudonglai.ldl@antgroup.com>
Stage 72-row input windows in VMEM while preserving aligned outputs and
backward residuals. Add opt-in Mosaic configuration, boundary/VJP coverage,
and an independent device-timed benchmark variant.

Co-authored-by: Fred33146 <liudonglai.ldl@antgroup.com>
Copy aligned sequence chunks sequentially through token-major VMEM scratch,
mask output padding, and retain gather for unsupported or oversized groups.
Preserve the public output and backward residual contracts.

Co-authored-by: Fred33146 <liudonglai.ldl@antgroup.com>
@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/pallas_mosaic_tpu_fwd_fused_test.py \
  tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_packed_test.py \
  tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_output_test.py

Result: 29 passed in 78.81s.

Findings and fixes made during validation:

  • Forward tests are scoped to the forward contract; packed-backward/VJP coverage lives in the stacked backward PR.
  • Replaced a rank-broadcasted i1 predicate that aborted TPU lowering.
  • Zeroed each copied sequence tail before the VMEM store, preventing poisoned aligned source tails from leaking into the compact output.
  • Singleton H/B/K-related tiled layouts compile and run without tpu.memref_squeeze failures.

Conclusion: the fused no-CP forward path, native packed input path, and packed output compaction are numerically correct on TPU v7x for the cases covered by this PR.

lambda self, device: True,
)


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] All forward equivalence cases built by _inputs use H=1, while the new fused kernel chooses a different head-group/VMEM layout when H>1. The benchmark runs H=8 but does not check numerical output. Please add at least one TPU correctness case with multiple heads (ideally a nontrivial mini_batch), including packed output and the VJP residual contract, before relying on the default fuse_forward=True path for those shapes.

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.

Addressed in 64622d9 (multi-head forward and VJP correctness coverage) and 4085015 (broadcast the packed segment map across the mini batch in the test). The added TPU cases cover H=2/4, B=2, nontrivial mini batches, packed output, and the VJP residual contract. On PR HEAD 4085015, the forward fused, packed, and output suites passed on TPU v7x: 33/33 in 110.19s. The PR description records the command and environment.

The suite only exercised heads=1, where Mosaic's (8, 128) tiling pads
the singleton head dimension and masks the multi-head layout entirely.
Add packed and unpacked cases with heads 2/4 and a nontrivial batch,
comparing fused vs staged kernels and against XLA.
_inputs hard-coded a [1, T] segment map, so the new multi-head cases
with batch=2 failed the 'B T' typecheck. Broadcast the shared map to
[B, T].
@fdz-1999
fdz-1999 marked this pull request as ready for review September 28, 2026 08:07
@Fred33146

Fred33146 commented Sep 29, 2026 •

Copy link
Copy Markdown
Collaborator

Non-blocking follow-up to the #1103 forward-megakernel plan: this PR does implement the fused non-CP forward path and direct packed-input path. The reported BF16 B=1/H=8/T=1024 comparisons against pinned openxla#1103 support the measured speedups. However, the representative KDA training specs in arg_specs.py are B=1/H=32/T=8192, including packed N=25, and I could not find like-for-like results for those shapes. Could you add fixed and packed forward comparisons on those specs against the same baseline (or explicitly scope the performance claim to the measured shapes)? Please also state whether packed_forward/packed_output were enabled: native packed reads are currently opt-in (packed_forward=False by default). This is a performance-evidence/documentation follow-up, not a correctness blocker.

@Fred33146

Copy link
Copy Markdown
Collaborator

Non-blocking tracking question for the #1103 x_block buffer-reuse reply. The current forward solver still materializes jnp.stack(rows[:j], axis=1) and jnp.stack(rows, axis=1) in pallas_mosaic_tpu_fwd_kernel.py (around lines 489/494). The prior reply said the next KDA update would revisit reuse and determine whether a buffer-reuse formulation is worthwhile; it did not promise that reuse would necessarily be implemented. Has that evaluation happened, and was the current formulation intentionally retained? If deferred, could you document the decision or link a follow-up so the item is traceable?

@fdz-1999

Copy link
Copy Markdown
Collaborator Author

@Fred33146 Thanks for tracking this. The current solver still materializes both jnp.stack(rows[:j], axis=1) and jnp.stack(rows, axis=1). We have not A/B evaluated a mutable-scratch/buffer-reuse alternative, so I cannot claim that the current formulation was retained after a measured comparison. I documented this as deferred in the #19 PR description, with numerical equivalence, compiler memory behavior, and representative fixed/packed latency as the criteria for revisiting it. No buffer-reuse benefit is claimed in this PR.

@fdz-1999

Copy link
Copy Markdown
Collaborator Author

@Fred33146 Thanks. I updated the #19 PR description to scope the measured forward speedups strictly to BF16 B=1/H=8/T=1024 (four non-64-aligned segments for packed); we have not run the representative B=1/H=32/T=8192 fixed or packed N=25 comparison, so no speedup is claimed for those specs. The benchmark's fixed fused variant uses Config(fuse_forward=True); its native-packed variant explicitly sets packed_forward=True and packed_output=True, both opt-in, against the pinned openxla#1103 default Mosaic baseline. The like-for-like representative-shape comparison remains a performance follow-up rather than an inferred result.

@fdz-1999

Copy link
Copy Markdown
Collaborator Author

@Fred33146 Follow-up to your representative-shape performance question (#issuecomment-5884250971): I parameterized the benchmark script (7ad795a) and updated the #19 description with reproducible commands and two same-Pod runs against pinned openxla#1103 939da5c. BF16 B=1/H=32/T=8192 fixed forward measured openxla#1103 1.990/1.970 ms vs #19 1.906/1.870 ms; packed N=25 measured 4.052/4.035 ms vs 2.891/2.902 ms. These are forward-only run medians (10 and 20 synchronized iterations), not full training throughput. The packed run explicitly sets packed_forward=True and packed_output=True; those flags remain opt-in, so the result is not a default-dispatch claim.

@fdz-1999

Copy link
Copy Markdown
Collaborator Author

@Fred33146 Follow-up to your x_block reuse question (#issuecomment-5884254847): I evaluated a fixed-size-block/row-concatenate candidate on TPU. It passed the 33/33 forward, packed, and output correctness suite, but representative H=32/T=8192 A/B showed no stable improvement: fixed 1.885 vs 1.906 ms and packed N=25 2.903 vs 2.891 ms (candidate vs current), with unchanged estimated peak memory. Direct scatter-add and dynamic_update_slice variants fail current Pallas TPU lowering. I reverted the candidate and documented the results in #19; the existing stack formulation is retained based on this evidence, not on an untested assumption.

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