Spaces:
Running on Zero
Running on Zero
| #!/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" | |