ReferenceFull Library ReferenceNemo AutomodelNemo AutomodelComponentsDistributednemo_automodel.components.distributed.context_parallel

nemo_automodel.components.distributed.context_parallel

View as Markdown

Context-parallel batch sharding.

Submodules

Package Contents

Classes

NameDescription
ContextParallelSharderCP backend description: how a batch is sharded and where local tokens live.

API

class nemo_automodel.components.distributed.context_parallel.sharder.ContextParallelSharder(
model: torch.nn.Module | None = None,
device_mesh: torch.distributed.device_mesh.DeviceMesh | None = None,
batch: dict[str, typing.Any] | None = None,
shard_batch: collections.abc.Callable[..., tuple[collections.abc.Callable, dict[str, typing.Any], 'ShardLayout | None']] | None = None,
local_token_global_indices: collections.abc.Callable[..., torch.Tensor] | None = None,
shard_layout: 'ShardLayout | None' = None,
padding_token_id: int = 0,
num_chunks: int = 1,
loss_mask: torch.Tensor | None = None,
invoke_pre_embed: bool = True,
extra_seq_buffers: dict[str, int] | None = None
)

CP backend description: how a batch is sharded and where local tokens live.

_cp_mesh
Any = resolved._cp_mesh
_loss_mask
Tensor | None = resolved._loss_mask
_padding_token_id
int = resolved._padding_token_id
_tp_mesh
Any = resolved._tp_mesh
local_token_global_indices
Callable[..., Tensor] | None = resolved.local_token_global_indices

(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).

shard_batch
Callable[..., tuple[Callable, dict[str, Any], 'ShardLayout | None']] = resolved.shard_batch

(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.

shard_layout
'ShardLayout | None' = resolved.shard_layout

The ShardLayout of the last shard_batch, set by shard. Sharders are built per resolution/hook call, so the layout never leaks across steps.

nemo_automodel.components.distributed.context_parallel.sharder.ContextParallelSharder._indices(
padded_seq_len: int,
device
) -> torch.Tensor
nemo_automodel.components.distributed.context_parallel.sharder.ContextParallelSharder.gather_token_tensor(
tensor: torch.Tensor,
seq_dim: int = 1,
trim: bool = False,
fill: float | int | None = None
) -> torch.Tensor

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).

nemo_automodel.components.distributed.context_parallel.sharder.ContextParallelSharder.shard(
batch: dict[str, typing.Any]
) -> tuple[collections.abc.Callable, dict[str, typing.Any]]

Shard a batch and retain its layout for token-aligned tensors.

nemo_automodel.components.distributed.context_parallel.sharder.ContextParallelSharder.shard_token_tensor(
tensor: torch.Tensor,
seq_dim: int = 1,
fill: float | int | None = None
) -> torch.Tensor

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, fill filling the pad slots;
  • tensors matching the reported pre-flatten input_row_shape on a flat-stream (THD) layout are flattened first (the returned shard is in the model’s local stream coordinate);
  • tensors of original_seq_len are right-padded to padded_seq_len with the explicit fill value;
  • tensors already at padded_seq_len shard directly.

Any other length raises instead of silently sharding the wrong slice.