ai-3d-generator / app.py
sarahyung's picture
Update app.py
96d9839 verified
Raw
History Blame Contribute Delete
3.77 kB
"""
AI 3D Generator - TripoSG + rembg (Fixed for HF Space)
"""
import os
import sys
import uuid
import time
import logging
import tempfile
from pathlib import Path
from typing import List
import gradio as gr
import numpy as np
import torch
from PIL import Image
import trimesh
from rembg import remove
# ZeroGPU support
try:
import spaces
except ImportError:
class spaces:
@staticmethod
def GPU(fn=None, duration=None):
if fn is None:
return lambda f: f
return fn
logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
logger = logging.getLogger(__name__)
# Config
MODEL_ID = "VAST-AI/TripoSG"
TRIPOSG_REPO_URL = "https://github.com/VAST-AI-Research/TripoSG.git"
TRIPOSG_CODE_DIR = "triposg_repo"
OUTPUT_DIR = Path(tempfile.gettempdir()) / "triposg_output"
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
# Clone repo
if not os.path.exists(TRIPOSG_CODE_DIR):
logger.info("Cloning TripoSG...")
os.system(f"git clone {TRIPOSG_REPO_URL} {TRIPOSG_CODE_DIR}")
sys.path.insert(0, TRIPOSG_CODE_DIR)
sys.path.insert(0, os.path.join(TRIPOSG_CODE_DIR, "scripts"))
# Global
triposg_pipeline = None
device = "cuda" if torch.cuda.is_available() else "cpu"
def load_models():
global triposg_pipeline
if triposg_pipeline is not None:
return
logger.info("Loading TripoSG...")
try:
from huggingface_hub import snapshot_download
from triposg.pipelines.pipeline_triposg import TripoSGPipeline
model_path = snapshot_download(MODEL_ID)
triposg_pipeline = TripoSGPipeline.from_pretrained(model_path).to(device, torch.float16 if device == "cuda" else torch.float32)
logger.info("TripoSG loaded")
except Exception as e:
logger.error(f"Load failed: {e}")
triposg_pipeline = None
def remove_background(image: Image.Image) -> Image.Image:
try:
return remove(image)
except Exception as e:
logger.warning(f"rembg failed: {e}")
return image
# ================== 下面貼返你原本其他函數 ==================
# 請你從原本 app.py copy 下面這些函數貼入去:
# select_best_view, create_composite_view, extract_dimensions_from_prompt, repair_mesh, scale_mesh, segment_mesh_parts, export_mesh
# (如果你冇,我可以提供最小版 placeholder)
# ================== Main Pipeline ==================
@spaces.GPU(duration=300)
def generate_3d_model(input_images: list, text_prompt: str, num_faces: int = 50000, export_format: str = "glb", progress=gr.Progress()):
load_models()
if triposg_pipeline is None:
raise gr.Error("Model load failed. Check logs.")
# ... 你原本 generate 邏輯 ...
# 暫時用 placeholder
return None, None, "Pipeline running... (add your full logic here)"
# ================== UI ==================
with gr.Blocks(title="AI 3D Generator") as demo:
gr.Markdown("# 🚀 TripoSG 3D Generator (rembg version)")
with gr.Row():
with gr.Column():
input_images = gr.File(label="Upload Images", file_count="multiple", file_types=["image"])
text_prompt = gr.Textbox(label="Text Prompt (dimensions etc)", lines=3)
num_faces = gr.Slider(5000, 90000, value=50000, step=5000, label="Target Faces")
btn = gr.Button("Generate 3D Model", variant="primary")
with gr.Column():
preview = gr.Model3D(label="3D Preview")
download = gr.File(label="Download Model")
status = gr.Textbox(label="Status", lines=8)
btn.click(
generate_3d_model,
inputs=[input_images, text_prompt, num_faces],
outputs=[preview, download, status]
)
demo.queue(max_size=5)
demo.launch(show_api=True)