Repository navigation
Conversation
bcbf1e4 to
e2cf254
Compare
TPU validation resultEnvironment:
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.pyResult: 56/56 passed. Full existing KDA kernel regression sweep: pytest -q tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_kernel_test.pyResult: 297 passed in 5374.83s (1:29:34). Validation notes:
Conclusion: saved-state fusion, rematerialized backward, packed backward input, packed gradient compaction, singleton shapes, and the existing KDA regression matrix pass on TPU v7x. |
0df3505 to
fe6db93
Compare
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>
fe6db93 to
3d5ba6a
Compare
|
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. |
- 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
|
@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. |
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:
Result: 56/56 passed.
Full KDA regression sweep:
Result: 297 passed in 5374.83s (1:29:34).
Findings and conclusion
1.86e-9observed); FP32 uses1e-7tolerance while BF16 retains exact comparison.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_vjpincludes 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.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
qwixshim is not committed.2026-09-24 review follow-up (TPU v7x)
3d5ba6a, the four backward feature suites passed on TPU v7x: 56/56 in 317.29s.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.