Skip to content

4. The loss

losses/loss_functions/my_loss.py → a term inside cfg.criterion_params

criterion_name is required. Even a single-term loss goes through gatle_ignite.losses.composite, a weighted sum built from config:

cfg.criterion_name = "gatle_ignite.losses.composite"
cfg.criterion_params = {"dict_of_loss_params": {
    "ce":  {"cls_name": "losses.loss_functions.cross_entropy",
            "loss_params": {"src_name": "logits", "tgt_name": ("targets", "labels")},
            "weight": 1.0},
    "aux": {"cls_name": "losses.loss_functions.aux", "loss_params": {...},
            "weight": 0.1, "start_iteration": 1000},
}}

src_name resolves against the model's output dict; tgt_name against prep_batch's whole return, which is why the paths start ("targets", ...). A string is a top-level key of that return, a tuple is a path into it. start_iteration / end_iteration switch a term on and off as training progresses.

There are no built-in loss terms, because a loss term describes your task. For cross-entropy, copy the synthetic example's:

examples/synthetic/losses/loss_functions/cross_entropy.py
import torch.nn as nn

from gatle_ignite import get_value


class Loss(nn.Module):
    def __init__(self, src_name="logits", tgt_name=("targets", "labels"), **kwargs):
        super().__init__()
        self.src_name = src_name
        self.tgt_name = tgt_name
        self.loss = nn.CrossEntropyLoss(**kwargs)

    def forward(self, y_pred, target, iteration=None):
        return self.loss(get_value(y_pred, self.src_name), get_value(target, self.tgt_name))

Writing one

class Loss(nn.Module):
    def forward(self, y_pred, target, iteration=None) -> Tensor
examples/contrastive/losses/loss_functions/ntxent.py
import torch
import torch.nn as nn
import torch.nn.functional as F

from gatle_ignite import get_value


class Loss(nn.Module):
    def __init__(self, src_a="z1", src_b="z2", temperature=0.2):
        super().__init__()
        # Two fields, not a tuple: get_value's tuple form is a PATH into nested dicts.
        self.src_a, self.src_b = src_a, src_b
        self.temperature = temperature

    def forward(self, y_pred, target=None, iteration=None):
        z1 = get_value(y_pred, self.src_a)  # (B, D), already L2-normalised by the model
        z2 = get_value(y_pred, self.src_b)  # (B, D)
        b = z1.shape[0]
        if b < 2:
            raise ValueError(
                f"NT-Xent needs at least 2 samples per batch, got {b}: with one sample "
                "there are no negatives and the loss is identically zero."
            )

        z = torch.cat([z1, z2], dim=0)  # (2B, D)
        sim = (z @ z.T) / self.temperature  # (2B, 2B)

        # Mask the diagonal with -inf: a sample is not its own negative. Out-of-place, since
        # in-place on an autograd-tracked output risks a version-counter error.
        eye = torch.eye(2 * b, dtype=torch.bool, device=z.device)
        sim = sim.masked_fill(eye, float("-inf"))

        # THE TARGETS: i's positive is at i+B and (i+B)'s at i, indices into this batch.
        targets = torch.cat(
            [torch.arange(b, 2 * b, device=z.device), torch.arange(0, b, device=z.device)]
        )
        return F.cross_entropy(sim, targets)

A sub-loss returns a bare tensor; a whole criterion does not

Replacing criterion_name itself means returning (total, {"loss_<key>": tensor}) and carrying a crit_keys attribute. Inside the composite, just return the scalar.

The weight scales the term as it is logged, not only as it is summed

So loss_aux charts its real contribution, but two terms with different weights are on different scales and are not comparable by eye. A gated-off term logs as exactly 0.0.

CTC: log-softmax in float32, and CTCLoss with autocast off

Take log_softmax of logits.float() and call nn.CTCLoss inside torch.autocast(..., enabled=False): cuDNN's CTC kernel can produce NaN under bf16.

Templates: loss