Repository navigation
Conversation
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>
TPU validation resultEnvironment:
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.pyResult: 29 passed in 78.81s. Findings and fixes made during validation:
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, | ||
| ) | ||
|
|
||
|
|
There was a problem hiding this comment.
[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.
There was a problem hiding this comment.
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].
|
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 |
|
Non-blocking tracking question for the #1103 x_block buffer-reuse reply. The current forward solver still materializes |
|
@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. |
|
@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. |
|
@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. |
|
@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. |
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.
Result: 29 passed in 78.81s.
Findings and conclusion
tpu.memref_squeezefailures.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.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
qwixcompatibility shim is not committed.2026-09-24 review follow-up (TPU v7x)
4085015, forward fused, packed, and output suites passed on TPU v7x: 33/33 in 110.19s.[B,T]segment-map test input.2026-09-29 review follow-ups: performance scope and x_block reuse
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.Config(fuse_forward=True). The native-packed variant usesConfig(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.jnp.stack(rows[:j], axis=1)andjnp.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(commit7ad795a). On the same TPU v7x Pod and JAX 0.11.2, we compared this PR's7ad795akernel tree with pinned openxla/tokamax openxla#1103939da5c. Inputs match the B=1/H=32/T=8192/K=V=128 BF16 training shapes inarg_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.fuse_forward=Truefuse_forward=True, packed_forward=True, packed_output=TrueThe 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_blockfollow-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. Directscatter-addanddynamic_update_sliceupdate 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 -qonpallas_mosaic_tpu_fwd_fused_test.py,pallas_mosaic_tpu_packed_test.py, andpallas_mosaic_tpu_output_test.py). This is after the benchmark-shape script commit; no forward kernel change was landed in this follow-up.