| 38 | return self.fp_senet_gt(img).unsqueeze(1) |
| 39 | |
| 40 | class reconstructor_loss(nn.Module): |
| 41 | def __init__(self): |
| 42 | super(reconstructor_loss, self).__init__() |
| 43 | |
| 44 | def forward(self, pred, gt): |
| 45 | left_loss=F.mse_loss(pred, gt, reduce=False) |
| 46 | return torch.mean(left_loss) |
| 47 | |
| 48 | |
| 49 | if __name__=='__main__': |