Skip to content

[megatron] KDA context parallelism: opt-in all-gather exchange (SKYRL_KDA_CP_EXCHANGE=allgather) - #2437

Draft
avigyabb wants to merge 2 commits into
glm5p3-flash/context-parallelfrom
glm5p3-flash/kda-cp-allgather-exchange
Draft

avigyabb wants to merge 2 commits into
glm5p3-flash/context-parallelfrom
glm5p3-flash/kda-cp-allgather-exchange

Conversation

@avigyabb

@avigyabb avigyabb commented Oct 7, 2026 •

Copy link
Copy Markdown
Collaborator

Stacked on #2418 (KDA context parallelism, all-to-all exchange). Review only the top commit.

Summary

Under context parallelism, KDA moves from sequence to head parallelism. #2418 does that with megatron-core's GatedDeltaNet all-to-all of the projected q/k/v/f/g/beta, which are ~1.25x wider than the hidden state and run on TP-gathered rows.

SKYRL_KDA_CP_EXCHANGE=allgather (opt-in) instead all-gathers the sequence-parallel hidden-state shards over CP (the only cross-node hop when CP spans nodes), then over TP, and projects only this rank's 1/cp slice of the local heads. One permutation of the narrow projected tensors puts tokens into natural order; backward reduce-scatters in reverse. About 5x fewer bytes cross the CP group.

Performance

1M tokens/seq, TP8 x CP2, EP32 x ETP2, 64 B200 all-to-all (#2418) all-gather (this PR)
NCCL over TCP fwd+bwd 1482 s 1251 s
NCCL over EFA 335 s/step, 12.5k tok/s 307 s/step, 13.7k tok/s (−8%)

Over EFA the step is 8% faster at the same peak memory (144.6 vs 144.5 GiB); EFA rows are the mean warm step of the full step (fwd+bwd+optimizer) with the rest of the long-context recipe on. On a 4-layer GLM-5.3-Flash slice at 1M tokens, CP2 across 2 nodes over EFA: 34.9 -> 32.8 s/step (−6%), same memory.

Testing

Split out of #2418.

🤖 Generated with Claude Code

…_KDA_CP_EXCHANGE=allgather)

Under CP, KDA moves from sequence to head parallelism with megatron-core's GatedDeltaNet
all-to-all of the projected q/k/v/f/g/beta, which are ~1.25x wider than the hidden state and run
on TP-gathered rows. SKYRL_KDA_CP_EXCHANGE=allgather instead all-gathers the sequence-parallel
hidden-state shards over CP (the only cross-node hop when CP spans nodes), then over TP, and
projects only this rank's 1/cp slice of the local heads: ~5x fewer bytes across the CP group.
Measured over TCP NCCL at 1M tokens, TP8 CP2: fwd+bwd 1482 -> 1251 s.

Tests: test_kda_context_parallel.py covers both exchanges (CP=2 vs CP=1).

Split out of #2418.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
@avigyabb

avigyabb commented Oct 8, 2026

Copy link
Copy Markdown
Collaborator Author

all gather was better in all tests - maybe we should fold into one path? maybe just wait until megatron PR gets pushed?

…llgather-exchange

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01K4BmK8yuJTMbprizvT4Rja

This branch was successfully deployed

1 active deployment
Preview — 7218281f Deployed Oct 8, 2026 by vercel[bot]
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.

1 participant