Repository navigation
Conversation
0a2c989 to
82a7af3
Compare
bcbf1e4 to
e2cf254
Compare
82a7af3 to
ae03e55
Compare
TPU validation resultEnvironment:
Executed in the Tokamax checkout for this PR: pytest -q \
tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_cp_fused_test.py \
tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_cp_megakernel_test.pyFinal result: 30 passed in 131.69s (2:11). Findings and fixes made during validation:
Conclusion: the production BF16 CP megakernel path is numerically validated for CP sizes 2 and 4, saved/rematerialized state policies, and single/multiple-segment inputs. FP32 preserves correctness through the staged fallback without affecting the BF16 performance path. |
81a6653 to
323337a
Compare
809ed3e to
e2ec8d5
Compare
323337a to
0df3505
Compare
e2ec8d5 to
97935ac
Compare
0df3505 to
fe6db93
Compare
97935ac to
90cebad
Compare
fe6db93 to
3d5ba6a
Compare
Reuse local reverse fusion and omit the pre-collective dAqk output. Keep the experimental configuration opt-in pending TPU validation. Co-authored-by: Fred33146 <liudonglai.ldl@antgroup.com>
Require explicit CP and rematerialization fusion opt-ins. Preserve staged CP state reconstruction and its FP32 v_new across communication. Co-authored-by: Fred33146 <liudonglai.ldl@antgroup.com>
Fuse state reconstruction, CP preprocessing, remote-DMA ring and reverse gradients in one Pallas call, preserving Tokamax local CP initial-state mapping. Keep the new Config opt-in pending TPU validation. Co-authored-by: Fred33146 <liudonglai.ldl@antgroup.com>
90cebad to
bc1eb4d
Compare
| cp_active=True, cp_context=context_parallel_metadata, | ||
| cp_size=context_parallel_metadata.cp_size, | ||
| cp_axis_name=context_parallel_metadata.axis_name, | ||
| N_MAX=max_num_segments, return_dh0=False, |
There was a problem hiding this comment.
[P1] When cp_megakernel is enabled with an initial_state, return_dh0=False discards its gradient. The VJP then substitutes zeros for the missing dh0, so training cannot propagate through the supplied state. Please request dh0 when an initial state is present and add a CP test that differentiates with respect to initial_state. The existing CP megakernel test only differentiates q/k/v/g/beta, so it cannot detect this.
There was a problem hiding this comment.
Fixed in the #18 stack. The VJP now requests dh0 when the caller supplied initial_state (return_dh0=has_initial_state), while the CP kernel still uses its prepared recurrence seed independently of that caller flag. Added CP=2/4 saved/remat initial-state-gradient coverage that differentiates with respect to initial_state and checks the gradient is nonzero. The CP fused + megakernel suites passed 34/34, and the complete KDA TPU sweep on f99f47a passed 297/297 (5383.49s).
| k_ref.at[:, 0, 0], | ||
| v_ref.at[:, 0, 0], | ||
| beta_tile, | ||
| beta_ref.at[:, 0, 0, :, 0], |
There was a problem hiding this comment.
[P1] This change reintroduces the squeezed beta Ref that #20 explicitly removed because Mosaic lowers it to an unsupported nested-tile memref_squeeze. The default non-CP fused backward still selects this path for H=1, so stacking #18 on #20 can break the singleton TPU case. Please retain the load-then-reshape form from #20 and rerun its H=1 backward tests on TPU.
There was a problem hiding this comment.
Fixed. The fused M4 beta load uses load-then-reshape, and f99f47a also restores the same safe beta-tile load in the inlined saved-state backward path; no squeezed beta Ref view is passed to Mosaic. The previously failing varlen/segment-fused family passed 39/39, the #20 H=1 backward regression checks passed, and the complete #18 KDA TPU sweep on f99f47a passed 297/297 (5383.49s).
| fuse_rematerialization | ||
| and fuse_backward | ||
| and not disable_recompute | ||
| and not _cp_active |
There was a problem hiding this comment.
[P1] The H > 1 guard from #20 was removed here, and this PR also removes the H=1 staged fallback in _rematerialize_states_pallas. #20 documents that the fused MB=1 recursive state update is not numerically equivalent to the staged kernel on current Mosaic. With fuse_rematerialization=True, H=1 now takes that known-bad path. Please restore the guard/fallback and include the singleton rematerialization regression test in TPU validation.
There was a problem hiding this comment.
Fixed. The H > 1 guard for fused rematerialization and the H=1 staged fallback in _rematerialize_states_pallas are restored in the #18 branch. The H=1 backward regression cases and CP saved/remat suites passed on TPU; the complete KDA sweep on f99f47a passed 297/297 (5383.49s).
- 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
The first attempt keyed return_dh0 on \initial_state is not None\, but CP forward always replaces a missing user state with the CP-prepared zeros state, so the check was always true and the shared empty-sequence fixup read an unbound dht_m4. - Request dh0 via the has_initial_state argument instead (the VJP's report of whether the user supplied a state) - Bind dht_m4 on the mega path so the shared fixup stays well-defined - Exercise the regression at the backward boundary: the public op still rejects CP with an initial_state, so the test drives chunk_kda_bwd_custom directly with a caller-reported state
The mega-path has_initial_state flag tells the kernel whether a recurrence seed exists, not whether the user wants dh0 back. In CP the forward always substitutes the CP-prepared state (merged from the previous rank) for a missing user state, so it must seed the recurrence even when the caller reported has_initial_state=False; passing the caller flag dropped the seed and corrupted remat+bf16 gradients (43.8% mismatch). Keep return_dh0=has_initial_state so the dh0 output still follows the caller's intent. The new initial-state grad test also used a 4-axis out_specs entry for the rank-3 db output and compared the full backward tuple against the staged path, which frees initial_state before its CP collective and never returns dh0. Compare the token gradients only and assert the fused path's dh0 directly.
shard_map keeps the backward tuple's None leaves, so out_specs cannot describe the raw 9-tuple. Return (dq, dk, dv, db, dg, dh0) instead, substituting zeros for the staged path's dh0 (it frees initial_state before its CP collective and never computes one) so both paths share one tree structure.
|
Non-blocking tracking for two openxla#1103 CP follow-ups: combining the first/last-segment metadata collectives and revisiting an |
|
@Fred33146 Thanks for flagging both openxla#1103 follow-ups. I documented their status in the #18 PR description: first_seg and last_seg are still gathered separately in cp_utils.py; the combined dS_ext/dM backward-data gather is unrelated. _merge_initial_state and _merge_dht still use lax.fori_loop. Neither a combined metadata gather nor associative_scan has been A/B benchmarked or fully validated here, so both are explicitly deferred from this migration, not described as completed. A follow-up should check CP=2/4, saved/remat, single/multi-segment numerical behavior, compile behavior, and representative forward+VJP latency before changing the recurrence. |
|
@Fred33146 Follow-up to the two openxla#1103 CP optimization questions (#issuecomment-5884256556): I tested a combined first_seg/last_seg metadata gather on the TPU; CP fused + megakernel correctness passed 34/34. End-to-end CP4 forward+VJP A/B was mixed: global T=512 megakernel 1.980 vs 2.022 ms (combined vs original, one run), while T=4096 was 2.522/2.555 vs 2.503/2.515 ms (two runs). With no consistent benefit, I reverted that candidate. I also prototyped associative_scan for the forward/backward affine merges: CP4 microbenchmarks were approximately tied with fori_loop, compile time increased, and FP32 reassociation produced about 1.2e-4 max-absolute difference on synthetic input. I did not land scan without a full numerical matrix or performance win. The #18 description now records the exact evaluation and scope; no CP kernel change was pushed. |
…da-mega-cp-backward-v2
Summary
PR C of the KDA megakernel migration, stacked on PR #20. Adds CP local backward fusion and the complete CP backward megakernel with Tokamax numerical tests and CI shard coverage.
Stack:
main-> PR #19 -> PR #20 -> this PR.Incremental size: 3,118 additions, 224 deletions across 6 files after the CP review regressions were fixed. The shared benchmark script remains in the #19 base and is not part of this incremental diff.
TPU validation
Environment: TPU v7x 4-chip host (2x2x1), 8 visible JAX TPU devices, JAX 0.11.2 with libtpu, Python 3.12.
Final result: 30 passed in 131.69s (2:11).
Findings and conclusion
3.42e-3fused-vs-staged accumulated error.The production BF16 CP megakernel is validated for CP sizes 2 and 4, saved/rematerialized policies, and single/multiple-segment inputs. FP32 preserves correctness through staged fallback without affecting the BF16 performance path.
TPU performance versus openxla/tokamax openxla#1103 (2026-09-23)
Baseline: openxla/tokamax #1103 pinned at
939da5c24b90dca4c73ad50fba9b9021d9f35a27, using its default CP forward + custom VJP. Both checkouts ran on the same TPU v7x Pod and JAX 0.11.2. This measures forward + VJP together, not backward-only time. Inputs: CP=4, BF16, H=8, B=2, K=V=128, one segment spanning the ranks. The benchmark script is inherited from PR #19; copy it into the pinned openxla#1103 checkout and use--baseline.Each case is independently compiled and timed with synchronized wall-clock iterations. T=512 was measured in three alternating 30-iteration runs; T=4096 in three alternating 20-iteration runs. The table uses the median of each run's median for those two shapes. T=2048 and T=8192 are single 20-iteration runs.
At T=4096 all three paired runs favored the megakernel: openxla#1103 run medians were 2.723/2.737/2.663 ms versus 2.361/2.496/2.500 ms. The short T=512 benchmark has only two 64-token chunks per rank, below the six-chunk DMA window; it cannot exercise cross-window ping-pong overlap. At T=4096 there are 16 chunks per rank and three windows. This is a plausible reason the optimizations become visible at longer lengths, not a phase-level profiling attribution. The separate
fuse_cp_backward, rematerialization and multi-segment packed cases are not covered by this sweep.Conclusion: the CP megakernel is near parity for the short case and shows a reproducible forward+VJP latency benefit at T=4096 for this shape; T=8192 is promising but has only one run. Estimated peak memory is lower throughout. This is one-host CP4, not multi-host throughput. FP32 uses the staged fallback. Compilation and host input preparation are excluded; the Pod's temporary
qwixshim is not committed.2026-09-24 review follow-up (TPU v7x)
f99f47a: load the beta tile before reshaping, instead of passing a squeezed Ref view to the inlined kernel.memref_squeezefailures and 77 passes; all 39 cases in the failing varlen/segment-fused family were rerun on the fixed code and passed (39/39, 432.23s).f99f47a:pytest -q tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_kernel_test.py --maxfail=1 --tb=short-> 297/297 passed in 5383.49s (1:29:43), exit code 0. All tests used JAX 0.11.2 with libtpu on the 4-chip TPU v7x Pod.2026-09-29 review follow-up: openxla#1103 CP optimization ideas
cp_utils.pystill gathersfirst_segandlast_segseparately. This PR's combined backward-data gather fordS_ext/dMis a different optimization; it does not implement or benchmark a combined metadata gather._merge_initial_stateand_merge_dhtstill usejax.lax.fori_loop. Anassociative_scanformulation has not been benchmarked or validated for recurrence order and numerical behavior.2026-09-29 CP metadata/scan evaluation
We prototyped a single
all_gather_into_tensorof[first_seg, last_seg]in place of the two scalar gathers. It passed the CP fused + megakernel TPU suites (34/34 in 148.98s), but forward+VJP A/B on the same CP4 TPU Pod did not show a consistent win. At global T=512, one 30-iteration run measured megakernel 1.980 ms (combined) vs 2.022 ms (original). At global T=4096, two runs measured 2.522/2.555 ms (combined) vs 2.503/2.515 ms (original). Peak memory was effectively unchanged. The candidate was reverted rather than landed as a blanket optimization.A separate CP4/H=8/B=2/K=V=128 FP32 affine-composition
associative_scanprototype was compared with the existingfori_loopmerge helpers. Forward and backward single-device microbenchmarks were approximately tied (0.1559 vs 0.1556 ms; 0.1586 vs 0.1566 ms, scan vs loop), while scan compilation was slower and reassociation changed results by about 1.2e-4 max absolute on the synthetic input. This prototype was not put into the CP training path or subjected to the full CP numerical matrix. With no clear latency gain and altered FP32 reduction order, the existing loops remain in place. These measurements are probes, not a claim about multi-host CP throughput.Commit
51a4139only synchronizes the benchmark script's shape parameters from #19; it does not change the CP implementation.2026-09-29 stacked-HEAD TPU verification
On HEAD
00ebb27, the CP fused and CP megakernel Tokamax suites passed 34/34 in 146.12s on the TPU v7x Pod. The merged base adds #19 review tests and the benchmark-shape script but no CP kernel change. The earlier 297/297 full KDA sweep was run onf99f47a; it was not rerun on this benchmark/test-only stack update.