"""SSIM as used in the article: 11-tap Gaussian (sigma 1.5), C1=(0.01*255)^2, C2=(0.03*255)^2, mean over R, G, B.
usage: ssim.py <reference> <candidate>...   prints bytes and SSIM for each candidate."""
import sys, os, subprocess, tempfile
import numpy as np
from PIL import Image
from scipy.ndimage import correlate1d

def load(path):
    try:
        im = Image.open(path); im.load()
    except Exception:
        tmp = tempfile.mktemp(suffix=".png")
        subprocess.run(["magick", path, tmp], check=True)
        im = Image.open(tmp); im.load(); os.unlink(tmp)
    if im.mode in ("RGBA", "LA", "P"):
        im = im.convert("RGBA"); bg = Image.new("RGBA", im.size, (255, 255, 255, 255)); im = Image.alpha_composite(bg, im)
    return np.asarray(im.convert("RGB"), dtype=np.float64)

x = np.arange(-5, 6); K = np.exp(-(x**2) / (2 * 1.5**2)); K /= K.sum()
def blur(a): return correlate1d(correlate1d(a, K, axis=0, mode="reflect"), K, axis=1, mode="reflect")
C1, C2 = (0.01 * 255) ** 2, (0.03 * 255) ** 2
def ssim(a, b):
    vals = []
    for c in range(3):
        p, q = a[..., c], b[..., c]
        mp, mq = blur(p), blur(q)
        vp, vq, cov = blur(p * p) - mp * mp, blur(q * q) - mq * mq, blur(p * q) - mp * mq
        vals.append((((2 * mp * mq + C1) * (2 * cov + C2)) / ((mp * mp + mq * mq + C1) * (vp + vq + C2))).mean())
    return float(np.mean(vals))

