"""Render the saved splats without reading the source photograph."""
import argparse
from pathlib import Path
import numpy as np
import torch
from PIL import Image
from gsplat import rasterization

p = argparse.ArgumentParser()
p.add_argument('--parameters', default='site/gaussians.npz')
p.add_argument('--output', default='results/rerender.png')
args = p.parse_args()
d = np.load(args.parameters, allow_pickle=False)
w, h = int(d['width']), int(d['height'])
xy, rgb, log_scales, angle, opacity = [torch.as_tensor(d[k],device='cuda')
    for k in ['xy','rgb','log_scales','angle','opacity']]
n = len(xy)
f = w / 2
z = torch.full((n,1),3.,device='cuda')
means = torch.cat([(xy*torch.tensor([w,h],device='cuda')-torch.tensor([w/2,h/2],device='cuda'))*(3/f),z],-1)
scales = torch.cat([log_scales.exp()*(3/f),torch.full((n,1),.0001,device='cuda')],-1)
zero = torch.zeros(n,device='cuda')
quat = torch.stack([torch.cos(angle/2),zero,zero,torch.sin(angle/2)],-1)
K = torch.tensor([[f,0,w/2],[0,f,h/2],[0,0,1]],device='cuda')[None]
with torch.no_grad():
    im = rasterization(means,quat,scales,opacity.sigmoid(),rgb.sigmoid(),
                       torch.eye(4,device='cuda')[None],K,w,h,packed=False,eps2d=.1)[0][0]
Path(args.output).parent.mkdir(parents=True,exist_ok=True)
Image.fromarray((im.clamp(0,1).cpu().numpy()*255).round().astype('uint8')).save(args.output)
print(f'Rendered {n:,} Gaussians at {w}x{h} from saved parameters only.')
