File size: 1,377 Bytes
338c3e4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from typing import List

import torch
from torch import Tensor, nn
import torch.nn.functional as F


class MseLoss(nn.Module):
    def __init__(self, normalize: bool, is_masked: bool = False):
        super().__init__()
        self.normalize = normalize
        self.is_masked = is_masked

    def get_score_names(self) -> List[str]:
        names = [
            'mse', 'rmse', 'mae'
        ]
        if self.normalize:
            names += ['nmse']
        return names

    def forward(self, preds: Tensor, labels: Tensor) -> dict[str, Tensor]:
        '''
        Args:
        - mask: 1 for valid pixels, 0 for invalid pixels.
        '''
        mse = F.mse_loss(input=preds, target=labels)
        mae = F.l1_loss(input=preds, target=labels)
        result = dict(
            mse=mse,
            rmse=torch.sqrt(mse),
            mae=mae,
        )
        if self.normalize:
            result['nmse'] = mse / torch.square(labels).mean()
            # result['mre'] = mae / torch.abs(labels).mean()
        return result


def loss_name_to_fn(name: str, masked: bool = False) -> MseLoss:
    name = name.lower()
    if masked:
        raise NotImplementedError
    else:
        if name == 'mse':
            return MseLoss(normalize=False)
        elif name == 'nmse':
            return MseLoss(normalize=True)
        else:
            raise NotImplementedError