nemo_automodel.components.loss.linear_ce_base
nemo_automodel.components.loss.linear_ce_base
Module Contents
Classes
API
Bases: Module
Losses consuming final hidden states and an LM-head weight instead of logits.
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-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.
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.