simonpick/TRELLIS_2_4K
2
1import gradio as gr2from gradio_client import Client, handle_file3import spaces4 5import os6os.environ["OPENCV_IO_ENABLE_OPENEXR"] = '1'7os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"8os.environ["ATTN_BACKEND"] = "flash_attn_3"9os.environ["FLEX_GEMM_AUTOTUNE_CACHE_PATH"] = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'autotune_cache.json')10os.environ["FLEX_GEMM_AUTOTUNER_VERBOSE"] = '1'11from datetime import datetime12import shutil13import cv214from typing import *15import torch16import numpy as np17from PIL import Image18import base6419import io20import tempfile21from trellis2.modules.sparse import SparseTensor22from trellis2.pipelines import Trellis2ImageTo3DPipeline23from trellis2.renderers import EnvMap24from trellis2.utils import render_utils25import o_voxel26 27 28MAX_SEED = np.iinfo(np.int32).max29TMP_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), 'tmp')30MODES = [31 {"name": "Normal", "icon": "assets/app/normal.png", "render_key": "normal"},32 {"name": "Clay", "icon": "assets/app/clay.png", "render_key": "clay"},33 {"name": "Color", "icon": "assets/app/basecolor.png", "render_key": "base_color"},34 {"name": "Forest", "icon": "assets/app/hdri_forest.png", "render_key": "shaded_forest"},35 {"name": "Sunset", "icon": "assets/app/hdri_sunset.png", "render_key": "shaded_sunset"},36 {"name": "Courtyard", "icon": "assets/app/hdri_courtyard.png", "render_key": "shaded_courtyard"},37]38STEPS = 839DEFAULT_MODE = 340DEFAULT_STEP = 341 42 43css = """44/* ═══════════════════════════════════════════════════════════════45 TRELLIS.2 — Modern Dark Theme46 ═══════════════════════════════════════════════════════════════ */47 48:root {49 --accent: #6366f1;50 --accent-hover: #818cf8;51 --accent-glow: rgba(99, 102, 241, 0.3);52 --surface-0: #0a0a0b;53 --surface-1: #111113;54 --surface-2: #1a1a1d;55 --surface-3: #242428;56 --border: rgba(255, 255, 255, 0.06);57 --text-primary: #fafafa;58 --text-secondary: rgba(255, 255, 255, 0.5);59 --radius: 16px;60 --radius-sm: 10px;61}62 63/* Global Overrides */64.gradio-container {65 background: var(--surface-0) !important;66 width: 100% !important;67 min-width: 800px !important;68 max-width: 1800px !important;69 margin: 0 auto !important;70 padding: 0 40px !important;71 box-sizing: border-box !important;72}73 74.gradio-container > .main {75 gap: 0 !important;76 width: 100% !important;77 max-width: none !important;78}79 80.contain {81 display: flex !important;82 flex-direction: column !important;83 max-width: none !important;84}85 86.dark {87 --block-background-fill: var(--surface-1) !important;88 --block-border-color: var(--border) !important;89 --body-background-fill: var(--surface-0) !important;90 --color-accent: var(--accent) !important;91}92 93/* Header */94.app-header {95 text-align: center;96 padding: 48px 20px 36px;97 border-bottom: 1px solid var(--border);98 margin-bottom: 32px;99 width: 100%;100}101 102.app-header h1 {103 font-family: 'SF Pro Display', -apple-system, BlinkMacSystemFont, sans-serif;104 font-size: 2.5rem;105 font-weight: 600;106 letter-spacing: -0.03em;107 background: linear-gradient(135deg, #fff 0%, rgba(255,255,255,0.7) 100%);108 -webkit-background-clip: text;109 -webkit-text-fill-color: transparent;110 margin: 0 0 8px 0;111}112 113.app-header p {114 color: var(--text-secondary);115 font-size: 1rem;116 margin: 0;117 font-weight: 400;118}119 120/* Panels */121.panel {122 background: var(--surface-1) !important;123 border: 1px solid var(--border) !important;124 border-radius: var(--radius) !important;125 overflow: hidden;126}127 128.panel-title {129 font-size: 0.7rem;130 text-transform: uppercase;131 letter-spacing: 0.1em;132 color: var(--text-secondary);133 padding: 16px 20px 8px;134 font-weight: 600;135}136 137/* Upload Area */138.upload-zone {139 min-height: 280px !important;140 border: 2px dashed var(--border) !important;141 border-radius: var(--radius) !important;142 background: var(--surface-2) !important;143 transition: all 0.3s ease;144}145 146.upload-zone:hover {147 border-color: var(--accent) !important;148 background: rgba(99, 102, 241, 0.05) !important;149}150 151/* Buttons */152.primary-btn {153 background: var(--accent) !important;154 border: none !important;155 border-radius: var(--radius-sm) !important;156 color: white !important;157 font-weight: 600 !important;158 padding: 14px 28px !important;159 font-size: 0.95rem !important;160 transition: all 0.2s ease !important;161 box-shadow: 0 4px 20px var(--accent-glow) !important;162}163 164.primary-btn:hover {165 background: var(--accent-hover) !important;166 transform: translateY(-1px);167 box-shadow: 0 6px 30px var(--accent-glow) !important;168}169 170.secondary-btn {171 background: var(--surface-3) !important;172 border: 1px solid var(--border) !important;173 border-radius: var(--radius-sm) !important;174 color: var(--text-primary) !important;175 font-weight: 500 !important;176 transition: all 0.2s ease !important;177}178 179.secondary-btn:hover {180 background: var(--surface-2) !important;181 border-color: var(--accent) !important;182}183 184/* Sliders & Inputs */185input[type="range"] {186 accent-color: var(--accent) !important;187}188 189.wrap input, .wrap textarea {190 background: var(--surface-2) !important;191 border: 1px solid var(--border) !important;192 border-radius: var(--radius-sm) !important;193 color: var(--text-primary) !important;194}195 196/* Radio Buttons */197.gr-radio-row {198 gap: 8px !important;199}200 201.gr-radio-row label {202 background: var(--surface-2) !important;203 border: 1px solid var(--border) !important;204 border-radius: var(--radius-sm) !important;205 padding: 10px 18px !important;206 transition: all 0.2s ease !important;207}208 209.gr-radio-row label:hover {210 border-color: var(--accent) !important;211}212 213.gr-radio-row label.selected {214 background: var(--accent) !important;215 border-color: var(--accent) !important;216}217 218/* Accordion */219.gr-accordion {220 border: 1px solid var(--border) !important;221 border-radius: var(--radius-sm) !important;222 background: var(--surface-2) !important;223}224 225/* Walkthrough/Stepper */226.stepper-wrapper { padding: 0; }227.stepper-container { padding: 0; align-items: center; }228.step-button { flex-direction: row; }229.step-connector { transform: none; }230.step-number { width: 16px; height: 16px; }231.step-label { position: relative; bottom: 0; }232 233/* Loading States */234.wrap.center.full { inset: 0; height: 100%; }235.wrap.center.full.translucent { background: var(--surface-1); }236 237/* ═══════════════════════════════════════════════════════════════238 3D PREVIEWER COMPONENT239 ═══════════════════════════════════════════════════════════════ */240 241.previewer-container {242 position: relative;243 width: 100%;244 height: 720px;245 display: flex;246 flex-direction: column;247 align-items: center;248 justify-content: center;249 padding: 24px;250 background: radial-gradient(ellipse at center, var(--surface-2) 0%, var(--surface-1) 100%);251 border-radius: var(--radius);252}253 254/* Viewport */255.previewer-container .display-row {256 flex: 1;257 width: 100%;258 display: flex;259 justify-content: center;260 align-items: center;261 min-height: 0;262}263 264.previewer-container .previewer-main-image {265 max-width: 100%;266 max-height: 100%;267 object-fit: contain;268 display: none;269 border-radius: var(--radius-sm);270 box-shadow: 0 20px 60px rgba(0, 0, 0, 0.4);271}272 273.previewer-container .previewer-main-image.visible {274 display: block;275 animation: fadeIn 0.3s ease;276}277 278@keyframes fadeIn {279 from { opacity: 0; transform: scale(0.98); }280 to { opacity: 1; transform: scale(1); }281}282 283/* Mode Selector */284.previewer-container .mode-row {285 display: flex;286 gap: 10px;287 margin-top: 20px;288 padding: 8px;289 background: var(--surface-0);290 border-radius: 50px;291 border: 1px solid var(--border);292}293 294.previewer-container .mode-btn {295 width: 32px;296 height: 32px;297 border-radius: 50%;298 cursor: pointer;299 opacity: 0.4;300 transition: all 0.25s cubic-bezier(0.4, 0, 0.2, 1);301 border: 2px solid transparent;302 object-fit: cover;303}304 305.previewer-container .mode-btn:hover {306 opacity: 0.8;307 transform: scale(1.1);308}309 310.previewer-container .mode-btn.active {311 opacity: 1;312 border-color: var(--accent);313 transform: scale(1.15);314 box-shadow: 0 0 20px var(--accent-glow);315}316 317/* Rotation Slider */318.previewer-container .slider-row {319 width: 100%;320 max-width: 320px;321 margin-top: 16px;322}323 324.previewer-container input[type=range] {325 -webkit-appearance: none;326 width: 100%;327 background: transparent;328 cursor: pointer;329}330 331.previewer-container input[type=range]::-webkit-slider-runnable-track {332 width: 100%;333 height: 6px;334 background: var(--surface-0);335 border-radius: 3px;336 border: 1px solid var(--border);337}338 339.previewer-container input[type=range]::-webkit-slider-thumb {340 -webkit-appearance: none;341 height: 18px;342 width: 18px;343 border-radius: 50%;344 background: var(--accent);345 margin-top: -7px;346 box-shadow: 0 2px 10px var(--accent-glow);347 transition: transform 0.15s ease;348}349 350.previewer-container input[type=range]::-webkit-slider-thumb:hover {351 transform: scale(1.2);352}353 354/* Empty State */355.empty-state {356 display: flex;357 flex-direction: column;358 align-items: center;359 gap: 16px;360 color: var(--text-secondary);361}362 363.empty-state svg {364 opacity: 0.3;365}366 367.empty-state p {368 font-size: 0.9rem;369 margin: 0;370}371 372/* Block Label Override */373.gradio-container .padded:has(.previewer-container) { padding: 0 !important; }374.gradio-container:has(.previewer-container) [data-testid="block-label"] {375 position: absolute;376 top: 0;377 left: 0;378}379 380/* GLB Viewer */381.model3d-container {382 background: var(--surface-2) !important;383 border-radius: var(--radius) !important;384}385 386/* Footer Note */387.footer-note {388 text-align: center;389 color: var(--text-secondary);390 font-size: 0.8rem;391 padding: 20px;392 border-top: 1px solid var(--border);393 margin-top: 32px;394 width: 100%;395}396 397/* Main Layout - Force side by side */398#main-row {399 width: 100% !important;400 max-width: none !important;401 margin: 0 !important;402 display: flex !important;403 flex-direction: row !important;404 flex-wrap: nowrap !important;405 gap: 32px !important;406 align-items: flex-start !important;407}408 409#main-row.row {410 flex-wrap: nowrap !important;411 max-width: none !important;412}413 414#input-col {415 flex: 0 0 400px !important;416 width: 400px !important;417 min-width: 350px !important;418 max-width: 450px !important;419}420 421#preview-col {422 flex: 1 1 auto !important;423 min-width: 500px !important;424}425 426@media (max-width: 900px) {427 #main-row {428 flex-direction: column !important;429 }430 #input-col,431 #preview-col {432 flex: 1 1 auto !important;433 width: 100% !important;434 max-width: 100% !important;435 min-width: 0 !important;436 }437}438"""439 440 441head = """442<link rel="preconnect" href="https://fonts.googleapis.com">443<link rel="preconnect" href="https://fonts.gstatic.com" crossorigin>444<link href="https://fonts.googleapis.com/css2?family=Inter:wght@400;500;600;700&display=swap" rel="stylesheet">445 446<script>447 function refreshView(mode, step) {448 const allImgs = document.querySelectorAll('.previewer-main-image');449 for (let i = 0; i < allImgs.length; i++) {450 const img = allImgs[i];451 if (img.classList.contains('visible')) {452 const id = img.id;453 const [_, m, s] = id.split('-');454 if (mode === -1) mode = parseInt(m.slice(1));455 if (step === -1) step = parseInt(s.slice(1));456 break;457 }458 }459 460 allImgs.forEach(img => img.classList.remove('visible'));461 const targetId = 'view-m' + mode + '-s' + step;462 const targetImg = document.getElementById(targetId);463 if (targetImg) targetImg.classList.add('visible');464 465 const allBtns = document.querySelectorAll('.mode-btn');466 allBtns.forEach((btn, idx) => {467 if (idx === mode) btn.classList.add('active');468 else btn.classList.remove('active');469 });470 }471 472 function selectMode(mode) { refreshView(mode, -1); }473 function onSliderChange(val) { refreshView(-1, parseInt(val)); }474</script>475"""476 477 478empty_html = """479<div class="previewer-container">480 <div class="empty-state">481 <svg width="64" height="64" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="1.5" stroke-linecap="round" stroke-linejoin="round">482 <path d="M21 16V8a2 2 0 0 0-1-1.73l-7-4a2 2 0 0 0-2 0l-7 4A2 2 0 0 0 3 8v8a2 2 0 0 0 1 1.73l7 4a2 2 0 0 0 2 0l7-4A2 2 0 0 0 21 16z"></path>483 <polyline points="3.27 6.96 12 12.01 20.73 6.96"></polyline>484 <line x1="12" y1="22.08" x2="12" y2="12"></line>485 </svg>486 <p>Upload an image to generate 3D</p>487 </div>488</div>489"""490 491 492def image_to_base64(image):493 buffered = io.BytesIO()494 image = image.convert("RGB")495 image.save(buffered, format="jpeg", quality=85)496 img_str = base64.b64encode(buffered.getvalue()).decode()497 return f"data:image/jpeg;base64,{img_str}"498 499 500def start_session(req: gr.Request):501 user_dir = os.path.join(TMP_DIR, str(req.session_hash))502 os.makedirs(user_dir, exist_ok=True)503 504 505def end_session(req: gr.Request):506 user_dir = os.path.join(TMP_DIR, str(req.session_hash))507 shutil.rmtree(user_dir)508 509 510def remove_background(input: Image.Image) -> Image.Image:511 with tempfile.NamedTemporaryFile(suffix='.png') as f:512 input = input.convert('RGB')513 input.save(f.name)514 output = rmbg_client.predict(handle_file(f.name), api_name="/image")[0][0]515 output = Image.open(output)516 return output517 518 519def preprocess_image(input: Image.Image) -> Image.Image:520 """Preprocess the input image."""521 has_alpha = False522 if input.mode == 'RGBA':523 alpha = np.array(input)[:, :, 3]524 if not np.all(alpha == 255):525 has_alpha = True526 max_size = max(input.size)527 scale = min(1, 1024 / max_size)528 if scale < 1:529 input = input.resize((int(input.width * scale), int(input.height * scale)), Image.Resampling.LANCZOS)530 if has_alpha:531 output = input532 else:533 output = remove_background(input)534 output_np = np.array(output)535 alpha = output_np[:, :, 3]536 bbox = np.argwhere(alpha > 0.8 * 255)537 bbox = np.min(bbox[:, 1]), np.min(bbox[:, 0]), np.max(bbox[:, 1]), np.max(bbox[:, 0])538 center = (bbox[0] + bbox[2]) / 2, (bbox[1] + bbox[3]) / 2539 size = max(bbox[2] - bbox[0], bbox[3] - bbox[1])540 size = int(size * 1)541 bbox = center[0] - size // 2, center[1] - size // 2, center[0] + size // 2, center[1] + size // 2542 output = output.crop(bbox)543 output = np.array(output).astype(np.float32) / 255544 output = output[:, :, :3] * output[:, :, 3:4]545 output = Image.fromarray((output * 255).astype(np.uint8))546 return output547 548 549def pack_state(latents: Tuple[SparseTensor, SparseTensor, int]) -> dict:550 shape_slat, tex_slat, res = latents551 return {552 'shape_slat_feats': shape_slat.feats.cpu().numpy(),553 'tex_slat_feats': tex_slat.feats.cpu().numpy(),554 'coords': shape_slat.coords.cpu().numpy(),555 'res': res,556 }557 558 559def unpack_state(state: dict) -> Tuple[SparseTensor, SparseTensor, int]:560 shape_slat = SparseTensor(561 feats=torch.from_numpy(state['shape_slat_feats']).cuda(),562 coords=torch.from_numpy(state['coords']).cuda(),563 )564 tex_slat = shape_slat.replace(torch.from_numpy(state['tex_slat_feats']).cuda())565 return shape_slat, tex_slat, state['res']566 567 568@spaces.GPU(duration=180)569def generate_and_extract(570 image: Image.Image,571 req: gr.Request,572 progress=gr.Progress(track_tqdm=True),573) -> Tuple[str, str, str]:574 """575 Combined function: Generate 3D from image AND extract GLB in one GPU session.576 This avoids issues with chaining multiple @spaces.GPU functions.577 """578 user_dir = os.path.join(TMP_DIR, str(req.session_hash))579 os.makedirs(user_dir, exist_ok=True)580 581 # Hardcoded values582 seed = np.random.randint(0, MAX_SEED)583 decimation_target = 300000584 texture_size = 4096585 586 # === STAGE 1: Generate 3D ===587 outputs, latents = pipeline.run(588 image,589 seed=seed,590 preprocess_image=False,591 sparse_structure_sampler_params={592 "steps": 12,593 "guidance_strength": 7.5,594 "guidance_rescale": 0.7,595 "rescale_t": 5.0,596 },597 shape_slat_sampler_params={598 "steps": 12,599 "guidance_strength": 7.5,600 "guidance_rescale": 0.5,601 "rescale_t": 3.0,602 },603 tex_slat_sampler_params={604 "steps": 12,605 "guidance_strength": 1.0,606 "guidance_rescale": 0.0,607 "rescale_t": 3.0,608 },609 pipeline_type="1024_cascade",610 return_latent=True,611 )612 mesh = outputs[0]613 mesh.simplify(16777216)614 615 # Render preview images616 images = render_utils.render_snapshot(mesh, resolution=1024, r=2, fov=36, nviews=STEPS, envmap=envmap)617 618 # Build preview HTML619 images_html = ""620 for m_idx, mode in enumerate(MODES):621 for s_idx in range(STEPS):622 unique_id = f"view-m{m_idx}-s{s_idx}"623 is_visible = (m_idx == DEFAULT_MODE and s_idx == DEFAULT_STEP)624 vis_class = "visible" if is_visible else ""625 img_base64 = image_to_base64(Image.fromarray(images[mode['render_key']][s_idx]))626 images_html += f'<img id="{unique_id}" class="previewer-main-image {vis_class}" src="{img_base64}" loading="eager">'627 628 btns_html = ""629 for idx, mode in enumerate(MODES): 630 active_class = "active" if idx == DEFAULT_MODE else ""631 btns_html += f'<img src="{mode["icon_base64"]}" class="mode-btn {active_class}" onclick="selectMode({idx})" title="{mode["name"]}">'632 633 preview_html = f"""634 <div class="previewer-container">635 <div class="display-row">{images_html}</div>636 <div class="mode-row">{btns_html}</div>637 <div class="slider-row">638 <input type="range" min="0" max="{STEPS - 1}" value="{DEFAULT_STEP}" step="1" oninput="onSliderChange(this.value)">639 </div>640 </div>641 """642 643 # === STAGE 2: Extract GLB ===644 shape_slat, tex_slat, res = latents645 mesh = pipeline.decode_latent(shape_slat, tex_slat, res)[0]646 mesh.simplify(16777216)647 648 glb = o_voxel.postprocess.to_glb(649 vertices=mesh.vertices,650 faces=mesh.faces,651 attr_volume=mesh.attrs,652 coords=mesh.coords,653 attr_layout=pipeline.pbr_attr_layout,654 grid_size=res,655 aabb=[[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]],656 decimation_target=decimation_target,657 texture_size=texture_size,658 remesh=True,659 remesh_band=1,660 remesh_project=0,661 use_tqdm=True,662 )663 664 now = datetime.now()665 timestamp = now.strftime("%Y-%m-%dT%H%M%S") + f".{now.microsecond // 1000:03d}"666 glb_path = os.path.join(user_dir, f'sample_{timestamp}.glb')667 glb.export(glb_path, extension_webp=True)668 669 torch.cuda.empty_cache()670 671 # Return: preview_html, glb_path (for viewer), glb_path (for download)672 return preview_html, glb_path, glb_path673 674 675# ═══════════════════════════════════════════════════════════════676# GRADIO INTERFACE677# ═══════════════════════════════════════════════════════════════678 679with gr.Blocks(theme=gr.themes.Base(primary_hue="indigo"), delete_cache=(600, 600)) as demo:680 681 # Header682 gr.HTML("""683 <div class="app-header">684 <h1>TRELLIS.2</h1>685 <p>Transform any image into a high-quality 3D asset</p>686 </div>687 """)688 689 with gr.Row(equal_height=False, elem_id="main-row"):690 # Left Panel — Input (span 1)691 with gr.Column(scale=1, min_width=320, elem_id="input-col"):692 693 # Image Upload694 image_prompt = gr.Image(695 label="Input Image",696 format="png",697 image_mode="RGBA",698 type="pil",699 height=400,700 elem_classes=["upload-zone"]701 )702 703 # Generate Button704 generate_btn = gr.Button("Generate 3D", variant="primary", elem_classes=["primary-btn"], size="lg")705 706 # Right Panel — Preview (span 2)707 with gr.Column(scale=2, elem_id="preview-col"):708 with gr.Walkthrough(selected=0) as walkthrough:709 with gr.Step("Preview", id=0):710 preview_output = gr.HTML(empty_html, label="3D Preview", show_label=False)711 712 with gr.Step("Export", id=1):713 glb_output = gr.Model3D(714 label="GLB Model",715 height=640,716 show_label=False,717 display_mode="solid",718 clear_color=(0.06, 0.06, 0.07, 0.0) # Alpha = 0 for transparent background719 )720 download_btn = gr.DownloadButton("Download GLB", elem_classes=["primary-btn"], size="lg")721 722 # Footer723 gr.HTML('<div class="footer-note">Generation includes automatic GLB extraction. This may take 90+ seconds total.</div>')724 725 # Event Handlers726 demo.load(start_session)727 demo.unload(end_session)728 729 image_prompt.upload(730 preprocess_image,731 inputs=[image_prompt],732 outputs=[image_prompt],733 )734 735 # Single GPU call: Generate 3D + Extract GLB736 generate_btn.click(737 generate_and_extract,738 inputs=[image_prompt],739 outputs=[preview_output, glb_output, download_btn],740 ).then(741 lambda: gr.Walkthrough(selected=1), outputs=walkthrough742 )743 744 745# ═══════════════════════════════════════════════════════════════746# LAUNCH747# ═══════════════════════════════════════════════════════════════748 749if __name__ == "__main__":750 os.makedirs(TMP_DIR, exist_ok=True)751 752 # Load mode icons753 for i in range(len(MODES)):754 icon = Image.open(MODES[i]['icon'])755 MODES[i]['icon_base64'] = image_to_base64(icon)756 757 rmbg_client = Client("briaai/BRIA-RMBG-2.0")758 pipeline = Trellis2ImageTo3DPipeline.from_pretrained('microsoft/TRELLIS.2-4B')759 pipeline.rembg_model = None760 pipeline.low_vram = False761 pipeline.cuda()762 763 envmap = {764 'forest': EnvMap(torch.tensor(765 cv2.cvtColor(cv2.imread('assets/hdri/forest.exr', cv2.IMREAD_UNCHANGED), cv2.COLOR_BGR2RGB),766 dtype=torch.float32, device='cuda'767 )),768 'sunset': EnvMap(torch.tensor(769 cv2.cvtColor(cv2.imread('assets/hdri/sunset.exr', cv2.IMREAD_UNCHANGED), cv2.COLOR_BGR2RGB),770 dtype=torch.float32, device='cuda'771 )),772 'courtyard': EnvMap(torch.tensor(773 cv2.cvtColor(cv2.imread('assets/hdri/courtyard.exr', cv2.IMREAD_UNCHANGED), cv2.COLOR_BGR2RGB),774 dtype=torch.float32, device='cuda'775 )),776 }777 778 demo.launch(css=css, head=head)