#!/usr/bin/env python3
"""Minimal LaMa inference using downloaded big-lama.pt"""
import torch
import cv2
import numpy as np
from pathlib import Path

# Load model
model_path = Path("/data/passport/models/big-lama.pt")
model = torch.jit.load(model_path, map_location="cpu")
model.eval()

# Load image & mask
img = cv2.imread("/passport/模板/m03.jpeg")
mask = cv2.imread("/root/m03_mask.png", 0)

# Preprocess
img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
h, w = img.shape[:2]

# Pad to multiple of 8
def pad_to_multiple(x, m=8):
    h, w = x.shape[:2]
    nh, nw = ((h + m - 1) // m) * m, ((w + m - 1) // m) * m
    pad_h, pad_w = nh - h, nw - w
    return np.pad(x, ((0, pad_h), (0, pad_w), (0, 0)) if x.ndim == 3 else ((0, pad_h), (0, pad_w)), mode="reflect"), (pad_h, pad_w)

img_pad, (ph, pw) = pad_to_multiple(img)
mask_pad, _ = pad_to_multiple(mask)

# Normalize
img_tensor = torch.from_numpy(img_pad).float() / 255.0
img_tensor = img_tensor.permute(2, 0, 1).unsqueeze(0)
mask_tensor = torch.from_numpy(mask_pad).float() / 255.0
mask_tensor = mask_tensor.unsqueeze(0).unsqueeze(0)

# Inference
with torch.no_grad():
    out = model(img_tensor, mask_tensor)

# Postprocess
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)
out = out[:h, :w]  # Remove padding

cv2.imwrite("/root/m03_clean.png", out)
print("✓ Saved /root/m03_clean.png")