nemo_automodel.components.distributed.context_parallel
nemo_automodel.components.distributed.context_parallel
Context-parallel batch sharding.
Submodules
nemo_automodel.components.distributed.context_parallel.maginemo_automodel.components.distributed.context_parallel.mambanemo_automodel.components.distributed.context_parallel.shardernemo_automodel.components.distributed.context_parallel.utils
Package Contents
Classes
API
CP backend description: how a batch is sharded and where local tokens live.
(cp_mesh, padded_seq_len, device) -> LongTensor with the global position of each local token —
closed-form for contiguous/round-robin layouts; None for
data-dependent layouts, whose partition arrives with
shard_layout (their token verbs raise before the first shard).
(cp_mesh, tp_mesh, batch, *, loss_mask=None, padding_token_id=0) -> (ctx_factory, batch, ShardLayout | None).
Pads and shards the batch, installs any backend-owned attention
transport, and reports the shard layout it computed; shard
stores it as shard_layout for the token verbs.
The ShardLayout of the last shard_batch, set
by shard. Sharders are built per resolution/hook call, so
the layout never leaks across steps.
Differentiably gather a token-aligned local shard to global order.
With trim=True the result is returned in the caller’s original
coordinates using the shard layout: sliced back to original_seq_len,
un-flattened to input_row_shape (THD), or mapped through the
reported position map (fill for input positions whose tokens were
dropped, e.g. re-padded pack slots). Raises when no layout is present
(nothing to trim to).
Shard a batch and retain its layout for token-aligned tensors.
Shard a full-length token-aligned tensor exactly like the model inputs.
When shard layout are present (after the first shard_batch), the
caller may pass tensors in its own coordinates and the verb applies the
same transform the batch went through:
[B, S_in]tensors on a repositioned-row layout (reported position map, e.g. DSV4 packed repad) are scattered into the padded rows,fillfilling the pad slots;- tensors matching the reported pre-flatten
input_row_shapeon a flat-stream (THD) layout are flattened first (the returned shard is in the model’s local stream coordinate); - tensors of
original_seq_lenare right-padded topadded_seq_lenwith the explicitfillvalue; - tensors already at
padded_seq_lenshard directly.
Any other length raises instead of silently sharding the wrong slice.