nemo_automodel.components.loss.linear_ce_base

View as Markdown

Module Contents

Classes

NameDescription
LinearCrossEntropyLosses consuming final hidden states and an LM-head weight instead of logits.

API

class nemo_automodel.components.loss.linear_ce_base.LinearCrossEntropy()

Bases: Module

Losses consuming final hidden states and an LM-head weight instead of logits.

nemo_automodel.components.loss.linear_ce_base.LinearCrossEntropy.materialize_lm_weight(
lm_weight: torch.Tensor,
grad_reduce_group: torch.distributed.ProcessGroup | None = None
) -> torch.Tensor
staticmethod

Materialize an LM-head DTensor with gradient-correct reduction semantics.

Linear CE consumes the LM-head weight outside the owning FSDP module’s forward. Each data/context-parallel rank therefore computes a rank-local full-weight gradient. A plain DTensor.full_tensor() marks that gradient as replicated, so backward only slices the local result into the owned shard instead of combining peer contributions.

Parameters:

lm_weight
torch.Tensor

LM-head weight with global shape [vocab, hidden]. A regular tensor is returned unchanged. A DTensor may have any FSDP sharding placement over its device mesh and is gathered to a rank-local regular tensor with the global shape, device, and dtype.

grad_reduce_group
dist.ProcessGroup | NoneDefaults to None

Process group whose ranks contribute independent token losses. Its size must match the LM-head DTensor mesh.

Returns: torch.Tensor

Regular tensor with shape [vocab, hidden]. For a DTensor input,

Raises:

  • ValueError: If a trainable sharded weight has no matching reduction group. This fails closed instead of producing rank-local shards.