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"