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
|