(self, device='cuda')
| 66 | """LPIPS metric wrapper.""" |
| 67 | |
| 68 | def __init__(self, device='cuda'): |
| 69 | self.device = device |
| 70 | self.model = LPIPS_VTP().to(device).eval() |
| 71 | |
| 72 | def __call__(self, img1: torch.Tensor, img2: torch.Tensor) -> torch.Tensor: |
| 73 | """Calculate LPIPS between two images. |