nemo_automodel.recipes.kd_utils
nemo_automodel.recipes.kd_utils
Distributed topology and tensor transport helpers for KD recipes.
Module Contents
Classes
Functions
Data
API
Student/teacher setups plus their global-rank assignments.
Whether the assignments are disjoint.
Resolved student topology and policies.
Ordered global ranks assigned to the student.
Resolved teacher topology and policies.
Ordered global ranks assigned to the teacher.
Move batches and teacher logits between disjoint model meshes.
Parameters:
Resolved disjoint student and teacher setups.
Device used for transport tensors and collectives.
Broadcast a nested tensor tree along one route.
Tensor leaves may have arbitrary shapes and axis order.
Parameters:
Nested source value on route.src and None elsewhere.
Source, membership, and process group for the broadcast.
Returns: Any
Reconstructed value on route members and None elsewhere.
Broadcast one worker command from the first student rank.
Parameters:
Command supplied on student ranks. Teacher ranks pass
None while waiting for the broadcast.
Returns: int
Broadcast command value on every student and teacher rank.
Match full teacher logits to a TP-sharded student vocabulary.
Parameters:
Tensor of global shape
[batch, sequence, vocab] containing student logits. A
vocabulary-sharded DTensor has local shape
[batch, sequence, local_vocab] and Shard(-1) placement.
Replicated tensor of shape
[batch, sequence, vocab] containing teacher logits.
Returns: torch.Tensor
Replicated tensor of shape [batch, sequence, vocab] for a
Move every tensor leaf in a nested value to the bridge device.
Tensor leaves may have arbitrary shapes and axis order.
Parameters:
Nested dictionaries, lists, tuples, scalar values, and tensor leaves.
Returns: Any
Equivalent nested value whose tensors preserve shape and dtype on
Send one nested batch from each active student replica to a teacher.
Batch tensor leaves preserve their original shapes and axis order.
Parameters:
Routing wave index.
Nested student batch on student ranks and None on teacher
ranks.
Returns: Any | None
Assigned nested batch on teacher ranks and None on student ranks.
Send full teacher logits back to the assigned student replica.
Parameters:
Routing wave index.
Tensor of shape [batch, sequence, vocab] containing full
replicated teacher logits on the teacher output rank, otherwise
None.
Returns: torch.Tensor | None
Replicated tensor of shape [batch, sequence, vocab] on ranks in
Wait for both model roles without using the default process group.
Minimal recipe-config interface consumed by KD topology setup.
Return one configuration value or default.
Return the explicitly requested mesh size for a separate KD model.
Rebuild a nested value from tensor leaves described by _tree_spec.
Tensor leaves preserve the exact shapes and dtypes encoded in spec and
alias the corresponding entries in tensors.
Parameters:
Metadata returned by _tree_spec.
Tensor leaves with the exact recorded shapes and dtypes.
Returns: Any
Reconstructed nested value whose tensor leaves alias tensors.
Describe a nested value while extracting tensor leaves without copies.
Tensor leaves may have arbitrary rank and axis semantics; each leaf’s exact shape and dtype are recorded for allocation on receiving ranks.
Parameters:
Nested dictionaries, lists, tuples, scalar values, and tensor leaves with arbitrary shape and axis order.
Output list populated with aliases of tensor leaves in traversal order.
Returns: Any
Pickle-compatible nested metadata describing value.
Adapt NEAT-packed teachers and check both roles consume the same mask layout.
Parameters:
Locally owned teacher stages, empty on student-only ranks.
Locally owned student stages, empty on teacher-only ranks.
Shared group for separate-mesh KD; all its ranks must call. Same-mesh KD passes None and performs only a local check.
Raises:
ValueError: If teacher and student packed mask layouts disagree.
Build shared or explicitly disjoint student and teacher setups.
Parameters:
Recipe configuration containing distributed and optional
teacher_distributed sections.
Total global process count available to both models.
Returns: KDDistributedSetups
Resolved model setups and their ordered global-rank assignments.
Reconstruct full teacher logits across TP and CP before mesh transport.
Parameters:
Tensor of global shape [batch, sequence, vocab] containing
teacher logits. It may be a vocabulary-sharded DTensor with
Shard(-1) placement and/or have a per-rank load-balanced
local_sequence extent under CP.
Teacher mesh containing optional tp and cp axes.
Unpadded global sequence length to retain.
Returns: torch.Tensor
Detached contiguous tensor of shape