UDREAM / model /utils /metrics.py
Vansh Chugh
initial deploy
d98780c
Raw
History Blame Contribute Delete
2.14 kB
#!/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"