Skip to content

Add CP KDA backward megakernel - #18

Open
fdz-1999 wants to merge 13 commits into
kda-mega-backward-packedfrom
kda-mega-cp-backward-v2
Open

fdz-1999 wants to merge 13 commits into
kda-mega-backward-packedfrom
kda-mega-cp-backward-v2

Conversation

@fdz-1999

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

Copy link
Copy Markdown
Collaborator

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.

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.py

Final result: 30 passed in 131.69s (2:11).

Findings and conclusion

  • Initial validation found all BF16 megakernel cases passing, while eight FP32 cases showed up to about 3.42e-3 fused-vs-staged accumulated error.
  • The pallas-kernel source CP megakernel and its hardware validation have a BF16 contract. Tokamax now dispatches the CP megakernel only for BF16.
  • FP32 safely retains the established staged CP backward path, and the test verifies that enabling the option does not alter FP32 staged results.
  • BF16 compares both megakernel and staged paths against the independent recurrent XLA reference.

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.

# In the pinned #1103 checkout; omit --baseline on this PR:
for t in 512 2048 4096 8192; do
  python -m tokamax.benchmarks.kda_megakernel_compare --suite cp --cp-tokens "$t" --baseline --iterations 20
done

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.

Global T (per rank) openxla#1103 latency CP megakernel latency Change Estimated peak memory
512 (128) 2.062 ms 2.052 ms ~0.5% lower; within noise 29.9 -> 22.6 MB
2048 (512) 2.235 ms 2.179 ms 2.5% lower 89.8 -> 58.1 MB
4096 (1024) 2.723 ms 2.496 ms 8.3% lower 169.5 -> 108.8 MB
8192 (2048) 3.439 ms 3.079 ms 10.5% lower 329.1 -> 215.6 MB

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 qwix shim is not committed.

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

  • Fixed the saved-state fused backward beta Ref regression in commit f99f47a: load the beta tile before reshaping, instead of passing a squeezed Ref view to the inlined kernel.
  • The previous broad backward sweep was stopped after 39 Mosaic memref_squeeze failures 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).
  • CP fused plus CP megakernel suites passed on the fixed code (34/34, 147.70s), including the initial-state-gradient regression.
  • The complete KDA regression sweep was then rerun on the fixed PR HEAD 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.py still gathers first_seg and last_seg separately. This PR's combined backward-data gather for dS_ext/dM is a different optimization; it does not implement or benchmark a combined metadata gather.
  • _merge_initial_state and _merge_dht still use jax.lax.fori_loop. An associative_scan formulation has not been benchmarked or validated for recurrence order and numerical behavior.
  • Both Implement KDA Pallas Kernel openxla/tokamax#1103 suggestions are intentionally deferred from this CP megakernel migration, rather than presented as completed optimizations. A follow-up should compare the combined metadata gather and a scan formulation against the current code on CP=2/4, saved/rematerialized state, single/multi-segment correctness, compile behavior, and representative forward+VJP latency. The performance numbers above do not isolate or claim improvements from either idea.

2026-09-29 CP metadata/scan evaluation

We prototyped a single all_gather_into_tensor of [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_scan prototype was compared with the existing fori_loop merge 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 51a4139 only 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 on f99f47a; it was not rerun on this benchmark/test-only stack update.

@fdz-1999
fdz-1999 force-pushed the kda-mega-cp-backward-v2 branch from 0a2c989 to 82a7af3 Compare September 22, 2026 16:04
@fdz-1999
fdz-1999 force-pushed the kda-mega-backward-packed branch from bcbf1e4 to e2cf254 Compare September 22, 2026 16:04
@fdz-1999
fdz-1999 force-pushed the kda-mega-cp-backward-v2 branch from 82a7af3 to ae03e55 Compare September 23, 2026 04:10
@fdz-1999

Copy link
Copy Markdown
Collaborator Author

TPU validation result

Environment:

  • TPU v7x 4-chip host, topology 2x2x1
  • 8 visible JAX TPU devices
  • 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_cp_fused_test.py \
  tokamax/_src/ops/experimental/kda/pallas_mosaic_tpu_cp_megakernel_test.py

Final result: 30 passed in 131.69s (2:11).

Findings and fixes made during validation:

  • Initial validation found all BF16 megakernel cases passing, while eight FP32 cases showed up to about 3.42e-3 fused-vs-staged accumulated error.
  • The pallas-kernel source CP megakernel and its hardware validation have a BF16 contract. Tokamax now dispatches the CP megakernel only for BF16.
  • FP32 requests safely retain the established staged CP backward path. The test verifies that enabling the option does not alter FP32 staged results.
  • BF16 continues to compare the megakernel and staged paths against the independent recurrent XLA reference.
  • After the dispatch and test-contract fix, the complete CP local-fusion and CP megakernel suite passed.

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.

@fdz-1999
fdz-1999 force-pushed the kda-mega-backward-packed branch from 81a6653 to 323337a Compare September 23, 2026 07:35
@fdz-1999
fdz-1999 force-pushed the kda-mega-cp-backward-v2 branch from 809ed3e to e2ec8d5 Compare September 23, 2026 07:35
@fdz-1999
fdz-1999 force-pushed the kda-mega-backward-packed branch from 323337a to 0df3505 Compare September 23, 2026 07:37
@fdz-1999
fdz-1999 force-pushed the kda-mega-cp-backward-v2 branch from e2ec8d5 to 97935ac Compare September 23, 2026 07:37
@fdz-1999
fdz-1999 force-pushed the kda-mega-backward-packed branch from 0df3505 to fe6db93 Compare September 23, 2026 10:09
@fdz-1999
fdz-1999 force-pushed the kda-mega-cp-backward-v2 branch from 97935ac to 90cebad Compare September 23, 2026 10:09
@fdz-1999
fdz-1999 force-pushed the kda-mega-backward-packed branch from fe6db93 to 3d5ba6a Compare September 23, 2026 11:02
fdz-1999 and others added 6 commits September 23, 2026 19:03
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>
@fdz-1999
fdz-1999 force-pushed the kda-mega-cp-backward-v2 branch from 90cebad to bc1eb4d Compare September 23, 2026 11:03
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,

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.

[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.

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.

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],

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.

[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.

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.

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

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.

[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.

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.

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.
@fdz-1999
fdz-1999 marked this pull request as ready for review September 28, 2026 08:08
@Fred33146

Copy link
Copy Markdown
Collaborator

Non-blocking tracking for two openxla#1103 CP follow-ups: combining the first/last-segment metadata collectives and revisiting an associative_scan formulation. This PR combines a different backward-data gather (dS_ext/dM), but cp_utils.py still performs separate all_gather_into_tensor calls for first_seg and last_seg, and _merge_initial_state / _merge_dht still use jax.lax.fori_loop. Were the metadata-gather and scan options evaluated and intentionally deferred? Please document the outcome or link a follow-up. The upstream wording was to consider/revisit these optimizations, so I am not treating their absence as a correctness or merge blocker for this CP megakernel PR.

@fdz-1999

Copy link
Copy Markdown
Collaborator Author

@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.

@fdz-1999

Copy link
Copy Markdown
Collaborator Author

@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.

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