#!/usr/bin/env python
"""Genera la geometría (shape) de la cabeza con Hunyuan3D-2 sobre MPS (Apple Silicon).
Solo malla — sin textura (eso necesita CUDA; y MetaHuman pone su propia piel).

Uso:
  python gen_shape.py <imagen> <octree_res> <steps>
Defaults: front_hero, 384, 50
"""
import os, sys, time
# Cualquier op no soportada en MPS cae a CPU en vez de romper
os.environ.setdefault("PYTORCH_ENABLE_MPS_FALLBACK", "1")

import torch
from PIL import Image

IMG    = sys.argv[1] if len(sys.argv) > 1 else "/Users/danicosta/Desktop/Atapuerca/3D/prep/single/head_crop.png"
OCTREE = int(sys.argv[2]) if len(sys.argv) > 2 else 384
STEPS  = int(sys.argv[3]) if len(sys.argv) > 3 else 50
GUIDANCE = float(sys.argv[4]) if len(sys.argv) > 4 else 5.0
BOX_V    = float(sys.argv[5]) if len(sys.argv) > 5 else 1.5
OUTDIR = "/Users/danicosta/Desktop/Atapuerca/3D/export"
os.makedirs(OUTDIR, exist_ok=True)

DEVICE = "mps" if torch.backends.mps.is_available() else "cpu"
DTYPE  = torch.float32
print(f"[gen] device={DEVICE} dtype={DTYPE} img={IMG} octree={OCTREE} steps={STEPS}")

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

# 1) Imagen de entrada con fondo recortado (RGBA)
img = Image.open(IMG)
if img.mode != "RGBA":
    print("[gen] quitando fondo (rembg)...")
    img = BackgroundRemover()(img.convert("RGB"))
print(f"[gen] imagen lista: {img.size} {img.mode}")

# 2) Pipeline de shape
print("[gen] cargando Hunyuan3D-2 (descarga pesos la 1a vez)...")
t0 = time.time()
pipe = Hunyuan3DDiTFlowMatchingPipeline.from_pretrained(
    "tencent/Hunyuan3D-2", device=DEVICE, dtype=DTYPE)
print(f"[gen] pipeline cargado en {time.time()-t0:.0f}s")

# 3) Generación
print("[gen] generando geometría...")
t0 = time.time()
print(f"[gen] guidance={GUIDANCE} box_v={BOX_V}")
gen = torch.Generator(device="cpu").manual_seed(42)
mesh = pipe(image=img, num_inference_steps=STEPS, octree_resolution=OCTREE,
            guidance_scale=GUIDANCE, box_v=BOX_V, generator=gen)[0]
print(f"[gen] geometría generada en {time.time()-t0:.0f}s")
print(f"[gen] malla: {len(mesh.vertices)} verts, {len(mesh.faces)} caras")

# 4) Exportar
base = os.path.join(OUTDIR, f"neandertal_head_raw_o{OCTREE}")
mesh.export(base + ".glb")
mesh.export(base + ".obj")
print(f"[gen] EXPORTADO: {base}.glb  y  .obj")
