nemo_automodel.components.distributed.blockdiag_cp
nemo_automodel.components.distributed.blockdiag_cp
Block-diagonal (per-document) varlen context parallelism for packed sequences.
Packed-sequence training concatenates many documents into one long sequence; a
correct attention mask is block-causal per document. The stock DTensor
context_parallel path assumes a single causal document, so packed VLM/LLM
batches need a CP implementation that reshards by contiguous chunks and rebuilds
per-document masking on every rank.
This package provides that implementation, split by responsibility:
.state— runtime knob normalization + activation-checkpoint-safe step state..kernels— dense-mask and varlen (FlashAttention / TransformerEngine) kernels..exchange— differentiable K/V collectives (all-gather, left-halo, all-to-all-v)..runtime— the SDPA entry point and collective-safe path selection..batch— batch padding/sharding and the per-step train context..packed— the cp_size==1 packed-sequence varlen SDPA hook.
Integration follows the model-owned CP convention of
nemo_automodel.components.distributed.cp_utils.make_cp_batch_and_ctx: a model
returns a
ContextParallelSharder
whose batch verb is make_cp_blockdiag_batch_and_ctx, and routes its softmax
attention through cp_blockdiag_sdpa while the returned context is active.
For cp_size == 1 packed runs, the model scopes the varlen SDPA patch to its
attention forwards with attach_cp1_packed_varlen_hooks and arms the
per-forward state with enable_cp1_packed_varlen /
disable_cp1_packed_varlen.
Only these integration entry points are exported; everything else (knob normalization, varlen metadata precompute, fire counters, kernels) is an internal detail of the package’s modules. Model wiring lands in follow-up PRs.
Submodules
nemo_automodel.components.distributed.blockdiag_cp.batchnemo_automodel.components.distributed.blockdiag_cp.exchangenemo_automodel.components.distributed.blockdiag_cp.kernelsnemo_automodel.components.distributed.blockdiag_cp.packednemo_automodel.components.distributed.blockdiag_cp.runtimenemo_automodel.components.distributed.blockdiag_cp.state
Package Contents
Classes
Functions
API
Typed model-facing view of an active block-diagonal CP step.
Parameters:
Process group that shards the sequence across context-parallel ranks.
Global packed-document boundaries of shape
[num_documents + 1] on the compute device.
CPU copy of packed_cu_seqlens with shape
[num_documents + 1] for FLA’s host-side CP planning.
Scope the cp1 packed varlen SDPA patch to the model’s attention forwards.
Registers a forward pre-hook / post-hook pair on every self_attn module
that installs _packed_varlen_sdpa as F.scaled_dot_product_attention
for the duration of that module’s forward and restores stock SDPA afterwards
(always_call=True, so a raising forward cannot leak the patch). Hooks are
attached to the checkpoint-wrapped INNER module because CheckpointWrapper’s
recompute bypasses __call__ on the wrapper — this is what keeps the
varlen path active during activation-checkpointing recompute in backward.
Same bounded-patch pattern as
nemo_automodel.components.distributed.cp_utils.attach_cp_sdpa_hooks.
Outside these hooks, F.scaled_dot_product_attention is untouched.
Parameters:
The model whose self_attn submodules route softmax attention
through F.scaled_dot_product_attention.
Configure the block-diagonal CP attention path from parsed runtime config.
Parameters:
Varlen kernel selection; any synonym accepted by
normalize_attn_backend (default "flash").
K/V delivery mode; any synonym accepted by
normalize_kv_exchange (default "allgather").
The configured cp1 packed varlen backend (‘te’/‘flash’), or None if disabled (dense).
Block-diagonal context-parallel SDPA.
Drop-in replacement for torch.nn.functional.scaled_dot_product_attention
while the context returned by
make_cp_blockdiag_batch_and_ctx
is active; a plain pass-through to stock SDPA otherwise. K/V are exchanged
across the CP group (all-gather, or a needed-only halo/a2a exchange) and one
local attention runs the local queries against the delivered keys with a
per-document causal mask. The passed attn_mask / is_causal are ignored
on the CP path — masking is rebuilt from the document ids so packed sequences
never attend across document boundaries.
Parameters:
This rank’s LOCAL query shard [B, Hq, L, D] (B = batch,
Hq = query heads, L = local sequence length, D = head dim).
This rank’s LOCAL key shard [B, Hkv, L, D].
This rank’s LOCAL value shard [B, Hkv, L, D].
Ignored on the CP path (forwarded to stock SDPA otherwise).
Dropout probability.
Ignored on the CP path (forwarded to stock SDPA otherwise).
Softmax scale (None -> D**-0.5).
Grouped-query attention flag as passed by HF’s sdpa path.
Ignored; accepted for SDPA signature compatibility.
Returns: torch.Tensor
Attention output [B, Hq, L, D] for this rank’s local rows.
Return the active block-diagonal CP step state, or None.
The state is published by
make_cp_blockdiag_batch_and_ctx
for the duration of forward and backward. Model-owned recurrent-attention
implementations use it to share the same contiguous layout and packed
document boundaries as the softmax-attention transport.
Returns: BlockdiagCpModelState | None
The immutable model-facing state, whose packed-boundary tensors have
Disarm stale cp1 state before starting a new outer model forward.
Decoder-layer activation-checkpoint recomputation happens before the next outer forward, so the preceding step’s state remains available for backward and is cleared before vision or any other attention in the next batch runs.
Arm cp1 packed block-diagonal varlen for the rest of this step.
Sets the per-forward doc_ids/backend state read by the SDPA patch that
attach_cp1_packed_varlen_hooks scopes to the attention forwards. The
state remains armed through backward’s activation-checkpoint recomputation;
the next outer model forward clears it via disable_cp1_packed_varlen.
Parameters:
Per-position document ids [1, S] or [S] (0 == padding)
over the full packed sequence.
Varlen kernel backend, "flash" or "te".
Sequentially shard a batch for block-diagonal CP.
Pads the sequence to a multiple of twice the CP world size, slices each
selected sequence-aligned tensor to this rank’s contiguous chunk, and
returns a context whose lifetime activates per-document CP SDPA state.
shard_primary=False leaves token ids or embeddings untouched for models
that embed multimodal inputs inside forward.
Softmax attention must route through
cp_blockdiag_sdpa
while this context is active. A model opts in by returning a
ContextParallelSharder
whose batch verb is this callable.
Parameters:
The context-parallel device (sub)mesh.
Accepted for the shared sharder signature; unused (block-diagonal CP shards only the sequence dimension).
The training batch. Contains exactly one primary stream:
inputs_embeds of shape [batch, sequence, hidden] or
input_ids of shape [batch, sequence]. The batch is mutated in
place: attention_mask is dropped, padding_mask of shape
[batch, sequence] (bool, True == pad) is added, and auxiliary
sequence-aligned tensors are padded then sliced to this rank’s
[row_offset, row_offset + sequence/cp) chunk.
Optional per-token loss mask [B, S]; padded with 0 and
sharded like the other sequence-aligned tensors (stored back into
batch["loss_mask"]).
Fill value for input_ids padding when the primary
stream is sharded here.
Whether to pad and shard the primary stream. Leave False
for models that embed multimodal inputs and shard the resulting
embeddings inside forward so FSDP hooks own the vision/embedding
parameter lifecycle.
Returns: Callable[[], ContextManager]
(train_ctx, batch, layout): a zero-arg callable returning the