#!/usr/bin/env python3 # -*- coding: utf-8 -*- """ Created on Tue Jul 9 15:47:48 2024 @author: louis """ # dropped EDR/SRMR/DRR metrics: training-only, unused at inference. import torch from torch import nn class IgnoreLastSamplesMetricWrapper(nn.Module): def __init__(self, base_metric: nn.Module, num_ignored_first_samples: int = 0, num_ignored_last_samples: int = 256): super().__init__() self.base_metric = base_metric self.num_ignored_last_samples = num_ignored_last_samples self.num_ignored_first_samples = num_ignored_first_samples def forward(self, pred, target): if self.num_ignored_last_samples > 0: return self.base_metric( pred[..., self.num_ignored_first_samples : -self.num_ignored_last_samples], target[..., self.num_ignored_first_samples : -self.num_ignored_last_samples], ) else: return self.base_metric( pred[..., self.num_ignored_first_samples :], target[..., self.num_ignored_first_samples :], ) def __str__(self): return f"{type(self.base_metric).__name__}[{self.num_ignored_first_samples}:-{self.num_ignored_last_samples}]" class DifferenceMetric(nn.Module): def __init__(self, property_to_measure): super().__init__() self.property_to_measure = property_to_measure def forward(self, pred, target): if isinstance(target, tuple): target = target[0] pred = pred[..., : target.size(-1)] target = target[..., : pred.size(-1)] pred_property = self.property_to_measure(pred) target_property = self.property_to_measure(target) pred_property_finite = pred_property[torch.logical_and(pred_property.isfinite(), target_property.isfinite())] target_property_finite = target_property[ torch.logical_and(pred_property.isfinite(), target_property.isfinite()) ] return (pred_property_finite - target_property_finite).abs().mean() def __str__(self): return f"{type(self.property_to_measure).__name__} difference"