Context Parallelism in MegaDLMs: How It Handles Long Sequences
Context Parallelism in MegaDLMs partitions input sequences across multiple GPUs to enable training and inference on sequences far longer than a single GPU can accommodate, while maintaining full attention computation through strategic key-value communication.
Context Parallelism (CP) is a model-parallelism dimension introduced in MegaDLMs (Megatron-based Deep Language Models) that splits the input sequence rather than tensor dimensions across GPU groups. This approach, implemented in the jinjieni/megadlms repository, allows the framework to scale to extremely long token sequences by distributing memory load while preserving computational correctness through specialized communication patterns.
What Is Context Parallelism in MegaDLMs?
Unlike Tensor Parallelism which splits weight matrices or Pipeline Parallelism which splits layers, Context Parallelism divides the sequence length dimension across participating GPUs. In megatron/core/parallel_state.py, the initialize_model_parallel function accepts a context_parallel_size parameter (defaulting to 1) and creates context-parallel groups that partition the sequence among available GPUs【/cache/repos/github.com/jinjieni/megadlms/main/megatron/core/parallel_state.py#L90-L60】.
Each GPU in a context-parallel group processes only a slice of the full token stream, reducing per-GPU activation memory from O(seq_len) to O(seq_len / cp_size). The full sequence context is reconstructed on-the-fly during attention computation through efficient GPU-to-GPU communication.
How Context Parallelism Handles Long Sequences
Sequence Partitioning Across GPUs
The fundamental mechanism for handling long sequences begins in megatron/training/initialize.py, where the input shape is explicitly divided by args.context_parallel_size【/cache/repos/github.com/jinjieni/megadlms/main/megatron/training/initialize.py#L244-L247】. This means if you configure a global sequence length of 8192 tokens with context_parallel_size=4, each GPU processes only 2048 tokens locally.
The effective per-GPU sequence length calculation follows:
from megatron.core.parallel_state import get_context_parallel_world_size
from megatron.core.transformer.transformer_config import TransformerConfig
cfg = TransformerConfig(seq_length=8192, context_parallel_size=4)
per_gpu_seq_len = cfg.seq_length // cfg.context_parallel_size
print(f"Per-GPU sequence length: {per_gpu_seq_len}") # Output: 2048
Key-Value Communication Patterns
To maintain full attention computation while sequences are partitioned, MegaDLMs implements specialized communication for exchanging key (K) and value (V) tensors. The TransformerConfig class in megatron/core/transformer/transformer_config.py stores the cp_comm_type field, supporting four communication schemes: p2p, all_gather, a2a (all-to-all), and a2a+p2p【/cache/repos/github.com/jinjieni/megadlms/main/megatron/core/transformer/transformer_config.py#L358-L368】.
During layer construction in megatron/core/transformer/transformer_layer.py, the chosen CP communication type is forwarded to the attention modules【/cache/repos/github.com/jinjieni/megadlms/main/megatron/core/transformer/transformer_layer.py#L15-L20】. Before attention computation, each GPU exchanges KV chunks belonging to other sequence partitions using the selected pattern, reconstructing the full-length context needed for self-attention while keeping memory usage distributed.
Important implementation note: The standard DotProductAttention in megatron/core/transformer/dot_product_attention.py explicitly asserts context_parallel_size == 1, meaning only the Transformer Engine (TE) based attention variant (TEDotProductAttention) supports Context Parallelism【/cache/repos/github.com/jinjieni/megadlms/main/megatron/core/transformer/dot_product_attention.py#L50-L53】.
Hierarchical Context Parallelism
For extremely long sequences that exceed single-node memory or bandwidth capacities, MegaDLMs supports hierarchical partitioning via hierarchical_context_parallel_sizes. This parameter, parsed in megatron/training/arguments.py, accepts a list such as [2,2] to create a 4-way split across two levels【/cache/repos/github.com/jinjieni/megadlms/main/megatron/training/arguments.py#L255-L259】.
This hierarchical approach enables multi-level splitting strategies, such as using high-bandwidth NVLink for intra-node context parallelism and InfiniBand for inter-node communication, optimizing the latency-bandwidth tradeoff for very long sequences.
Configuring Context Parallelism in MegaDLMs
To enable Context Parallelism in your MegaDLM training runs, specify the parallelism degree and communication type via command-line arguments:
python pretrain_difflm.py \
--tensor-model-parallel-size 4 \
--pipeline-model-parallel-size 1 \
--context-parallel-size 4 \
--cp-comm-type p2p \
--seq-length 8192 \
--use-te-distributed-activation-checkpointing \
... other args ...
For hierarchical configurations:
python pretrain_difflm.py \
--context-parallel-size 4 \
--hierarchical-context-parallel-sizes 2 2 \
--cp-comm-type a2a \
...
When building custom models programmatically, ensure you use the Transformer Engine-based attention and configure the TransformerConfig correctly:
from megatron.core.transformer.transformer_layer import TransformerLayer
from megatron.core.transformer.transformer_config import TransformerConfig
cfg = TransformerConfig(
hidden_size=4096,
num_attention_heads=32,
num_layers=24,
context_parallel_size=2, # Split sequence into 2 parts
cp_comm_type="all_gather", # Use all-gather for KV exchange
use_te=True, # Required for CP support
)
# The CP type is automatically passed to attention blocks
layer = TransformerLayer(
submodules=my_submodules,
layer_number=0,
config=cfg,
)
Key Implementation Files
The Context Parallelism mechanism in MegaDLMs spans several core modules:
| File | Role in Context Parallelism |
|---|---|
megatron/core/parallel_state.py |
Creates CP groups via initialize_model_parallel and stores context_parallel_size【/cache/repos/github.com/jinjieni/megadlms/main/megatron/core/parallel_state.py#L90-L60】 |
megatron/core/transformer/transformer_config.py |
Defines cp_comm_type field with supported communication schemes【/cache/repos/github.com/jinjieni/megadlms/main/megatron/core/transformer/transformer_config.py#L358-L368】 |
megatron/core/transformer/transformer_layer.py |
Forwards CP settings to self- and cross-attention modules【/cache/repos/github.com/jinjieni/megadlms/main/megatron/core/transformer/transformer_layer.py#L15-L20】 |
megatron/core/transformer/dot_product_attention.py |
Standard attention asserts context_parallel_size == 1; only TE variant supports CP【/cache/repos/github.com/jinjieni/megadlms/main/megatron/core/transformer/dot_product_attention.py#L50-L53】 |
megatron/training/arguments.py |
Parses CLI arguments for --context-parallel-size, --cp-comm-type, and hierarchical sizes【/cache/repos/github.com/jinjieni/megadlms/main/megatron/training/arguments.py#L255-L259】 |
megatron/training/initialize.py |
Divides input shapes by context_parallel_size to establish per-GPU sequence length【/cache/repos/github.com/jinjieni/megadlms/main/megatron/training/initialize.py#L244-L247】 |
Summary
- Context Parallelism in MegaDLMs partitions input sequences across GPU groups rather than model weights, enabling training on sequences longer than single-GPU memory limits.
- The
context_parallel_sizeparameter ininitialize_model_parallelcreates CP groups that divide sequence length, reducing per-GPU activation memory by a factor ofcp_size. - Communication patterns (
p2p,all_gather,a2a,a2a+p2p) exchange KV tensors between GPUs before attention computation, allowing each GPU to reconstruct full context while storing only local sequence slices. - Hierarchical CP supports multi-level splitting (e.g.,
[2,2]) to optimize bandwidth across NVLink and InfiniBand hierarchies for extremely long sequences. - Transformer Engine requirement: Only
TEDotProductAttentionsupports CP; the standardDotProductAttentionexplicitly requirescontext_parallel_size == 1.
Frequently Asked Questions
What is the difference between Context Parallelism and Tensor Parallelism in MegaDLMs?
Tensor Parallelism splits individual weight matrices and activations across GPUs (typically partitioning hidden dimensions), while Context Parallelism splits the input sequence length across GPUs. In MegaDLMs, Tensor Parallelism requires all GPUs to hold the full sequence, whereas Context Parallelism allows each GPU to store only seq_len / cp_size tokens, enabling much longer sequences at the cost of additional KV communication between GPUs.
How does Context Parallelism affect memory usage during training?
Context Parallelism reduces activation memory linearly with the parallelism degree. When using context_parallel_size=4, each GPU stores only 25% of the sequence activations, reducing memory from O(seq_len) to O(seq_len / cp_size). However, the trade-off involves temporary memory spikes during the KV exchange phase (depending on cp_comm_type), though the net memory savings typically enable training on sequences 2-8× longer than otherwise possible.
Which communication type should I use for Context Parallelism?
The optimal cp_comm_type depends on your hardware topology and sequence length. P2P (point-to-point) works well for small CP sizes within a single node. All_gather is suitable when bandwidth permits full KV replication. A2a (all-to-all) minimizes data movement for large CP groups across nodes. A2a+p2p combines both for hierarchical setups. For NVLink-connected GPUs, p2p typically offers lowest latency, while a2a optimizes inter-node InfiniBand utilization.
Can Context Parallelism be combined with other parallelism strategies?
Yes, Context Parallelism is designed to compose with Tensor, Pipeline, and Data Parallelism in MegaDLMs. The initialize_model_parallel function in megatron/core/parallel_state.py creates orthogonal process groups for each parallelism dimension, allowing configurations like 4-way Tensor Parallelism combined with 2-way Context Parallelism. When combining strategies, the total GPU count must satisfy tp_size × cp_size × pp_size × dp_size = total_gpus, and hierarchical CP can further optimize communication across different network fabrics.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →