Spaces:
Sleeping
Sleeping
File size: 2,140 Bytes
d98780c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 | #!/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"
|