"""Segmentation losses.""" import torch import torch.nn.functional as F def cross_entropy_loss(logits, targets, ignore_index=255): """Standard per-pixel cross entropy. logits: [B, C, H, W], targets: [B, H, W].""" if logits.shape[2:] != targets.shape[1:]: logits = F.interpolate(logits, size=targets.shape[1:], mode="bilinear", align_corners=False) return F.cross_entropy(logits, targets, ignore_index=ignore_index)