Skip to content

Add packed KDA backward megakernel - #20

Open
fdz-1999 wants to merge 9 commits into
kda-mega-forward-packedfrom
kda-mega-backward-packed
Open

fdz-1999 wants to merge 9 commits into
kda-mega-forward-packedfrom
kda-mega-backward-packed

Conversation

@fdz-1999

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

Copy link
Copy Markdown
Collaborator

Summary

PR B of the KDA megakernel migration, stacked on PR #19. Adds saved-state and rematerialized no-CP backward fusion, native packed backward input, packed gradient compaction, singleton-tile handling, Tokamax numerical tests, and CI shard coverage.

Stack: main -> PR #19 -> this PR -> PR #18.

Incremental size: 1,476 additions, 156 deletions across 15 files.

TPU validation

Environment: TPU v7x 4-chip host (2x2x1), JAX 0.11.2 with libtpu, Python 3.12.

Feature-specific suites:

pytest -q \
  tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_bwd_fused_test.py \
  tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_remat_test.py \
  tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_packed_bwd_test.py \
  tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_packed_gradients_test.py

Result: 56/56 passed.

Full KDA regression sweep:

pytest -q tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_kernel_test.py

Result: 297 passed in 5374.83s (1:29:34).

Findings and conclusion

  • An initial sweep was interrupted by an unhealthy TPU device and Pod exit 137. The environment was recreated on a healthy host and the entire sweep was rerun from the beginning; the result above is from the clean rerun.
  • Packed gradient compaction differed from gather by at most a few FP32 ULPs (1.86e-9 observed); FP32 uses 1e-7 tolerance while BF16 retains exact comparison.
  • The singleton fix loads beta before dropping its trailing singleton dimension, avoiding tiled memref_squeeze.

Saved-state fusion, rematerialized backward, packed backward input, packed gradient compaction, singleton shapes, and the existing KDA regression matrix pass on TPU v7x.

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

Baseline: openxla/tokamax #1103 pinned at 939da5c24b90dca4c73ad50fba9b9021d9f35a27, using its default Pallas/Mosaic KDA forward + custom VJP. Compared with this PR on the same TPU v7x Pod and JAX 0.11.2. BF16, H=8, B=1, T=1024, K=V=128; packed case has four non-64-aligned segments. forward_and_vjp includes both forward and backward. Each variant is independently compiled, then timed for 10 synchronized wall-clock iterations; values are medians. This replaces the earlier within-branch staged comparison.

The benchmark script is inherited from PR #19. Copy it into the pinned openxla#1103 checkout and run with --baseline, then run without that flag on this PR:

python -m tokamax.benchmarks.kda_megakernel_compare --suite backward --baseline --iterations 10
python -m tokamax.benchmarks.kda_megakernel_compare --suite backward-packed --baseline --iterations 10
# Repeat on this PR without --baseline.
Workload openxla#1103 baseline This PR Latency change Estimated peak memory
Fixed forward + VJP 0.492 ms 0.415 ms 15.7% lower 78.7 -> 64.0 MB
Packed forward + VJP 0.819 ms 0.785 ms 4.1% lower 109.6 -> 97.7 MB

Conclusion: the fixed case improves clearly; the packed latency gain is modest and needs a broader shape sweep. This is an end-to-end forward + VJP microbenchmark of the stacked changes, so it does not isolate the backward kernel from PR #19's forward improvements. Compilation and host input preparation are excluded. The TPU Pod's temporary qwix shim is not committed.

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

  • On the PR HEAD 3d5ba6a, the four backward feature suites passed on TPU v7x: 56/56 in 317.29s.
  • Two additional backward regression cases for a single head and unaligned K/V passed (2/2 in 11.53s).
  • These are current-HEAD targeted checks. The earlier 297-case full sweep reported above was not rerun in this review follow-up; JAX 0.11.2 with libtpu.

2026-09-29 stacked-HEAD TPU verification

On HEAD b1bc4e6, the fused backward, rematerialization, packed backward, and packed-gradient Tokamax suites passed 56/56 in 315.83s on the TPU v7x Pod. The new merge commit synchronizes #19's review tests and benchmark-shape script; it does not change the backward kernel. The earlier 297-case full sweep was not rerun for this benchmark/test-only stack update.

@fdz-1999
fdz-1999 force-pushed the kda-mega-backward-packed branch from bcbf1e4 to e2cf254 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

Feature-specific backward suites:

pytest -q \
  tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_bwd_fused_test.py \
  tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_remat_test.py \
  tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_packed_bwd_test.py \
  tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_packed_gradients_test.py

Result: 56/56 passed.

Full existing KDA kernel regression sweep:

pytest -q tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_kernel_test.py

Result: 297 passed in 5374.83s (1:29:34).

Validation notes:

  • The first full sweep was interrupted when the original TPU Pod was evicted with Allocate failed due to no healthy devices present and container exit code 137. It showed no pytest failure before the infrastructure interruption.
  • The environment was recreated on a healthy TPU host and the full sweep was rerun from the beginning; the result above is from that clean rerun.
  • Packed gradient compaction differs from the gather path by at most a few FP32 ULPs (1.86e-9 observed). The test now uses 1e-7 for FP32 and retains exact comparison for BF16.
  • The backward singleton-tile fix loads beta before dropping its trailing singleton dimension, avoiding an illegal tiled memref_squeeze.

Conclusion: saved-state fusion, rematerialized backward, packed backward input, packed gradient compaction, singleton shapes, and the existing KDA regression matrix pass on TPU v7x.

@fdz-1999
fdz-1999 force-pushed the kda-mega-backward-packed branch 3 times, most recently from 0df3505 to fe6db93 Compare September 23, 2026 10:09
fdz-1999 and others added 7 commits September 23, 2026 19:02
Keep recomputed intermediates and dAqk/dv0 in VMEM while preserving the
existing arithmetic, residual contract, and CP/rematerialization fallbacks.

Co-authored-by: Fred33146 <liudonglai.ldl@antgroup.com>
Combine WY recomputation with state reconstruction, retaining FP32 v_new
from the running state. Reuse local backward fusion without materializing
w/u/qg/kg or dAqk/dv0 between kernels. Keep opt-in dispatch and CP fallback.

Co-authored-by: Fred33146 <liudonglai.ldl@antgroup.com>
Mask inactive aligned-token gradients before gate reductions and propagate
final-state cotangents through empty sequences. Compare complete gradients
with XLA autodiff and remove packed CPU xfails.

Co-authored-by: Fred33146 <liudonglai.ldl@antgroup.com>
Stage original normalized Q/K, V and beta through bounded packed windows
for the reverse chunk traversal. Preserve aligned residual and gradient
contracts, with opt-in dispatch for saved and fused-rematerialized states.

Co-authored-by: Fred33146 <liudonglai.ldl@antgroup.com>
Apply output compaction after gate reductions and normalization backward,
casting feature gradients to their public dtype first. Preserve beta gather,
state gradients and the existing aligned residual contract.

Co-authored-by: Fred33146 <liudonglai.ldl@antgroup.com>
@Fred33146

Copy link
Copy Markdown
Collaborator

Reviewed the complete incremental diff against #19, including saved-state/rematerialized backward dispatch, packed input loading, gradient compaction, and its numerical tests. I did not find another confirmed blocking issue in this PR. Important stack note: #18 currently reverts several H=1 Mosaic/ref-shape fixes from this branch; those regressions are annotated on #18 and should be resolved before merging the stack.

fdz-1999 added a commit that referenced this pull request Sep 24, 2026
- Restore the tiled beta/db loads and stores from #20 (squeeze regression)
- Restore the H>1 fuse_remat guard and the heads==1 staged fallback
  in _rematerialize_states_pallas
- Force return_dh0 when an initial_state is supplied so the VJP keeps
  propagating gradients through the supplied state
- Add regression test covering the initial-state gradient for the CP
  megakernel at cp_size 2 and 4
@fdz-1999
fdz-1999 marked this pull request as ready for review September 28, 2026 08:08
@fdz-1999

Copy link
Copy Markdown
Collaborator Author

@Fred33146 Thanks for reviewing the incremental #20 diff and flagging the stack risk. The H=1 Mosaic/ref-shape regressions introduced in the stacked #18 branch have now been fixed there: the beta Ref uses load-then-reshape instead of a squeezed tiled Ref, and the H>1 fused-rematerialization guard plus H=1 staged fallback are restored (latest #18 HEAD: f99f47a). I replied to the three specific #18 review threads with the fixes and regression coverage. On that fixed #18 HEAD, the previously failing varlen/segment-fused family passed 39/39, CP fused + megakernel suites passed 34/34, and the complete KDA TPU sweep passed 297/297 in 5383.49s. #20 itself also passed its 56/56 backward feature suite and the two targeted singleton/non-aligned checks on HEAD 3d5ba6a. The #18 stack regressions you flagged are therefore addressed in code and validated on TPU; the #18 review threads remain open for reviewer confirmation.

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