"""Single-view image fitting with gsplat's CUDA rasterizer (not 3D reconstruction).

All centers lie on one plane. Optimize their image-plane locations, anisotropic
sizes, in-plane orientations, colors, and opacities against the uploaded RGB image.
Adapted conceptually from gsplat/examples/image_fitting.py (Apache-2.0).
"""
import argparse
import hashlib
import json
import math
import time
from pathlib import Path

import numpy as np
from PIL import Image, ImageOps
import torch
from gsplat import rasterization
import gsplat


def main():
    p = argparse.ArgumentParser()
    p.add_argument('--image', default='site/original.jpg')
    p.add_argument('--out', default='site')
    p.add_argument('--steps', type=int, default=5000)
    p.add_argument('--num-points', type=int, default=24000,
                   help='Model capacity: number of Gaussians, not matrix rank.')
    p.add_argument('--results', default='results')
    p.add_argument('--seed', type=int, default=29)
    p.add_argument('--checkpoint-every', type=int, default=500)
    args = p.parse_args()
    if args.num_points < 4 or args.steps < 1:
        p.error('--num-points must be at least 4 and --steps at least 1')
    torch.manual_seed(args.seed)
    np.random.seed(args.seed)
    torch.set_num_threads(6)
    out = Path(args.out)
    (out / 'frames').mkdir(parents=True, exist_ok=True)
    results = Path(args.results)
    results.mkdir(parents=True, exist_ok=True)
    image = ImageOps.exif_transpose(Image.open(args.image)).convert('RGB')
    # Re-encode without embedded metadata; these pixels define the fitting target.
    image.save(out / 'reference.png')
    w, h = image.size
    target = torch.tensor(np.asarray(image).copy(), device='cuda', dtype=torch.float32) / 255
    f = w / 2
    K = torch.tensor([[f, 0, w/2], [0, f, h/2], [0, 0, 1]], device='cuda')[None]
    view = torch.eye(4, device='cuda')[None]

    # A stratified covering grid plus detail-weighted points. This is initialization,
    # not a learned generative prior; both the initialization and loss see the photo.
    n = args.num_points
    grid_budget = int(.8 * n)
    nx = max(1, min(grid_budget, round(math.sqrt(grid_budget * w / h))))
    ny = max(1, grid_budget // nx)
    yy, xx = torch.meshgrid((torch.arange(ny, device='cuda')+.5)/ny,
                            (torch.arange(nx, device='cuda')+.5)/nx, indexing='ij')
    grid = torch.stack([xx.flatten(), yy.flatten()], -1)
    gray = target.mean(-1)
    edges = torch.zeros_like(gray)
    edges[:, 1:] += (gray[:, 1:] - gray[:, :-1]).abs()
    edges[1:, :] += (gray[1:, :] - gray[:-1, :]).abs()
    n_extra = n - len(grid)
    indices = torch.multinomial((edges + .003).flatten(), n_extra,
                                replacement=n_extra > w*h)
    extra = torch.stack([(indices % w + .5)/w, (indices // w + .5)/h], -1)
    initial = torch.cat([grid, extra], 0)
    assert len(initial) == n
    xy = torch.nn.Parameter(initial.clone())
    rgb0 = target[(initial[:, 1]*h).long().clamp(0,h-1),
                  (initial[:, 0]*w).long().clamp(0,w-1)]
    rgb = torch.nn.Parameter(torch.logit(rgb0.clamp(.005,.995)))
    spacing = max(w / nx, h / ny)
    max_scale = max(32., spacing * 4)
    log_scales = torch.nn.Parameter(torch.full((n, 2), math.log(spacing*.55), device='cuda'))
    angle = torch.nn.Parameter(torch.zeros(n, device='cuda'))
    opacity = torch.nn.Parameter(torch.full((n,), math.log(.7/.3), device='cuda'))
    optimizer = torch.optim.Adam([
        {'params': [xy], 'lr': .00035 * min(10., math.sqrt(24000/n))},
        {'params': [rgb], 'lr': .035},
        {'params': [log_scales], 'lr': .025},
        {'params': [angle], 'lr': .015},
        {'params': [opacity], 'lr': .025},
    ])
    initial_lrs = [group['lr'] for group in optimizer.param_groups]
    z = torch.full((n, 1), 3., device='cuda')
    zero = torch.zeros(n, device='cuda')
    tiny_z = torch.full((n, 1), .0001, device='cuda')

    def render():
        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), tiny_z], -1)
        quat = torch.stack([torch.cos(angle/2), zero, zero, torch.sin(angle/2)], -1)
        return rasterization(means=means, quats=quat, scales=scales,
                             opacities=opacity.sigmoid(), colors=rgb.sigmoid(),
                             viewmats=view, Ks=K, width=w, height=h,
                             packed=False, eps2d=.1)[0][0]

    def png(tensor, path):
        Image.fromarray((tensor.detach().clamp(0,1).cpu().numpy()*255).round().astype('uint8')).save(path)

    # Warm up native kernels before measuring optimization runtime.
    print(json.dumps({'status':'warming_up', 'gsplat':gsplat.__version__, 'points':n,
                      'width':w, 'height':h}), flush=True)
    first = render()
    first.mean().backward()
    optimizer.zero_grad(set_to_none=True)
    torch.cuda.synchronize()
    torch.cuda.reset_peak_memory_stats()
    started = time.perf_counter()
    frames, history = [], []
    best_mse = float('inf')
    best_step = 0
    checkpoints = {0, 25, 100, 250, 500, 1000, 1500, 2000, 3000, 4000, args.steps}
    checkpoints.update(range(args.checkpoint_every, args.steps+1, args.checkpoint_every))
    for step in range(args.steps + 1):
        rendered = render()
        loss = (rendered-target).square().mean()
        if step % 100 == 0 or step == args.steps:
            mse = loss.item()
            row = {'step':step,'mse':mse,'psnr_db':-10*math.log10(mse),
                   'seconds':time.perf_counter()-started}
            history.append(row)
            print(json.dumps(row), flush=True)
        if step in checkpoints:
            with torch.no_grad():
                mse = loss.item()
                name = f'frames/step-{step:05d}.webp'
                Image.fromarray((rendered.clamp(0,1).cpu().numpy()*255).round().astype('uint8')).save(out/name, quality=92)
                frames.append({'step':step,'image':name,'psnr_db':-10*math.log10(mse)})
                if mse < best_mse:
                    best_mse = mse
                    best_step = step
                    best = {name: param.detach().clone() for name,param in
                            [('xy',xy),('rgb',rgb),('log_scales',log_scales),('angle',angle),('opacity',opacity)]}
        if step == args.steps:
            break
        optimizer.zero_grad(set_to_none=True)
        loss.backward()
        optimizer.step()
        with torch.no_grad():
            xy.clamp_(-.02, 1.02)
            log_scales.clamp_(math.log(.25), math.log(max_scale))
            opacity.clamp_(-8, 8)
        factor = .12**(step / args.steps)
        for group,lr in zip(optimizer.param_groups,initial_lrs):
            group['lr'] = lr*factor
    with torch.no_grad():
        for name, param in [('xy',xy),('rgb',rgb),('log_scales',log_scales),('angle',angle),('opacity',opacity)]:
            param.copy_(best[name])
        final = render()
        torch.cuda.synchronize()
        runtime = time.perf_counter()-started
        png(final, out/'reconstruction.png')
        Image.fromarray((final.clamp(0,1).cpu().numpy()*255).round().astype('uint8')).save(out/'reconstruction.webp',quality=95)
        absolute = (final-target).abs().mean(-1)
        # Orange heatmap; fixed 8x error scale, not individually normalized.
        e = (absolute*8).clamp(0,1)
        png(torch.stack([e, e*.45, e*.12],-1),out/'error.png')
        saved = {**{k:v.cpu() for k,v in best.items()},'width':w,'height':h,
                 'gsplat_version':gsplat.__version__,'seed':args.seed}
        torch.save(saved,results/'gaussians.pt')
        np.savez_compressed(out/'gaussians.npz',
                            **{k:v.numpy() if hasattr(v,'numpy') else v for k,v in saved.items()})
        stride = max(1, math.ceil(n/3000))
        selected = torch.arange(0,n,stride,device='cuda')
        data = torch.cat([xy[selected]*torch.tensor([w,h],device='cuda'),
                          log_scales[selected].exp(), angle[selected,None],
                          rgb[selected].sigmoid(), opacity[selected,None].sigmoid()],-1).cpu().numpy()
        (out/'splats.json').write_text(json.dumps({'columns':['x','y','sx','sy','angle','r','g','b','opacity'],
             'count_total':n,'subset_stride':stride,'splats':np.round(data,4).tolist()},separators=(',',':')))
        metrics = {'width':w,'height':h,'gaussians':n,'steps':args.steps,'seed':args.seed,
                   'psnr_db':-10*math.log10(((final-target)**2).mean().item()),
                   'mae_rgb':absolute.mean().item(), 'fit_seconds':runtime,
                   'peak_allocated_mb':torch.cuda.max_memory_allocated()/1024**2,
                   'gpu':torch.cuda.get_device_name(),'torch_version':torch.__version__,
                   'gsplat_version':gsplat.__version__,
                   'trainable_parameters':9*n,'best_checkpoint_step':best_step,
                   'initial_grid':[ny,nx],'detail_points':n_extra,
                   'max_scale_pixels':max_scale,'initial_learning_rates':initial_lrs,
                   'input_sha256':hashlib.sha256(Path(args.image).read_bytes()).hexdigest(),
                   'method':f'{n:,} coplanar anisotropic Gaussians; fixed camera; RGB MSE; image-informed initialization; Adam',
                   'evaluation':'In-sample reconstruction of the single supplied image, not held-out accuracy or 3D recovery',
                   'history':history,'frames':frames}
        (out/'metrics.json').write_text(json.dumps(metrics,indent=2))
        print(json.dumps({'status':'complete',**{k:v for k,v in metrics.items() if k not in ('history','frames')}}),flush=True)


if __name__ == '__main__':
    main()
