File size: 3,677 Bytes
4c3d957
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
"""Appearance metrics on object-masked renders (LOCKED spec).

Rankings (per PI decision): LPIPS primary, SSIM secondary, CLIP tertiary.
PSNR reported but NON-RANKING. NO FID/KID, NO PSNR ranking.

  * LPIPS(net='alex')  : lpips pkg, inputs in [-1,1].
  * SSIM               : skimage structural_similarity, win_size=11,
                         gaussian_weights=True, data_range=1, channel_axis=-1.
  * CLIP-similarity    : cosine of open_clip ViT-B-32 image embeddings.
  * PSNR               : torchmetrics PeakSignalNoiseRatio data_range=1.

Masking: every metric is computed on the object composited over a fixed
WHITE background using the render alpha (alpha>0 => object). Pred and GT use
the identical background, so silhouette disagreement is penalised (it shows up
as object-vs-background) while background pixels never bias the score.
Metrics are computed per view, then averaged over views, then over objects
(the averaging over objects happens in the CLI).
"""

from __future__ import annotations

import numpy as np
import torch

from skimage.metrics import structural_similarity as _ssim
from torchmetrics.image import PeakSignalNoiseRatio

_LPIPS = None
_CLIP = None
_CLIP_PRE = None
_PSNR = None
_DEVICE = "cuda"


def _lazy():
    global _LPIPS, _CLIP, _CLIP_PRE, _PSNR
    if _LPIPS is None:
        import lpips
        _LPIPS = lpips.LPIPS(net="alex").to(_DEVICE).eval()
    if _CLIP is None:
        import open_clip
        import os
        cache = os.environ.get("HF_HOME", None)
        model, _, pre = open_clip.create_model_and_transforms(
            "ViT-B-32", pretrained="openai", cache_dir=cache)
        _CLIP = model.to(_DEVICE).eval()
        _CLIP_PRE = pre
    if _PSNR is None:
        _PSNR = PeakSignalNoiseRatio(data_range=1.0).to(_DEVICE)


def composite_white(rgba: np.ndarray) -> np.ndarray:
    """Composite an (H,W,4) float[0,1] RGBA over white -> (H,W,3) float[0,1]."""
    rgb = rgba[..., :3]
    a = rgba[..., 3:4]
    return rgb * a + (1.0 - a)


def _prep_tensor(rgb):
    """(H,W,3) float[0,1] -> (1,3,H,W) tensor on device."""
    t = torch.as_tensor(rgb, dtype=torch.float32, device=_DEVICE)
    return t.permute(2, 0, 1)[None].contiguous()


@torch.no_grad()
def appearance_metrics(pred_rgba: np.ndarray, ref_rgba: np.ndarray):
    """Compare two RGBA renders. Returns dict lpips/ssim/clip/psnr for one view.

    Both inputs are (H,W,4) float in [0,1]. Composited over white internally.
    """
    _lazy()
    pred = composite_white(np.asarray(pred_rgba, np.float32))
    ref = composite_white(np.asarray(ref_rgba, np.float32))

    pt = _prep_tensor(pred)
    rt = _prep_tensor(ref)

    # LPIPS wants [-1,1]
    lp = float(_LPIPS(pt * 2 - 1, rt * 2 - 1).item())

    # SSIM (skimage, per channel gaussian)
    ss = float(_ssim(pred, ref, win_size=11, gaussian_weights=True,
                     data_range=1.0, channel_axis=-1))

    # PSNR (non-ranking)
    ps = float(_PSNR(pt, rt).item())

    # CLIP cosine
    from PIL import Image
    pi = _CLIP_PRE(Image.fromarray((pred * 255).astype(np.uint8))).unsqueeze(0).to(_DEVICE)
    ri = _CLIP_PRE(Image.fromarray((ref * 255).astype(np.uint8))).unsqueeze(0).to(_DEVICE)
    fe = _CLIP.encode_image(pi)
    ge = _CLIP.encode_image(ri)
    fe = fe / fe.norm(dim=-1, keepdim=True)
    ge = ge / ge.norm(dim=-1, keepdim=True)
    clip = float((fe * ge).sum(-1).item())

    return {"lpips": lp, "ssim": ss, "clip": clip, "psnr": ps}


def average_views(view_dicts):
    """Mean each metric across a list of per-view dicts."""
    if not view_dicts:
        return {}
    keys = view_dicts[0].keys()
    return {k: float(np.mean([d[k] for d in view_dicts])) for k in keys}