import torch import numpy as np from collections import OrderedDict class SegmentationMetrics: """ Computes: OA, mIoU, mFscore, mPrecision, mRecall, Kappa """ def __init__(self, num_classes, ignore_index=19): self.num_classes = num_classes self.ignore_index = ignore_index self.reset() def reset(self): self.total_intersect = torch.zeros(self.num_classes) self.total_union = torch.zeros(self.num_classes) self.total_pred_label = torch.zeros(self.num_classes) self.total_label = torch.zeros(self.num_classes) # For Kappa: full confusion matrix self.confusion_matrix = torch.zeros(self.num_classes, self.num_classes) def update(self, pred_logits, labels): pred_label = pred_logits.argmax(dim=1) # (B, H, W) for b in range(pred_label.shape[0]): intersect, union, pred_area, label_area = self._intersect_and_union( pred_label[b], labels[b] ) self.total_intersect += intersect self.total_union += union self.total_pred_label += pred_area self.total_label += label_area self._update_confusion(pred_label[b], labels[b]) def _intersect_and_union(self, pred, label): if pred.shape != label.shape: label = label.float().unsqueeze(0).unsqueeze(0) label = torch.nn.functional.interpolate( label, size=pred.shape, mode='nearest' ).squeeze().long() mask = label != self.ignore_index pred = pred[mask] label = label[mask] intersect = pred[pred == label] area_intersect = torch.histc(intersect.float(), bins=self.num_classes, min=0, max=self.num_classes-1).cpu() area_pred = torch.histc(pred.float(), bins=self.num_classes, min=0, max=self.num_classes-1).cpu() area_label = torch.histc(label.float(), bins=self.num_classes, min=0, max=self.num_classes-1).cpu() area_union = area_pred + area_label - area_intersect return area_intersect, area_union, area_pred, area_label def _update_confusion(self, pred, label): if pred.shape != label.shape: label = label.float().unsqueeze(0).unsqueeze(0) label = torch.nn.functional.interpolate( label, size=pred.shape, mode='nearest' ).squeeze().long() mask = label != self.ignore_index pred = pred[mask].cpu() label = label[mask].cpu() # Accumulate confusion matrix indices = self.num_classes * label + pred cm = torch.bincount(indices, minlength=self.num_classes**2) self.confusion_matrix += cm.reshape(self.num_classes, self.num_classes).float() def _compute_kappa(self): cm = self.confusion_matrix.numpy() n = cm.sum() if n == 0: return 0. po = np.diag(cm).sum() / n pe = (cm.sum(axis=0) * cm.sum(axis=1)).sum() / (n * n) if abs(1 - pe) < 1e-10: return 1. return (po - pe) / (1 - pe) def compute(self): eps = 1e-8 aAcc = self.total_intersect.sum() / (self.total_label.sum() + eps) iou = self.total_intersect / (self.total_union + eps) precision = self.total_intersect / (self.total_pred_label + eps) recall = self.total_intersect / (self.total_label + eps) fscore = 2 * precision * recall / (precision + recall + eps) valid = self.total_label > 0 valid[self.ignore_index] = False mIoU = iou[valid].mean().item() if valid.any() else 0. mFscore = fscore[valid].mean().item() if valid.any() else 0. mPrecision = precision[valid].mean().item() if valid.any() else 0. mRecall = recall[valid].mean().item() if valid.any() else 0. kappa = self._compute_kappa() return OrderedDict({ 'OA': round(aAcc.item() * 100, 2), 'mIoU': round(mIoU * 100, 2), 'mFscore': round(mFscore * 100, 2), 'mPrecision': round(mPrecision * 100, 2), 'mRecall': round(mRecall * 100, 2), 'Kappa': round(kappa * 100, 2), }) def per_class_iou(self): eps = 1e-8 return (self.total_intersect / (self.total_union + eps)).numpy()