Download main_method/code/losses.py from Dhruv1000/TSRDA: direct link, hf CLI and curl.
- Browser
- Download file 2.97 kB
-
https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/losses.py
- Command line
-
hf download hf://Dhruv1000/TSRDA/main_method/code/losses.py
-
curl -L -o losses.py https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/losses.py
2.97 kB
| """ | |
| Lovász-Softmax loss — NaN-safe version for small-batch training. | |
| Fixes over v1: | |
| - eps guard on union denominator (prevents /0 on rare classes) | |
| - support for multiple ignore indices (background + void) | |
| """ | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| def lovasz_grad(gt_sorted): | |
| """Gradient of Lovász extension w.r.t. sorted errors.""" | |
| p = len(gt_sorted) | |
| gts = gt_sorted.sum() | |
| intersection = gts - gt_sorted.float().cumsum(0) | |
| union = gts + (1.0 - gt_sorted).float().cumsum(0) | |
| jaccard = 1.0 - intersection / (union + 1e-6) # eps guard | |
| if p > 1: | |
| jaccard[1:p] = jaccard[1:p] - jaccard[0:-1] | |
| return jaccard | |
| def lovasz_softmax_flat(probas, labels, classes="present"): | |
| """ | |
| Multi-class Lovász-Softmax on flattened tensors. | |
| probas: (N, C) softmax probabilities | |
| labels: (N,) ground truth indices | |
| """ | |
| if probas.numel() == 0: | |
| return probas * 0.0 | |
| C = probas.shape[1] | |
| if classes == "all": | |
| class_list = list(range(C)) | |
| elif classes == "present": | |
| class_list = torch.unique(labels).tolist() | |
| else: | |
| class_list = list(classes) | |
| losses = [] | |
| for c in class_list: | |
| fg = (labels == c).float() | |
| if fg.sum() == 0: | |
| continue | |
| class_pred = probas[:, c] | |
| errors = (fg - class_pred).abs() | |
| errors_sorted, perm = torch.sort(errors, dim=0, descending=True) | |
| fg_sorted = fg[perm.data] | |
| grad = lovasz_grad(fg_sorted) | |
| loss_c = torch.dot(errors_sorted, grad) | |
| if torch.isfinite(loss_c): | |
| losses.append(loss_c) | |
| if not losses: | |
| return probas.sum() * 0.0 | |
| return sum(losses) / len(losses) | |
| class LovaszSoftmaxLoss(nn.Module): | |
| """ | |
| Lovász-Softmax for 2D segmentation. | |
| logits: (B, C, H, W) | |
| targets: (B, H, W) long | |
| ignore_indices: list of class indices to exclude (e.g. [0, 19]) | |
| """ | |
| def __init__(self, ignore_indices=None, classes="present"): | |
| super().__init__() | |
| self.ignore_indices = ignore_indices or [] | |
| self.classes = classes | |
| def forward(self, logits, targets): | |
| probas = F.softmax(logits, dim=1) | |
| B, C, H, W = probas.shape | |
| probas_flat = probas.permute(0, 2, 3, 1).contiguous().view(-1, C) | |
| targets_flat = targets.contiguous().view(-1) | |
| # Mask out all ignore indices | |
| if self.ignore_indices: | |
| valid = torch.ones(targets_flat.shape[0], dtype=torch.bool, | |
| device=targets_flat.device) | |
| for idx in self.ignore_indices: | |
| valid &= (targets_flat != idx) | |
| probas_flat = probas_flat[valid] | |
| targets_flat = targets_flat[valid] | |
| if probas_flat.numel() == 0: | |
| return logits.sum() * 0.0 | |
| loss = lovasz_softmax_flat(probas_flat, targets_flat, classes=self.classes) | |
| return torch.nan_to_num(loss, nan=0.0) | |