#!/usr/bin/env python3
import argparse, torch, cv2, numpy as np
from pathlib import Path

def pad(x, m=8):
    h,w = x.shape[:2]
    nh,nw = ((h+m-1)//m)*m, ((w+m-1)//m)*m
    ph,pw = nh-h, nw-w
    if x.ndim==3: return np.pad(x,((0,ph),(0,pw),(0,0)),mode="reflect"),(ph,pw)
    return np.pad(x,((0,ph),(0,pw)),mode="reflect"),(ph,pw)

def main():
    p = argparse.ArgumentParser()
    p.add_argument("--img", required=True); p.add_argument("--mask", required=True); p.add_argument("--out", required=True)
    p.add_argument("--model", default="/data/port/models/big-lama.pt")
    args = p.parse_args()
    model = torch.jit.load(args.model, map_location="cpu"); model.eval()
    img = cv2.imread(args.img); mask = cv2.imread(args.mask, 0)
    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB); h,w = img.shape[:2]
    img_pad,_ = pad(img); mask_pad,_ = pad(mask)
    img_t = torch.from_numpy(img_pad).float()/255.0; img_t = img_t.permute(2,0,1).unsqueeze(0)
    mask_t = torch.from_numpy(mask_pad).float()/255.0; mask_t = mask_t.unsqueeze(0).unsqueeze(0)
    with torch.no_grad(): out = model(img_t, mask_t)
    out = out[0].permute(1,2,0).cpu().numpy()
    out = np.clip(out*255,0,255).astype(np.uint8)
    out = cv2.cvtColor(out, cv2.COLOR_RGB2BGR)[:h,:w]
    cv2.imwrite(args.out, out)
    print(f"✓ {args.out}")

if __name__ == "__main__": main()
