nemo_automodel.components.distributed.blockdiag_cp

View as Markdown

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

Package Contents

Classes

NameDescription
BlockdiagCpModelStateTyped model-facing view of an active block-diagonal CP step.

Functions

NameDescription
attach_cp1_packed_varlen_hooksScope the cp1 packed varlen SDPA patch to the model’s attention forwards.
configure_cp_varlenConfigure the block-diagonal CP attention path from parsed runtime config.
cp1_packed_varlen_backendThe configured cp1 packed varlen backend (‘te’/‘flash’), or None if disabled (dense).
cp_blockdiag_sdpaBlock-diagonal context-parallel SDPA.
current_blockdiag_cp_stateReturn the active block-diagonal CP step state, or None.
disable_cp1_packed_varlenDisarm stale cp1 state before starting a new outer model forward.
enable_cp1_packed_varlenArm cp1 packed block-diagonal varlen for the rest of this step.
make_cp_blockdiag_batch_and_ctxSequentially shard a batch for block-diagonal CP.

API

class nemo_automodel.components.distributed.blockdiag_cp.state.BlockdiagCpModelState(
group: torch.distributed.ProcessGroup,
packed_cu_seqlens: torch.Tensor,
packed_cu_seqlens_cpu: torch.Tensor
)
Dataclass

Typed model-facing view of an active block-diagonal CP step.

Parameters:

group
ProcessGroup

Process group that shards the sequence across context-parallel ranks.

packed_cu_seqlens
Tensor

Global packed-document boundaries of shape [num_documents + 1] on the compute device.

packed_cu_seqlens_cpu
Tensor

CPU copy of packed_cu_seqlens with shape [num_documents + 1] for FLA’s host-side CP planning.

group
ProcessGroup
packed_cu_seqlens
Tensor
packed_cu_seqlens_cpu
Tensor
nemo_automodel.components.distributed.blockdiag_cp.packed.attach_cp1_packed_varlen_hooks(
model: torch.nn.Module
) -> None

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:

model
torch.nn.Module

The model whose self_attn submodules route softmax attention through F.scaled_dot_product_attention.

nemo_automodel.components.distributed.blockdiag_cp.state.configure_cp_varlen(
attn_backend: str = 'flash',
kv_exchange: str = 'allgather'
) -> None

Configure the block-diagonal CP attention path from parsed runtime config.

Parameters:

attn_backend
strDefaults to 'flash'

Varlen kernel selection; any synonym accepted by normalize_attn_backend (default "flash").

kv_exchange
strDefaults to 'allgather'

K/V delivery mode; any synonym accepted by normalize_kv_exchange (default "allgather").

nemo_automodel.components.distributed.blockdiag_cp.packed.cp1_packed_varlen_backend() -> str | None

The configured cp1 packed varlen backend (‘te’/‘flash’), or None if disabled (dense).

nemo_automodel.components.distributed.blockdiag_cp.runtime.cp_blockdiag_sdpa(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_mask: torch.Tensor | None = None,
dropout_p: float = 0.0,
is_causal: bool = False,
scale: float | None = None,
enable_gqa: bool = False,
kwargs = {}
) -> torch.Tensor

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:

query
torch.Tensor

This rank’s LOCAL query shard [B, Hq, L, D] (B = batch, Hq = query heads, L = local sequence length, D = head dim).

key
torch.Tensor

This rank’s LOCAL key shard [B, Hkv, L, D].

value
torch.Tensor

This rank’s LOCAL value shard [B, Hkv, L, D].

attn_mask
torch.Tensor | NoneDefaults to None

Ignored on the CP path (forwarded to stock SDPA otherwise).

dropout_p
floatDefaults to 0.0

Dropout probability.

is_causal
boolDefaults to False

Ignored on the CP path (forwarded to stock SDPA otherwise).

scale
float | NoneDefaults to None

Softmax scale (None -> D**-0.5).

enable_gqa
boolDefaults to False

Grouped-query attention flag as passed by HF’s sdpa path.

**kwargs
Defaults to {}

Ignored; accepted for SDPA signature compatibility.

Returns: torch.Tensor

Attention output [B, Hq, L, D] for this rank’s local rows.

nemo_automodel.components.distributed.blockdiag_cp.state.current_blockdiag_cp_state() -> nemo_automodel.components.distributed.blockdiag_cp.state.BlockdiagCpModelState | None

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

nemo_automodel.components.distributed.blockdiag_cp.packed.disable_cp1_packed_varlen() -> None

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.

nemo_automodel.components.distributed.blockdiag_cp.packed.enable_cp1_packed_varlen(
doc_ids: torch.Tensor,
backend: str
) -> None

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:

doc_ids
torch.Tensor

Per-position document ids [1, S] or [S] (0 == padding) over the full packed sequence.

backend
str

Varlen kernel backend, "flash" or "te".

nemo_automodel.components.distributed.blockdiag_cp.batch.make_cp_blockdiag_batch_and_ctx(
cp_mesh: torch.distributed.device_mesh.DeviceMesh,
tp_mesh: torch.distributed.device_mesh.DeviceMesh | None,
batch: dict[str, typing.Any],
loss_mask: torch.Tensor | None = None,
padding_token_id: int = 0,
shard_primary: bool = True
) -> tuple[typing.Callable[[], typing.ContextManager], dict[str, typing.Any], nemo_automodel.components.distributed.context_parallel.sharder.ShardLayout | None]

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:

cp_mesh
DeviceMesh

The context-parallel device (sub)mesh.

tp_mesh
DeviceMesh | None

Accepted for the shared sharder signature; unused (block-diagonal CP shards only the sequence dimension).

batch
dict[str, Any]

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.

loss_mask
torch.Tensor | NoneDefaults to None

Optional per-token loss mask [B, S]; padded with 0 and sharded like the other sequence-aligned tensors (stored back into batch["loss_mask"]).

padding_token_id
intDefaults to 0

Fill value for input_ids padding when the primary stream is sharded here.

shard_primary
boolDefaults to True

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