import os, sys, time
os.environ.setdefault("PYTORCH_ENABLE_MPS_FALLBACK", "1")
import torch, numpy as np
from PIL import Image

DEVICE = sys.argv[1] if len(sys.argv) > 1 else "cpu"
DTYPE  = torch.float16 if (len(sys.argv) > 2 and sys.argv[2] == "fp16") else torch.float32
STEPS  = int(sys.argv[3]) if len(sys.argv) > 3 else 20
OCT    = int(sys.argv[4]) if len(sys.argv) > 4 else 160

from hy3dgen.rembg import BackgroundRemover
from hy3dgen.shapegen import Hunyuan3DDiTFlowMatchingPipeline

IMG = sys.argv[5] if len(sys.argv) > 5 else "/Users/danicosta/Desktop/Atapuerca/3D/prep/single/front_hero.jpg"
img = Image.open(IMG)
if img.mode != "RGBA":
    img = BackgroundRemover()(img.convert("RGB"))
print(f"[test] device={DEVICE} dtype={DTYPE} steps={STEPS} octree={OCT}")
pipe = Hunyuan3DDiTFlowMatchingPipeline.from_pretrained("tencent/Hunyuan3D-2", device=DEVICE, dtype=DTYPE)
t0 = time.time()
gen = torch.Generator(device="cpu").manual_seed(42)
mesh = pipe(image=img, num_inference_steps=STEPS, octree_resolution=OCT, guidance_scale=7.5, generator=gen)[0]
dt = time.time() - t0
v = np.asarray(mesh.vertices); lo = v.min(0); hi = v.max(0); eps = (hi - lo).max() * 0.01
on_bbox = (((np.abs(v - lo) < eps) | (np.abs(v - hi) < eps)).any(1)).mean()
verdict = "CUBO/FALLO" if on_bbox > 0.5 else "CABEZA OK"
print(f"[test] {dt:.0f}s  verts={len(v)}  frac_on_bbox={on_bbox:.3f}  => {verdict}")
