nemo_automodel.recipes.kd_utils

View as Markdown

Distributed topology and tensor transport helpers for KD recipes.

Module Contents

Classes

NameDescription
KDDistributedSetupsStudent/teacher setups plus their global-rank assignments.
KDMeshBridgeMove batches and teacher logits between disjoint model meshes.
_ConfigLikeMinimal recipe-config interface consumed by KD topology setup.
_Replica-
_Route-

Functions

NameDescription
_mesh_sizeReturn the explicitly requested mesh size for a separate KD model.
_model_replicas-
_section_to_dict-
_tree_from_specRebuild a nested value from tensor leaves described by _tree_spec.
_tree_specDescribe a nested value while extracting tensor leaves without copies.
configure_kd_teacher_packingAdapt NEAT-packed teachers and check both roles consume the same mask layout.
create_kd_distributed_setupsBuild shared or explicitly disjoint student and teacher setups.
materialize_teacher_logitsReconstruct full teacher logits across TP and CP before mesh transport.

Data

RUN_TEACHER

STOP_TEACHER

API

class nemo_automodel.recipes.kd_utils.KDDistributedSetups(
student_ranks: tuple[int, ...],
teacher_ranks: tuple[int, ...],
separate: bool
)
Dataclass

Student/teacher setups plus their global-rank assignments.

separate
bool

Whether the assignments are disjoint.

student
DistributedSetup

Resolved student topology and policies.

student_ranks
tuple[int, ...]

Ordered global ranks assigned to the student.

teacher
DistributedSetup

Resolved teacher topology and policies.

teacher_ranks
tuple[int, ...]

Ordered global ranks assigned to the teacher.

class nemo_automodel.recipes.kd_utils.KDMeshBridge(
device: torch.device
)

Move batches and teacher logits between disjoint model meshes.

Parameters:

setups
KDDistributedSetups

Resolved disjoint student and teacher setups.

device
torch.device

Device used for transport tensors and collectives.

control_group
input_routes
list[list[_Route]] = []
is_student
bool
is_teacher
bool
num_waves
output_routes
list[list[_Route]] = []
rank
= dist.get_rank()
student_group
= dist.new_group(ranks=(list(self.student_ranks)))
student_ranks
= setups.student_ranks
student_replicas
= _model_replicas(setups.student)
teacher_group
= dist.new_group(ranks=(list(self.teacher_ranks)))
teacher_ranks
= setups.teacher_ranks
teacher_replicas
= _model_replicas(setups.teacher)
nemo_automodel.recipes.kd_utils.KDMeshBridge._broadcast_tree(
value: typing.Any,
) -> typing.Any

Broadcast a nested tensor tree along one route.

Tensor leaves may have arbitrary shapes and axis order.

Parameters:

value
Any

Nested source value on route.src and None elsewhere.

route
_Route

Source, membership, and process group for the broadcast.

Returns: Any

Reconstructed value on route members and None elsewhere.

nemo_automodel.recipes.kd_utils.KDMeshBridge.broadcast_command(
command: int | None = None
) -> int

Broadcast one worker command from the first student rank.

Parameters:

command
int | NoneDefaults to None

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.

nemo_automodel.recipes.kd_utils.KDMeshBridge.match_student_vocab_shard(
student_logits: torch.Tensor,
teacher_logits: torch.Tensor
) -> torch.Tensor
staticmethod

Match full teacher logits to a TP-sharded student vocabulary.

Parameters:

student_logits
torch.Tensor

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.

teacher_logits
torch.Tensor

Replicated tensor of shape [batch, sequence, vocab] containing teacher logits.

Returns: torch.Tensor

Replicated tensor of shape [batch, sequence, vocab] for a

nemo_automodel.recipes.kd_utils.KDMeshBridge.move_to_device(
value: typing.Any
) -> typing.Any

Move every tensor leaf in a nested value to the bridge device.

Tensor leaves may have arbitrary shapes and axis order.

Parameters:

value
Any

Nested dictionaries, lists, tuples, scalar values, and tensor leaves.

Returns: Any

Equivalent nested value whose tensors preserve shape and dtype on

nemo_automodel.recipes.kd_utils.KDMeshBridge.send_batch(
wave: int,
batch: typing.Any | None
) -> typing.Any | None

Send one nested batch from each active student replica to a teacher.

Batch tensor leaves preserve their original shapes and axis order.

Parameters:

wave
int

Routing wave index.

batch
Any | None

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.

nemo_automodel.recipes.kd_utils.KDMeshBridge.send_logits(
wave: int,
logits: torch.Tensor | None
) -> torch.Tensor | None

Send full teacher logits back to the assigned student replica.

Parameters:

wave
int

Routing wave index.

logits
torch.Tensor | None

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

nemo_automodel.recipes.kd_utils.KDMeshBridge.synchronize() -> None

Wait for both model roles without using the default process group.

class nemo_automodel.recipes.kd_utils._ConfigLike()
Protocol

Minimal recipe-config interface consumed by KD topology setup.

nemo_automodel.recipes.kd_utils._ConfigLike.get(
key: str,
default: typing.Any = None
) -> typing.Any

Return one configuration value or default.

class nemo_automodel.recipes.kd_utils._Replica(
ranks: tuple[int, ...],
input_rank: int,
output_rank: int
)
Dataclass
input_rank
int
output_rank
int
ranks
tuple[int, ...]
class nemo_automodel.recipes.kd_utils._Route(
src: int,
ranks: tuple[int, ...],
group: torch.distributed.ProcessGroup
)
Dataclass
group
ProcessGroup
ranks
tuple[int, ...]
src
int
nemo_automodel.recipes.kd_utils._mesh_size(
distributed_cfg: typing.Any,
label: str
) -> int

Return the explicitly requested mesh size for a separate KD model.

nemo_automodel.recipes.kd_utils._section_to_dict(
section: typing.Any
) -> dict
nemo_automodel.recipes.kd_utils._tree_from_spec(
spec: typing.Any,
tensors: list[torch.Tensor]
) -> typing.Any

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:

spec
Any

Metadata returned by _tree_spec.

tensors
list[torch.Tensor]

Tensor leaves with the exact recorded shapes and dtypes.

Returns: Any

Reconstructed nested value whose tensor leaves alias tensors.

nemo_automodel.recipes.kd_utils._tree_spec(
value: typing.Any,
tensors: list[torch.Tensor]
) -> typing.Any

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:

value
Any

Nested dictionaries, lists, tuples, scalar values, and tensor leaves with arbitrary shape and axis order.

tensors
list[torch.Tensor]

Output list populated with aliases of tensor leaves in traversal order.

Returns: Any

Pickle-compatible nested metadata describing value.

nemo_automodel.recipes.kd_utils.configure_kd_teacher_packing(
teacher_parts: collections.abc.Sequence[torch.nn.Module],
student_parts: collections.abc.Sequence[torch.nn.Module],
control_group: torch.distributed.ProcessGroup | None = None
) -> None

Adapt NEAT-packed teachers and check both roles consume the same mask layout.

Parameters:

teacher_parts
Sequence[torch.nn.Module]

Locally owned teacher stages, empty on student-only ranks.

student_parts
Sequence[torch.nn.Module]

Locally owned student stages, empty on teacher-only ranks.

control_group
dist.ProcessGroup | NoneDefaults to None

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.
nemo_automodel.recipes.kd_utils.create_kd_distributed_setups(
world_size: int

Build shared or explicitly disjoint student and teacher setups.

Parameters:

cfg
_ConfigLike

Recipe configuration containing distributed and optional teacher_distributed sections.

world_size
int

Total global process count available to both models.

Returns: KDDistributedSetups

Resolved model setups and their ordered global-rank assignments.

nemo_automodel.recipes.kd_utils.materialize_teacher_logits(
logits: torch.Tensor,
device_mesh: 'DeviceMesh',
sequence_length: int
) -> torch.Tensor

Reconstruct full teacher logits across TP and CP before mesh transport.

Parameters:

logits
torch.Tensor

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.

device_mesh
'DeviceMesh'

Teacher mesh containing optional tp and cp axes.

sequence_length
int

Unpadded global sequence length to retain.

Returns: torch.Tensor

Detached contiguous tensor of shape

nemo_automodel.recipes.kd_utils.RUN_TEACHER = 1
nemo_automodel.recipes.kd_utils.STOP_TEACHER = 0