Spaces:
Runtime error
Runtime error
Download demo_release.py from Merlinus001/JoyAI-Image-Edit-Space: direct link, hf CLI and curl.
- Browser
- Download file 133 kB
-
https://huggingface.co/spaces/Merlinus001/JoyAI-Image-Edit-Space/resolve/main/demo_release.py
- Command line
-
hf download hf://spaces/Merlinus001/JoyAI-Image-Edit-Space/demo_release.py
-
curl -L -o demo_release.py https://huggingface.co/spaces/Merlinus001/JoyAI-Image-Edit-Space/resolve/main/demo_release.py
133 kB
| #!/usr/bin/env python3 | |
| # -*- coding: utf-8 -*- | |
| """JoyAI-Image spatial editing demo. | |
| This module provides the Gradio UI and runtime glue code for three tasks: | |
| object moving, object rotation, and camera control. It supports optional | |
| prompt rewriting, local three.js injection, bucket-based input resizing, and | |
| persistent run export for debugging or demo collection. | |
| """ | |
| import argparse | |
| import base64 | |
| import copy | |
| import hashlib | |
| import io | |
| import json | |
| import math | |
| import os | |
| import random | |
| import re | |
| import sys | |
| import tempfile | |
| import threading | |
| import time | |
| import warnings | |
| from datetime import datetime | |
| from pathlib import Path | |
| from typing import Any, Dict, List, Optional, Tuple | |
| import gradio as gr | |
| ROOT_DIR = Path(__file__).resolve().parent | |
| SRC_DIR = ROOT_DIR / "src" | |
| if str(SRC_DIR) not in sys.path: | |
| sys.path.insert(0, str(SRC_DIR)) | |
| try: | |
| import spaces | |
| except Exception: | |
| class _SpacesFallback: | |
| def GPU(fn=None, *args, **kwargs): | |
| if fn is None: | |
| def decorator(func): | |
| return func | |
| return decorator | |
| return fn | |
| spaces = _SpacesFallback() | |
| import numpy as np | |
| import torch | |
| import torchvision.transforms.functional as TF | |
| from PIL import Image, ImageDraw, ImageFile, ImageFilter, ImageOps | |
| from modules.models.bucket import BucketGroup, generate_video_image_bucket | |
| for env_name, env_value in { | |
| "HF_HUB_OFFLINE": "1", | |
| "TRANSFORMERS_OFFLINE": "1", | |
| "HF_DATASETS_OFFLINE": "1", | |
| "HF_HUB_DISABLE_TELEMETRY": "1", | |
| "GRADIO_ANALYTICS_ENABLED": "False", | |
| "WANDB_DISABLED": "true", | |
| "DO_NOT_TRACK": "1", | |
| "TOKENIZERS_PARALLELISM": "false", | |
| }.items(): | |
| os.environ.setdefault(env_name, env_value) | |
| warnings.filterwarnings("ignore") | |
| def is_rank0() -> bool: | |
| return int(os.environ.get("RANK", "0")) == 0 | |
| def resolve_device() -> torch.device: | |
| if not torch.cuda.is_available(): | |
| return torch.device("cpu") | |
| local_rank = int(os.environ.get("LOCAL_RANK", "0")) | |
| torch.cuda.set_device(local_rank) | |
| return torch.device(f"cuda:{local_rank}") | |
| MAX_SEED = np.iinfo(np.int32).max | |
| BASE_SIZE = 1024 | |
| RUN_BUTTON_IDLE_TEXT = "🚀 Generate" | |
| RUN_BUTTON_BUSY_TEXT = "Processing..." | |
| TASK_MODE_OBJECT = "Object Moving" | |
| TASK_MODE_ROTATION = "Object Rotation" | |
| TASK_MODE_CAMERA = "Camera Control" | |
| ROTATION_VIEW_ORDER = [ | |
| "front", | |
| "front right side", | |
| "right side", | |
| "rear right side", | |
| "rear", | |
| "rear left side", | |
| "left side", | |
| "front left side", | |
| ] | |
| YAW_MIN = -90.0 | |
| YAW_MAX = 90.0 | |
| PITCH_MIN = -90.0 | |
| PITCH_MAX = 90.0 | |
| YAW_STEP = 45.0 | |
| PITCH_STEP = 22.5 | |
| AUTO_PE_DISABLE_YAW_MIN = -90.0 | |
| AUTO_PE_DISABLE_YAW_MAX = 90.0 | |
| DEFAULT_BOX_CX = 0.5 | |
| DEFAULT_BOX_CY = 0.5 | |
| DEFAULT_BOX_W = 0.3 | |
| DEFAULT_BOX_H = 0.3 | |
| def box_selector_html() -> str: | |
| return """ | |
| <style> | |
| #se-box-wrap { position: relative; width: 100%; max-width: 520px; } | |
| #se-box-img { | |
| width: 100%; display: none; border-radius: 8px; border: 1px solid #444; cursor: crosshair; | |
| user-select: none; -webkit-user-drag: none; touch-action: none; | |
| } | |
| #se-box-empty { | |
| width: 100%; min-height: 220px; display: flex; align-items: center; justify-content: center; | |
| border-radius: 8px; border: 1px dashed #666; color: #9a9a9a; font-size: 13px; text-align: center; | |
| box-sizing: border-box; padding: 16px; | |
| } | |
| #se-box-overlay { | |
| position: absolute; border: 2px solid #ff2d2d; box-sizing: border-box; pointer-events: none; display: none; | |
| box-shadow: 0 0 0 9999px rgba(255, 0, 0, 0.08); | |
| } | |
| #se-box-hint { margin-top: 6px; color: #bbb; font-size: 12px; } | |
| #se-box-debug { margin-top: 4px; color: #9a9a9a; font-size: 11px; } | |
| </style> | |
| <div id="se-box-wrap"> | |
| <div id="se-box-empty">Upload an image above, then drag here to draw the box.</div> | |
| <img id="se-box-img" alt="box selector" /> | |
| <div id="se-box-overlay"></div> | |
| </div> | |
| <div id="se-box-hint">Drag the mouse to draw the red box (real-time preview).</div> | |
| <div id="se-box-debug"></div> | |
| <script> | |
| (function(){ | |
| const scriptEl = document.currentScript; | |
| const root = (scriptEl && scriptEl.closest('.gradio-container')) || document; | |
| const img = document.getElementById('se-box-img'); | |
| const empty = document.getElementById('se-box-empty'); | |
| const debug = document.getElementById('se-box-debug'); | |
| const overlay = document.getElementById('se-box-overlay'); | |
| let dragging = false; | |
| let startX = 0, startY = 0; | |
| let currentBox = {cx:0.5, cy:0.5, w:0.3, h:0.3}; | |
| let lastHiddenValue = ""; | |
| let lastSourceSrc = ""; | |
| function clamp(v, lo, hi){ return Math.max(lo, Math.min(hi, v)); } | |
| function setDebug(msg){ if(debug) debug.textContent = msg; } | |
| function sanitizeBox(box){ | |
| const w = clamp(Number(box?.w ?? 0.3) || 0.3, 0.001, 1.0); | |
| const h = clamp(Number(box?.h ?? 0.3) || 0.3, 0.001, 1.0); | |
| const cx = clamp(Number(box?.cx ?? 0.5) || 0.5, w / 2, 1 - w / 2); | |
| const cy = clamp(Number(box?.cy ?? 0.5) || 0.5, h / 2, 1 - h / 2); | |
| return {cx, cy, w, h}; | |
| } | |
| function setOverlayFromNorm(box){ | |
| if(!img.src){ | |
| img.style.display = 'none'; | |
| if(empty) empty.style.display = 'flex'; | |
| overlay.style.display = 'none'; | |
| setDebug('source: none'); | |
| return; | |
| } | |
| img.style.display = 'block'; | |
| if(empty) empty.style.display = 'none'; | |
| const rect = img.getBoundingClientRect(); | |
| if(!rect.width || !rect.height){ | |
| overlay.style.display = 'none'; | |
| return; | |
| } | |
| const x = (box.cx - box.w / 2) * rect.width; | |
| const y = (box.cy - box.h / 2) * rect.height; | |
| const w = box.w * rect.width; | |
| const h = box.h * rect.height; | |
| overlay.style.left = `${x}px`; | |
| overlay.style.top = `${y}px`; | |
| overlay.style.width = `${w}px`; | |
| overlay.style.height = `${h}px`; | |
| overlay.style.display = 'block'; | |
| } | |
| function dispatchBox(box){ | |
| currentBox = sanitizeBox(box); | |
| const hidden = root.querySelector('#se_box_json textarea, #se_box_json input'); | |
| if(!hidden) return; | |
| hidden.value = JSON.stringify(currentBox); | |
| hidden.dispatchEvent(new Event('input', {bubbles:true})); | |
| hidden.dispatchEvent(new Event('change', {bubbles:true})); | |
| lastHiddenValue = hidden.value; | |
| } | |
| function pointToNorm(clientX, clientY){ | |
| const rect = img.getBoundingClientRect(); | |
| const x = clamp((clientX - rect.left) / rect.width, 0, 1); | |
| const y = clamp((clientY - rect.top) / rect.height, 0, 1); | |
| return {x,y}; | |
| } | |
| function beginDrag(clientX, clientY){ | |
| if(!img.src) return; | |
| dragging = true; | |
| const p = pointToNorm(clientX, clientY); | |
| startX = p.x; | |
| startY = p.y; | |
| setOverlayFromNorm(sanitizeBox({cx:startX, cy:startY, w:0.001, h:0.001})); | |
| } | |
| function updateDrag(clientX, clientY){ | |
| if(!dragging) return; | |
| const p = pointToNorm(clientX, clientY); | |
| const x1 = Math.min(startX, p.x), x2 = Math.max(startX, p.x); | |
| const y1 = Math.min(startY, p.y), y2 = Math.max(startY, p.y); | |
| setOverlayFromNorm(sanitizeBox({cx:(x1+x2)/2, cy:(y1+y2)/2, w:Math.max(0.001, x2-x1), h:Math.max(0.001, y2-y1)})); | |
| } | |
| function endDrag(clientX, clientY){ | |
| if(!dragging) return; | |
| dragging = false; | |
| const p = pointToNorm(clientX, clientY); | |
| const x1 = Math.min(startX, p.x), x2 = Math.max(startX, p.x); | |
| const y1 = Math.min(startY, p.y), y2 = Math.max(startY, p.y); | |
| const box = sanitizeBox({cx:(x1+x2)/2, cy:(y1+y2)/2, w:Math.max(0.001, x2-x1), h:Math.max(0.001, y2-y1)}); | |
| setOverlayFromNorm(box); | |
| dispatchBox(box); | |
| } | |
| img.addEventListener('mousedown', (ev)=>{ | |
| if(ev.button !== 0) return; | |
| beginDrag(ev.clientX, ev.clientY); | |
| }); | |
| window.addEventListener('mousemove', (ev)=>{ | |
| updateDrag(ev.clientX, ev.clientY); | |
| }); | |
| window.addEventListener('mouseup', (ev)=>{ | |
| endDrag(ev.clientX, ev.clientY); | |
| }); | |
| img.addEventListener('touchstart', (ev)=>{ | |
| if(!ev.touches || !ev.touches.length) return; | |
| const touch = ev.touches[0]; | |
| beginDrag(touch.clientX, touch.clientY); | |
| ev.preventDefault(); | |
| }, {passive:false}); | |
| window.addEventListener('touchmove', (ev)=>{ | |
| if(!dragging || !ev.touches || !ev.touches.length) return; | |
| const touch = ev.touches[0]; | |
| updateDrag(touch.clientX, touch.clientY); | |
| ev.preventDefault(); | |
| }, {passive:false}); | |
| window.addEventListener('touchend', (ev)=>{ | |
| const touch = (ev.changedTouches && ev.changedTouches[0]) || null; | |
| if(!touch) return; | |
| endDrag(touch.clientX, touch.clientY); | |
| ev.preventDefault(); | |
| }, {passive:false}); | |
| function syncHiddenBox(){ | |
| const hidden = root.querySelector('#se_box_json textarea, #se_box_json input'); | |
| if(!hidden) return; | |
| const value = String(hidden.value || "").trim(); | |
| if(!value || value === lastHiddenValue) return; | |
| try{ | |
| currentBox = sanitizeBox(JSON.parse(value)); | |
| setOverlayFromNorm(currentBox); | |
| lastHiddenValue = value; | |
| }catch(_err){} | |
| } | |
| function getSourceUrlFromHidden(){ | |
| const hiddenSource = | |
| root.querySelector('#se_box_src_url textarea, #se_box_src_url input') || | |
| document.querySelector('#se_box_src_url textarea, #se_box_src_url input'); | |
| if(!hiddenSource) return ""; | |
| return String(hiddenSource.value || "").trim(); | |
| } | |
| function getComponentImageSrc(componentId){ | |
| const container = | |
| root.querySelector(`#${componentId}`) || | |
| document.querySelector(`#${componentId}`); | |
| if(!container) return ""; | |
| const imgNode = container.querySelector('img'); | |
| if(imgNode && imgNode.src){ | |
| return String(imgNode.src || ""); | |
| } | |
| const canvasNode = container.querySelector('canvas'); | |
| if(canvasNode){ | |
| try{ | |
| return canvasNode.toDataURL('image/png'); | |
| }catch(_err){} | |
| } | |
| return ""; | |
| } | |
| function pullImageAndBox(){ | |
| const explicitSrc = getSourceUrlFromHidden(); | |
| if(explicitSrc){ | |
| if(img.src !== explicitSrc){ | |
| img.src = explicitSrc; | |
| lastSourceSrc = explicitSrc; | |
| setOverlayFromNorm(currentBox); | |
| } | |
| setDebug('source: hidden_data_url'); | |
| syncHiddenBox(); | |
| return; | |
| } | |
| const sourceFromHiddenImage = getComponentImageSrc('se_source_image'); | |
| const sourceFromUploader = getComponentImageSrc('se_upload_image'); | |
| const fallbackSrc = sourceFromHiddenImage || sourceFromUploader; | |
| if(!fallbackSrc){ | |
| if(lastSourceSrc){ | |
| img.removeAttribute('src'); | |
| img.style.display = 'none'; | |
| if(empty) empty.style.display = 'flex'; | |
| overlay.style.display = 'none'; | |
| lastSourceSrc = ""; | |
| } | |
| setDebug('source: unresolved'); | |
| return; | |
| } | |
| if(img.src !== fallbackSrc){ | |
| img.src = fallbackSrc; | |
| lastSourceSrc = fallbackSrc; | |
| setOverlayFromNorm(currentBox); | |
| } | |
| setDebug(sourceFromHiddenImage ? 'source: hidden_image_component' : 'source: uploader_component'); | |
| syncHiddenBox(); | |
| } | |
| setInterval(pullImageAndBox, 150); | |
| window.addEventListener('resize', ()=>setOverlayFromNorm(currentBox)); | |
| setOverlayFromNorm(currentBox); | |
| })(); | |
| </script> | |
| """ | |
| CAMERA_EXAMPLE_SAMPLES = [ | |
| { | |
| "image": "images/example_1.jpg", | |
| "yaw": -45.0, | |
| "pitch": 30.0, | |
| "zoom_delta": 0, | |
| "prompt": "", | |
| }, | |
| { | |
| "image": "images/example_2.png", | |
| "yaw": 0.0, | |
| "pitch": 60.0, | |
| "zoom_delta": -1, | |
| "prompt": "", | |
| }, | |
| { | |
| "image": "images/example_3.jpg", | |
| "yaw": 45.0, | |
| "pitch": 0.0, | |
| "zoom_delta": 1, | |
| "prompt": "", | |
| }, | |
| { | |
| "image": "images/example_4.png", | |
| "yaw": 45.0, | |
| "pitch": 0.0, | |
| "zoom_delta": -1, | |
| "prompt": "", | |
| }, | |
| ] | |
| natural_template_rot_only = """ | |
| Move the camera. | |
| - Camera rotation: Yaw {y_rot:.1f}°, Pitch {x_rot:.1f}°. | |
| """.strip() | |
| ZOOM_IN_PHRASES = ["- Camera zoom: in."] | |
| ZOOM_OUT_PHRASES = ["- Camera zoom: out."] | |
| ZOOM_SAME_PHRASES = ["- Camera zoom: unchanged."] | |
| TEMPLATE_END = "- Keep the 3D scene static; only change the viewpoint." | |
| def draw_red_box_on_image( | |
| image: Image.Image, | |
| cx: float, | |
| cy: float, | |
| w: float, | |
| h: float, | |
| box_thickness: int = 3, | |
| ) -> Image.Image: | |
| """Draw a red box on the image and return the updated image.""" | |
| img = image.copy().convert("RGB") | |
| width, height = img.size | |
| box_cx = cx * width | |
| box_cy = cy * height | |
| box_w = w * width | |
| box_h = h * height | |
| x1 = int(box_cx - box_w / 2) | |
| y1 = int(box_cy - box_h / 2) | |
| x2 = int(box_cx + box_w / 2) | |
| y2 = int(box_cy + box_h / 2) | |
| x1 = max(0, min(x1, width)) | |
| y1 = max(0, min(y1, height)) | |
| x2 = max(0, min(x2, width)) | |
| y2 = max(0, min(y2, height)) | |
| draw = ImageDraw.Draw(img) | |
| draw.rectangle([x1, y1, x2, y2], outline="red", width=box_thickness) | |
| return img | |
| def iphone_to_web_image( | |
| image, | |
| crop_ratio: float = 1.0, | |
| target_short_side: int = 1024, | |
| blur_radius: float = 0.9, | |
| contrast_gamma: float = 1.07, | |
| ): | |
| """ | |
| Preprocess an iPhone photo so it is closer to typical web-image statistics. | |
| """ | |
| if isinstance(image, str): | |
| image = Image.open(image) | |
| image = image.convert("RGB") | |
| image = ImageOps.exif_transpose(image) | |
| w, h = image.size | |
| if 0 < crop_ratio < 1.0: | |
| new_w = int(w * crop_ratio) | |
| new_h = int(h * crop_ratio) | |
| left = max((w - new_w) // 2, 0) | |
| top = max((h - new_h) // 2, 0) | |
| image = image.crop((left, top, left + new_w, top + new_h)) | |
| w, h = image.size | |
| short_side = min(w, h) | |
| if short_side != target_short_side: | |
| scale = target_short_side / short_side | |
| new_w = int(round(w * scale)) | |
| new_h = int(round(h * scale)) | |
| image = image.resize((new_w, new_h), Image.BICUBIC) | |
| blurred = image.filter(ImageFilter.GaussianBlur(radius=blur_radius)) | |
| image = Image.blend(image, blurred, alpha=0.22) | |
| arr = np.asarray(image).astype(np.float32) / 255.0 | |
| arr = np.clip(arr, 0.0, 1.0) | |
| arr = np.power(arr, contrast_gamma) | |
| arr = np.clip(arr * 255.0, 0, 255).astype(np.uint8) | |
| return Image.fromarray(arr) | |
| def pil_to_data_url(image: Image.Image) -> str: | |
| buf = io.BytesIO() | |
| image.convert("RGB").save(buf, format="PNG") | |
| img_str = base64.b64encode(buf.getvalue()).decode() | |
| return f"data:image/png;base64,{img_str}" | |
| def build_camera_preview_html(image: Optional[Image.Image]) -> str: | |
| if image is None: | |
| return """ | |
| <div style=" | |
| width: 100%; | |
| height: 320px; | |
| display: flex; | |
| align-items: center; | |
| justify-content: center; | |
| border: 1px solid #d8e1ee; | |
| border-radius: 12px; | |
| background: linear-gradient(180deg, #ffffff 0%, #f8fafc 100%); | |
| color: #64748b; | |
| font-size: 14px; | |
| "> | |
| Input image preview will appear here. | |
| </div> | |
| """ | |
| data_url = pil_to_data_url(image) | |
| return f""" | |
| <div style=" | |
| width: 100%; | |
| height: 320px; | |
| display: flex; | |
| align-items: center; | |
| justify-content: center; | |
| border: 1px solid #d8e1ee; | |
| border-radius: 12px; | |
| background: linear-gradient(180deg, #ffffff 0%, #f8fafc 100%); | |
| overflow: hidden; | |
| "> | |
| <img | |
| src="{data_url}" | |
| alt="Input preview" | |
| style=" | |
| max-width: 100%; | |
| max-height: 100%; | |
| width: auto; | |
| height: auto; | |
| object-fit: contain; | |
| display: block; | |
| " | |
| /> | |
| </div> | |
| """ | |
| def _build_conversation_instruction(y_rot: float, x_rot: float, d_delta: int, rng: random.Random) -> str: | |
| instr = natural_template_rot_only.format(x_rot=x_rot, y_rot=y_rot) | |
| if d_delta < 0: | |
| instr += "\n" + rng.choice(ZOOM_IN_PHRASES) | |
| elif d_delta > 0: | |
| instr += "\n" + rng.choice(ZOOM_OUT_PHRASES) | |
| else: | |
| instr += "\n" + rng.choice(ZOOM_SAME_PHRASES) | |
| return instr + "\n" + TEMPLATE_END | |
| def clamp(v: float, lo: float, hi: float) -> float: | |
| return max(lo, min(hi, v)) | |
| def snap_to_step(value: float, step: float, lo: float, hi: float) -> float: | |
| value = clamp(float(value), lo, hi) | |
| snapped = round(value / step) * step | |
| return clamp(snapped, lo, hi) | |
| def snap_yaw(value: float) -> float: | |
| return snap_to_step(value, YAW_STEP, YAW_MIN, YAW_MAX) | |
| def snap_pitch(value: float) -> float: | |
| return snap_to_step(value, PITCH_STEP, PITCH_MIN, PITCH_MAX) | |
| def snap_zoom_delta(v: int) -> int: | |
| return int(max(-1, min(1, int(round(v))))) | |
| def zoom_word(zoom_delta: int) -> str: | |
| if zoom_delta < 0: | |
| return "in" | |
| if zoom_delta > 0: | |
| return "out" | |
| return "unchanged" | |
| def build_camera_instruction(yaw: float, pitch: float, zoom_delta: int, seed: int = 0) -> str: | |
| yaw = snap_yaw(yaw) | |
| pitch = snap_pitch(pitch) | |
| zoom_delta = snap_zoom_delta(zoom_delta) | |
| rng = random.Random(int(seed)) | |
| return _build_conversation_instruction(y_rot=yaw, x_rot=pitch, d_delta=zoom_delta, rng=rng) | |
| def build_camera_display(yaw: float, pitch: float, zoom_delta: int) -> str: | |
| yaw = snap_yaw(yaw) | |
| pitch = snap_pitch(pitch) | |
| zoom_delta = snap_zoom_delta(zoom_delta) | |
| return f"Yaw {yaw:.1f}°, Pitch {pitch:.1f}°, Zoom: {zoom_word(zoom_delta)}." | |
| def default_rotation_value() -> Dict[str, Any]: | |
| return {"yaw": 0.0, "view": ROTATION_VIEW_ORDER[0]} | |
| def normalize_rotation_yaw(yaw: float) -> float: | |
| return float(((float(yaw) % 360.0) + 360.0) % 360.0) | |
| def snap_rotation_yaw(yaw: float) -> float: | |
| return float((round(normalize_rotation_yaw(yaw) / 45.0) * 45) % 360) | |
| def rotation_view_from_yaw(yaw: float) -> str: | |
| index = int(round(snap_rotation_yaw(yaw) / 45.0)) % len(ROTATION_VIEW_ORDER) | |
| return ROTATION_VIEW_ORDER[index] | |
| def normalize_rotation_value(rotation_value: Any) -> Dict[str, Any]: | |
| if not isinstance(rotation_value, dict): | |
| return default_rotation_value() | |
| yaw = snap_rotation_yaw(float(rotation_value.get("yaw", 0.0))) | |
| view = str(rotation_value.get("view", "") or "").strip().lower() | |
| normalized_view = rotation_view_from_yaw(yaw) | |
| if view in ROTATION_VIEW_ORDER: | |
| normalized_view = view | |
| yaw = snap_rotation_yaw(ROTATION_VIEW_ORDER.index(view) * 45.0) | |
| return {"yaw": yaw, "view": normalized_view} | |
| def normalize_object_description_for_prompt(object_name: str) -> str: | |
| object_name = str(object_name or "").strip() | |
| if not object_name: | |
| return "" | |
| if re.match(r"(?i)^the\s+", object_name): | |
| object_name = re.sub(r"(?i)^the\s+", "", object_name, count=1).strip() | |
| return object_name | |
| def build_object_moving_prompt(object_name: str) -> str: | |
| object_name = normalize_object_description_for_prompt(object_name) | |
| if not object_name: | |
| return "" | |
| return f"move the {object_name} into the red box, remove the red box, remove the {object_name}" | |
| def build_object_rotation_prompt(object_name: str, rotation_value: Any) -> str: | |
| normalized = normalize_rotation_value(rotation_value) | |
| object_name = normalize_object_description_for_prompt(object_name) | |
| if not object_name: | |
| return "" | |
| return f"rotate the {object_name} to show the {normalized['view']} view." | |
| def build_object_rotation_display(rotation_value: Any) -> str: | |
| normalized = normalize_rotation_value(rotation_value) | |
| return f"Object View: {normalized['view']}." | |
| def should_use_auto_pe(yaw: float) -> bool: | |
| yaw = snap_yaw(yaw) | |
| return yaw < AUTO_PE_DISABLE_YAW_MIN or yaw > AUTO_PE_DISABLE_YAW_MAX | |
| def resolve_local_three_js(explicit_js_path: Optional[str] = None) -> Optional[Path]: | |
| if explicit_js_path: | |
| p = Path(explicit_js_path).expanduser().resolve() | |
| if p.is_file(): | |
| return p | |
| raise FileNotFoundError(f"--three-js-path not found: {explicit_js_path}") | |
| script_dir = Path(__file__).resolve().parent | |
| local_default = script_dir / "three.min.js" | |
| if local_default.is_file(): | |
| return local_default | |
| return None | |
| def read_local_js_inline(js_path: Optional[Path]) -> str: | |
| if js_path is None: | |
| return "" | |
| content = js_path.read_text(encoding="utf-8") | |
| return f"<script>\n{content}\n</script>" | |
| def build_demo_examples_from_config() -> Tuple[List[List[Any]], List[List[Any]]]: | |
| """ | |
| examples_table: | |
| [example_id, image_name, yaw, pitch, zoom] | |
| examples_full: | |
| [image_path, yaw, pitch, zoom, prompt] | |
| """ | |
| example_rows: List[List[Any]] = [] | |
| valid_examples: List[List[Any]] = [] | |
| for idx, item in enumerate(CAMERA_EXAMPLE_SAMPLES): | |
| image_path = str(item.get("image", "")).strip() | |
| if not image_path: | |
| print(f"[Warning] Skip example {idx}: image path is empty.") | |
| continue | |
| image_path_obj = Path(image_path) | |
| if not image_path_obj.is_absolute(): | |
| image_path_obj = (Path(__file__).resolve().parent / image_path_obj).resolve() | |
| if not image_path_obj.is_file(): | |
| print(f"[Warning] Skip example {idx}: image file not found: {image_path_obj}") | |
| continue | |
| image_path = str(image_path_obj) | |
| yaw = snap_yaw(float(item.get("yaw", 0.0))) | |
| pitch = snap_pitch(float(item.get("pitch", 0.0))) | |
| zoom_delta = snap_zoom_delta(int(item.get("zoom_delta", 0))) | |
| prompt = str(item.get("prompt", "") or "").strip() | |
| if not prompt: | |
| prompt = build_camera_instruction(yaw, pitch, zoom_delta, seed=0) | |
| example_id = len(valid_examples) | |
| valid_examples.append([image_path, yaw, pitch, zoom_delta, prompt]) | |
| example_rows.append([example_id, Path(image_path).name, yaw, pitch, zoom_delta]) | |
| return example_rows, valid_examples | |
| def device_select() -> torch.device: | |
| return torch.device("cuda:3") if torch.cuda.is_available() else torch.device("cpu") | |
| def resize_center_crop(img: Image.Image, target_size: Tuple[int, int]) -> Image.Image: | |
| img = img.convert("RGB") | |
| w, h = img.size | |
| bh, bw = target_size | |
| if w == bw and h == bh: | |
| return img.convert("RGB") | |
| scale = max(bh / h, bw / w) | |
| resize_h, resize_w = math.ceil(h * scale), math.ceil(w * scale) | |
| img = img.convert("RGB").resize((resize_w, resize_h), Image.BICUBIC) | |
| img = TF.center_crop(img, target_size) | |
| return img | |
| def image_signature_sha1(image: Image.Image) -> str: | |
| buf = io.BytesIO() | |
| image.convert("RGB").save(buf, format="PNG") | |
| return hashlib.sha1(buf.getvalue()).hexdigest() | |
| def is_offline_mode() -> bool: | |
| return str(os.getenv("HF_HUB_OFFLINE", "0")).strip().lower() in {"1", "true", "yes", "y"} | |
| def save_png_no_compression(image: Image.Image) -> str: | |
| image = image.convert("RGB") | |
| tmp = tempfile.NamedTemporaryFile(prefix="camera_edit_", suffix=".png", delete=False) | |
| tmp_path = tmp.name | |
| tmp.close() | |
| image.save(tmp_path, format="PNG", compress_level=0, optimize=False) | |
| return tmp_path | |
| def make_timestamp_name() -> str: | |
| return datetime.now().strftime("%Y%m%d_%H%M%S_%f") | |
| def ensure_dir(path: str) -> Path: | |
| p = Path(path).expanduser().resolve() | |
| p.mkdir(parents=True, exist_ok=True) | |
| return p | |
| def save_generation_bundle( | |
| save_root: str, | |
| input_image: Image.Image, | |
| output_image: Image.Image, | |
| final_prompt: str, | |
| *, | |
| base_prompt: str = "", | |
| expected_caption: str = "", | |
| display_text: str = "", | |
| yaw: float = 0.0, | |
| pitch: float = 0.0, | |
| zoom_delta: int = 0, | |
| seed: int = 0, | |
| guidance_scale: float = 0.0, | |
| num_inference_steps: int = 0, | |
| enable_prompt_rewrite: bool = False, | |
| rewrite_backend: str = "", | |
| ) -> Optional[Path]: | |
| """ | |
| Save the inputs, outputs, and prompts for a single generation request. | |
| Directory layout: | |
| save_root/ | |
| 20260320_194530_123456/ | |
| input.png | |
| output.png | |
| prompt.txt | |
| meta.json | |
| """ | |
| save_root = str(save_root or "").strip() | |
| if not save_root: | |
| return None | |
| root = ensure_dir(save_root) | |
| timestamp = make_timestamp_name() | |
| run_dir = root / timestamp | |
| suffix = 1 | |
| while run_dir.exists(): | |
| run_dir = root / f"{timestamp}_{suffix:02d}" | |
| suffix += 1 | |
| run_dir.mkdir(parents=True, exist_ok=False) | |
| input_path = run_dir / "input.png" | |
| output_path = run_dir / "output.png" | |
| prompt_path = run_dir / "prompt.txt" | |
| meta_path = run_dir / "meta.json" | |
| input_image.convert("RGB").save( | |
| input_path, | |
| format="PNG", | |
| compress_level=0, | |
| optimize=False, | |
| ) | |
| output_image.convert("RGB").save( | |
| output_path, | |
| format="PNG", | |
| compress_level=0, | |
| optimize=False, | |
| ) | |
| prompt_path.write_text( | |
| str(final_prompt or "").strip() + "\n", | |
| encoding="utf-8", | |
| ) | |
| meta = { | |
| "timestamp": run_dir.name, | |
| "input_image": str(input_path), | |
| "output_image": str(output_path), | |
| "prompt_file": str(prompt_path), | |
| "camera_display": str(display_text or ""), | |
| "yaw": float(yaw), | |
| "pitch": float(pitch), | |
| "zoom_delta": int(zoom_delta), | |
| "seed": int(seed), | |
| "guidance_scale": float(guidance_scale), | |
| "num_inference_steps": int(num_inference_steps), | |
| "enable_prompt_rewrite": bool(enable_prompt_rewrite), | |
| "rewrite_backend": str(rewrite_backend or ""), | |
| "base_prompt": str(base_prompt or "").strip(), | |
| "expected_caption": str(expected_caption or "").strip(), | |
| "final_prompt": str(final_prompt or "").strip(), | |
| } | |
| meta_path.write_text( | |
| json.dumps(meta, ensure_ascii=False, indent=2), | |
| encoding="utf-8", | |
| ) | |
| return run_dir | |
| class EditorApp: | |
| def __init__( | |
| self, | |
| ckpt_root: str, | |
| config_path: Optional[str] = None, | |
| *, | |
| rewrite_model: str = "gpt-5", | |
| hsdp_shard_dim: Optional[int] = None, | |
| enable_prompt_rewrite: bool = False, | |
| basesize: int = BASE_SIZE, | |
| device: Optional[torch.device] = None, | |
| model_load_mode: str = "startup_preload", | |
| ): | |
| from infer_runtime.settings import load_settings | |
| self.ckpt_root = str(ckpt_root) | |
| self.config_path = str(config_path) if config_path else None | |
| self.rewrite_model = str(rewrite_model) | |
| self.hsdp_shard_dim = hsdp_shard_dim | |
| self.enable_prompt_rewrite = bool(enable_prompt_rewrite) | |
| self.basesize = int(basesize) | |
| self.device = device | |
| self.model_load_mode = str(model_load_mode or "startup_preload").strip().lower() | |
| self._preload_attempted = False | |
| self._preload_error = None | |
| self._cpu_preloaded = False | |
| self._move_back_to_cpu_after_infer = os.getenv("MODEL_RETURN_TO_CPU_AFTER_INFER", "0").lower() in {"1", "true", "yes", "y"} | |
| self._infer_lock = threading.Lock() | |
| self._model_lock = threading.Lock() | |
| self._scheduler_template = None | |
| self.model = None | |
| self.settings = load_settings( | |
| ckpt_root=self.ckpt_root, | |
| config_path=self.config_path, | |
| rewrite_model=self.rewrite_model, | |
| default_seed=0, | |
| ) | |
| if is_rank0(): | |
| if self.model_load_mode in {"cpu_preload", "cpu", "cpu_only", "cpu_global_preload"}: | |
| print("[Model] CPU preload mode is enabled. The app will build the model on CPU during startup and move it to GPU only during inference.") | |
| elif self.model_load_mode in {"startup", "startup_preload", "boot", "auto", "eager", "preload", "gpu", "gpu_preload", "cuda", "global_cuda"}: | |
| print("[Model] Direct GPU startup preload mode is enabled. The app will attempt to build the model globally on CUDA during startup.") | |
| else: | |
| print("[Model] Lazy GPU loading is enabled. The model will be built on the first inference call.") | |
| print(f"[Model] Config path: {self.settings.config_path}") | |
| print(f"[Model] Checkpoint path: {self.settings.ckpt_path}") | |
| print(f"[Model] Prompt rewrite enabled: {self.enable_prompt_rewrite}") | |
| print(f"[Model] Prompt rewrite model: {self.rewrite_model}") | |
| print(f"[Model] Basesize: {self.basesize}") | |
| if self.hsdp_shard_dim is not None: | |
| print(f"[Model] Override hsdp_shard_dim: {self.hsdp_shard_dim}") | |
| print(f"[Model] Load mode: {self.model_load_mode}") | |
| def _cpu_preload_enabled(self) -> bool: | |
| return self.model_load_mode in {"cpu_preload", "cpu", "cpu_only", "cpu_global_preload"} | |
| def _ensure_model_loaded(self, target_device: Optional[torch.device] = None) -> None: | |
| desired_device = torch.device(target_device) if target_device is not None else None | |
| if self.model is not None: | |
| if desired_device is None: | |
| return | |
| current_device = getattr(self.model, "current_device", lambda: self.device)() | |
| if current_device == desired_device: | |
| self.device = desired_device | |
| return | |
| with self._model_lock: | |
| if self.model is None: | |
| build_device = desired_device | |
| if build_device is None: | |
| if self._cpu_preload_enabled(): | |
| build_device = torch.device("cpu") | |
| else: | |
| if self.device is None or self.device.type == "cpu": | |
| self.device = resolve_device() | |
| build_device = self.device | |
| if is_rank0(): | |
| print(f"[Model] Building model on device: {build_device}") | |
| from infer_runtime.model import build_model | |
| self.model = build_model( | |
| self.settings, | |
| device=build_device, | |
| hsdp_shard_dim_override=self.hsdp_shard_dim, | |
| ) | |
| self.device = torch.device(build_device) | |
| self._cpu_preloaded = self.device.type == "cpu" | |
| try: | |
| pipeline = getattr(self.model, "pipeline", None) | |
| scheduler = getattr(pipeline, "scheduler", None) | |
| if scheduler is not None: | |
| self._scheduler_template = copy.deepcopy(scheduler) | |
| except Exception as e: | |
| if is_rank0(): | |
| print(f"[Model] Warning: failed to snapshot scheduler template: {e}") | |
| self._scheduler_template = None | |
| if is_rank0(): | |
| print(f"[Model] Device: {self.device}") | |
| print(f"[Model] Model load complete.") | |
| if desired_device is not None: | |
| current_device = getattr(self.model, "current_device", lambda: self.device)() | |
| if current_device != desired_device: | |
| if is_rank0(): | |
| print(f"[Model] Moving model from {current_device} to {desired_device}...") | |
| move_to = getattr(self.model, "move_to_device", None) | |
| if move_to is None: | |
| raise RuntimeError("Loaded model does not support runtime device migration.") | |
| move_to(desired_device) | |
| self.device = desired_device | |
| if is_rank0(): | |
| print(f"[Model] Model now on device: {self.device}") | |
| def _move_model_back_to_cpu(self) -> None: | |
| if self.model is None: | |
| return | |
| current_device = getattr(self.model, "current_device", lambda: self.device)() | |
| if current_device.type == "cpu": | |
| self.device = current_device | |
| return | |
| with self._model_lock: | |
| current_device = getattr(self.model, "current_device", lambda: self.device)() | |
| if current_device.type == "cpu": | |
| self.device = current_device | |
| return | |
| if is_rank0(): | |
| print("[Model] Moving model back to CPU after inference...") | |
| move_to_cpu = getattr(self.model, "move_to_cpu", None) | |
| if move_to_cpu is None: | |
| raise RuntimeError("Loaded model does not support returning to CPU.") | |
| move_to_cpu() | |
| self.device = torch.device("cpu") | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| if is_rank0(): | |
| print("[Model] Model is back on CPU.") | |
| def is_model_loaded(self) -> bool: | |
| return self.model is not None | |
| def is_model_on_cpu(self) -> bool: | |
| if self.model is None: | |
| return False | |
| current_device = getattr(self.model, "current_device", lambda: self.device)() | |
| return current_device.type == "cpu" | |
| def is_model_on_gpu(self) -> bool: | |
| if self.model is None: | |
| return False | |
| current_device = getattr(self.model, "current_device", lambda: self.device)() | |
| return current_device.type == "cuda" | |
| def maybe_preload_model(self) -> None: | |
| if self._preload_attempted: | |
| return | |
| self._preload_attempted = True | |
| if self.model_load_mode in {"lazy", "disabled", "off", "false", "0", "page_warmup", "warmup", "load_on_page", "onload"}: | |
| if is_rank0(): | |
| print("[Model] Startup preload disabled; using runtime loading or page warmup.") | |
| return | |
| if self._cpu_preload_enabled(): | |
| preload_device = torch.device("cpu") | |
| else: | |
| if torch.cuda.is_available(): | |
| preload_device = torch.device(f"cuda:{int(os.environ.get('LOCAL_RANK', '0'))}") | |
| else: | |
| preload_device = torch.device("cpu") | |
| if is_rank0(): | |
| if preload_device.type == "cpu" and self._cpu_preload_enabled(): | |
| print("[Model] Attempting startup CPU preload...") | |
| else: | |
| print(f"[Model] Attempting direct startup preload on {preload_device}...") | |
| try: | |
| self._ensure_model_loaded(target_device=preload_device) | |
| self._preload_error = None | |
| if is_rank0(): | |
| if preload_device.type == "cpu" and self._cpu_preload_enabled(): | |
| print("[Model] Startup CPU preload succeeded.") | |
| else: | |
| print(f"[Model] Startup GPU preload succeeded on {preload_device}.") | |
| except Exception as e: | |
| self._preload_error = str(e) | |
| if is_rank0(): | |
| if preload_device.type == "cpu" and self._cpu_preload_enabled(): | |
| print(f"[Model] Startup CPU preload failed; falling back to lazy loading: {e}") | |
| else: | |
| print(f"[Model] Startup GPU preload failed; falling back to lazy loading: {e}") | |
| self.model = None | |
| self._scheduler_template = None | |
| self.device = None | |
| try: | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| except Exception: | |
| pass | |
| def _reset_scheduler_runtime_state(self) -> None: | |
| pipeline = getattr(self.model, "pipeline", None) | |
| scheduler = getattr(pipeline, "scheduler", None) | |
| if scheduler is None: | |
| return | |
| for attr_name, attr_value in ( | |
| ("_step_index", None), | |
| ("step_index", None), | |
| ("_begin_index", None), | |
| ("begin_index", None), | |
| ("_timesteps", None), | |
| ): | |
| if hasattr(scheduler, attr_name): | |
| try: | |
| setattr(scheduler, attr_name, attr_value) | |
| except Exception: | |
| pass | |
| if hasattr(scheduler, "set_begin_index"): | |
| try: | |
| scheduler.set_begin_index(0) | |
| except Exception: | |
| try: | |
| scheduler.set_begin_index(None) | |
| except Exception: | |
| pass | |
| def _restore_fresh_scheduler(self) -> None: | |
| pipeline = getattr(self.model, "pipeline", None) | |
| if pipeline is None: | |
| return | |
| if self._scheduler_template is not None: | |
| try: | |
| pipeline.scheduler = copy.deepcopy(self._scheduler_template) | |
| return | |
| except Exception as e: | |
| if is_rank0(): | |
| print(f"[Model] Warning: failed to restore scheduler from template: {e}") | |
| self._reset_scheduler_runtime_state() | |
| def _choose_size(self, img: Image.Image, req_h: int, req_w: int) -> Tuple[int, int]: | |
| """ | |
| Best-effort preview size helper for UI rendering only. | |
| The actual inference-side bucket selection is handled internally by the | |
| new infer_runtime model through InferenceParams(basesize=...). | |
| """ | |
| base = max(int(self.basesize), int(req_h), int(req_w)) | |
| w, h = img.size | |
| if w <= 0 or h <= 0: | |
| return base, base | |
| scale = base / float(min(h, w)) | |
| out_h = max(64, int(round((h * scale) / 64.0) * 64)) | |
| out_w = max(64, int(round((w * scale) / 64.0) * 64)) | |
| return out_h, out_w | |
| def run( | |
| self, | |
| image: Image.Image, | |
| prompt: str, | |
| steps: int, | |
| guidance: float, | |
| seed: int, | |
| ) -> Tuple[Image.Image, str]: | |
| from infer_runtime.model import InferenceParams | |
| if image is None: | |
| raise ValueError("image is None") | |
| input_image = image.convert("RGB") | |
| target_device = resolve_device() | |
| self._ensure_model_loaded(target_device=target_device) | |
| effective_prompt = self.model.maybe_rewrite_prompt( | |
| prompt, | |
| input_image, | |
| self.enable_prompt_rewrite, | |
| ) | |
| with self._infer_lock: | |
| self._restore_fresh_scheduler() | |
| try: | |
| output_image = self.model.infer( | |
| InferenceParams( | |
| prompt=effective_prompt, | |
| image=input_image, | |
| height=self.basesize, | |
| width=self.basesize, | |
| steps=int(steps), | |
| guidance_scale=float(guidance), | |
| seed=int(seed), | |
| neg_prompt="", | |
| basesize=int(self.basesize), | |
| ) | |
| ) | |
| finally: | |
| self._reset_scheduler_runtime_state() | |
| if self._move_back_to_cpu_after_infer: | |
| self._move_model_back_to_cpu() | |
| return output_image, effective_prompt | |
| class CameraControl3D(gr.HTML): | |
| """ | |
| value: | |
| { "yaw": float[-90,90], "pitch": float[-90,90], "zoom_delta": int{-1,0,1} } | |
| """ | |
| def __init__(self, value=None, imageUrl=None, three_available: bool = False, **kwargs): | |
| if value is None: | |
| value = {"yaw": 0.0, "pitch": 0.0, "zoom_delta": 0} | |
| offline_msg = """ | |
| <div style="width:100%;height:450px;display:flex;align-items:center;justify-content:center;background:linear-gradient(180deg,#ffffff 0%,#f3f6fb 100%);border:1px solid #d8e1ee;border-radius:12px;color:#334155;padding:24px;box-sizing:border-box;text-align:center;"> | |
| <div> | |
| <div style="font-size:20px;font-weight:600;margin-bottom:10px;">3D viewport is disabled</div> | |
| <div style="font-size:14px;line-height:1.6;opacity:0.9;"> | |
| No local three.js was provided.<br/> | |
| You can still use the sliders below for camera control.<br/> | |
| Put <code>three.min.js</code> next to this <code>.py</code> file to enable 3D dragging. | |
| </div> | |
| </div> | |
| </div> | |
| """ | |
| html_template = """ | |
| <div id="camera-control-wrapper" style="width: 100%; height: 450px; position: relative; background: radial-gradient(circle at top, #ffffff 0%, #f4f7fb 58%, #e9eef6 100%); border: 1px solid #d8e1ee; border-radius: 12px; overflow: hidden;"> | |
| <div id="prompt-overlay" style=" | |
| position: absolute; | |
| bottom: 10px; | |
| left: 50%; | |
| transform: translateX(-50%); | |
| background: rgba(255,255,255,0.92); | |
| padding: 10px 14px; | |
| border-radius: 8px; | |
| border: 1px solid rgba(148,163,184,0.35); | |
| box-shadow: 0 12px 28px rgba(15,23,42,0.12); | |
| font-family: monospace; | |
| font-size: 12px; | |
| line-height: 1.2; | |
| color: #0f766e; | |
| white-space: nowrap; | |
| z-index: 10; | |
| max-width: calc(100% - 24px); | |
| box-sizing: border-box; | |
| text-align: center;"></div> | |
| </div> | |
| """ if three_available else offline_msg | |
| js_on_load = f""" | |
| (() => {{ | |
| function waitForThreeAndInit() {{ | |
| if (typeof THREE === 'undefined') {{ | |
| setTimeout(waitForThreeAndInit, 50); | |
| return; | |
| }} | |
| init3D(); | |
| }} | |
| function init3D() {{ | |
| const YAW_MIN = {YAW_MIN}; | |
| const YAW_MAX = {YAW_MAX}; | |
| const PITCH_MIN = {PITCH_MIN}; | |
| const PITCH_MAX = {PITCH_MAX}; | |
| const YAW_STEP = {YAW_STEP}; | |
| const PITCH_STEP = {PITCH_STEP}; | |
| const wrapper = element.querySelector('#camera-control-wrapper'); | |
| const promptOverlay = element.querySelector('#prompt-overlay'); | |
| if (!wrapper || !promptOverlay) return; | |
| function clamp(v, lo, hi) {{ | |
| return Math.min(hi, Math.max(lo, v)); | |
| }} | |
| function snapToStep(value, step, lo, hi) {{ | |
| const v = clamp(value, lo, hi); | |
| const snapped = Math.round(v / step) * step; | |
| return clamp(snapped, lo, hi); | |
| }} | |
| function snapYaw(value) {{ | |
| return snapToStep(value, YAW_STEP, YAW_MIN, YAW_MAX); | |
| }} | |
| function snapPitch(value) {{ | |
| return snapToStep(value, PITCH_STEP, PITCH_MIN, PITCH_MAX); | |
| }} | |
| function snapZoom(v) {{ | |
| return Math.round(clamp(v, -1, 1)); | |
| }} | |
| function zoomWord(z) {{ | |
| if (z < 0) return 'in'; | |
| if (z > 0) return 'out'; | |
| return 'unchanged'; | |
| }} | |
| function buildCompactDisplay(yaw, pitch, zoomDelta) {{ | |
| const y = snapYaw(yaw); | |
| const x = snapPitch(pitch); | |
| const z = snapZoom(zoomDelta); | |
| return `- Camera rotation: Yaw ${{y.toFixed(1)}}°, Pitch ${{x.toFixed(1)}}°, zoom: ${{zoomWord(z)}}.`; | |
| }} | |
| const scene = new THREE.Scene(); | |
| scene.background = new THREE.Color(0xf8fafc); | |
| const camera = new THREE.PerspectiveCamera(50, wrapper.clientWidth / wrapper.clientHeight, 0.1, 1000); | |
| camera.position.set(4.5, 3, 4.5); | |
| camera.lookAt(0, 0.75, 0); | |
| const renderer = new THREE.WebGLRenderer({{ antialias: true }}); | |
| renderer.setSize(wrapper.clientWidth, wrapper.clientHeight); | |
| renderer.setPixelRatio(Math.min(window.devicePixelRatio, 2)); | |
| wrapper.insertBefore(renderer.domElement, promptOverlay); | |
| scene.add(new THREE.AmbientLight(0xffffff, 0.88)); | |
| const dirLight = new THREE.DirectionalLight(0xffffff, 0.6); | |
| dirLight.position.set(5, 10, 5); | |
| scene.add(dirLight); | |
| const fillLight = new THREE.DirectionalLight(0xdbeafe, 0.45); | |
| fillLight.position.set(-4, 5, -2); | |
| scene.add(fillLight); | |
| scene.add(new THREE.GridHelper(8, 16, 0xcbd5e1, 0xe2e8f0)); | |
| const CENTER = new THREE.Vector3(0, 0.75, 0); | |
| const BASE_DISTANCE = 1.6; | |
| const AZIMUTH_RADIUS = 2.4; | |
| const ELEVATION_RADIUS = 1.8; | |
| let yawAngle = props.value?.yaw ?? 0.0; | |
| let pitchAngle = props.value?.pitch ?? 0.0; | |
| let zoomValue = props.value?.zoom_delta ?? 0.0; | |
| function createPlaceholderTexture() {{ | |
| const canvas = document.createElement('canvas'); | |
| canvas.width = 256; | |
| canvas.height = 256; | |
| const ctx = canvas.getContext('2d'); | |
| ctx.fillStyle = '#e2e8f0'; | |
| ctx.fillRect(0, 0, 256, 256); | |
| ctx.fillStyle = '#f8b4a6'; | |
| ctx.beginPath(); | |
| ctx.arc(128, 128, 80, 0, Math.PI * 2); | |
| ctx.fill(); | |
| ctx.fillStyle = '#334155'; | |
| ctx.beginPath(); | |
| ctx.arc(100, 110, 10, 0, Math.PI * 2); | |
| ctx.arc(156, 110, 10, 0, Math.PI * 2); | |
| ctx.fill(); | |
| ctx.strokeStyle = '#334155'; | |
| ctx.lineWidth = 3; | |
| ctx.beginPath(); | |
| ctx.arc(128, 130, 35, 0.2, Math.PI - 0.2); | |
| ctx.stroke(); | |
| return new THREE.CanvasTexture(canvas); | |
| }} | |
| const planeMaterial = new THREE.MeshBasicMaterial({{ | |
| map: createPlaceholderTexture(), | |
| side: THREE.DoubleSide | |
| }}); | |
| let targetPlane = new THREE.Mesh(new THREE.PlaneGeometry(1.2, 1.2), planeMaterial); | |
| targetPlane.position.copy(CENTER); | |
| scene.add(targetPlane); | |
| function updateTextureFromUrl(url) {{ | |
| if (!url) {{ | |
| planeMaterial.map = createPlaceholderTexture(); | |
| planeMaterial.needsUpdate = true; | |
| scene.remove(targetPlane); | |
| targetPlane = new THREE.Mesh(new THREE.PlaneGeometry(1.2, 1.2), planeMaterial); | |
| targetPlane.position.copy(CENTER); | |
| scene.add(targetPlane); | |
| return; | |
| }} | |
| const loader = new THREE.TextureLoader(); | |
| loader.load(url, (texture) => {{ | |
| texture.minFilter = THREE.LinearFilter; | |
| texture.magFilter = THREE.LinearFilter; | |
| planeMaterial.map = texture; | |
| planeMaterial.needsUpdate = true; | |
| const img = texture.image; | |
| if (img && img.width && img.height) {{ | |
| const aspect = img.width / img.height; | |
| const maxSize = 1.5; | |
| let planeWidth, planeHeight; | |
| if (aspect > 1) {{ | |
| planeWidth = maxSize; | |
| planeHeight = maxSize / aspect; | |
| }} else {{ | |
| planeHeight = maxSize; | |
| planeWidth = maxSize * aspect; | |
| }} | |
| scene.remove(targetPlane); | |
| targetPlane = new THREE.Mesh( | |
| new THREE.PlaneGeometry(planeWidth, planeHeight), | |
| planeMaterial | |
| ); | |
| targetPlane.position.copy(CENTER); | |
| scene.add(targetPlane); | |
| }} | |
| }}, undefined, (err) => {{ | |
| console.error('Failed to load texture:', err); | |
| }}); | |
| }} | |
| if (props.imageUrl) updateTextureFromUrl(props.imageUrl); | |
| const cameraGroup = new THREE.Group(); | |
| const bodyMat = new THREE.MeshStandardMaterial({{ color: 0x6699cc, metalness: 0.5, roughness: 0.3 }}); | |
| cameraGroup.add(new THREE.Mesh(new THREE.BoxGeometry(0.3, 0.22, 0.38), bodyMat)); | |
| const lens = new THREE.Mesh( | |
| new THREE.CylinderGeometry(0.09, 0.11, 0.18, 16), | |
| new THREE.MeshStandardMaterial({{ color: 0x6699cc, metalness: 0.5, roughness: 0.3 }}) | |
| ); | |
| lens.rotation.x = Math.PI / 2; | |
| lens.position.z = 0.26; | |
| cameraGroup.add(lens); | |
| scene.add(cameraGroup); | |
| const azimuthPoints = []; | |
| for (let i = 0; i <= 64; i++) {{ | |
| const angle = THREE.MathUtils.degToRad(YAW_MIN + ((YAW_MAX - YAW_MIN) * i / 64)); | |
| azimuthPoints.push(new THREE.Vector3( | |
| AZIMUTH_RADIUS * Math.sin(angle), | |
| 0.05, | |
| AZIMUTH_RADIUS * Math.cos(angle) | |
| )); | |
| }} | |
| const azimuthCurve = new THREE.CatmullRomCurve3(azimuthPoints); | |
| const azimuthArc = new THREE.Mesh( | |
| new THREE.TubeGeometry(azimuthCurve, 64, 0.04, 8, false), | |
| new THREE.MeshStandardMaterial({{ color: 0x00ff88, emissive: 0x00ff88, emissiveIntensity: 0.3 }}) | |
| ); | |
| scene.add(azimuthArc); | |
| const azimuthHandle = new THREE.Mesh( | |
| new THREE.SphereGeometry(0.18, 16, 16), | |
| new THREE.MeshStandardMaterial({{ color: 0x00ff88, emissive: 0x00ff88, emissiveIntensity: 0.5 }}) | |
| ); | |
| azimuthHandle.userData.type = 'yaw'; | |
| scene.add(azimuthHandle); | |
| const arcPoints = []; | |
| for (let i = 0; i <= 64; i++) {{ | |
| const angle = THREE.MathUtils.degToRad(PITCH_MIN + ((PITCH_MAX - PITCH_MIN) * i / 64)); | |
| arcPoints.push(new THREE.Vector3( | |
| -0.8, | |
| ELEVATION_RADIUS * Math.sin(angle) + CENTER.y, | |
| ELEVATION_RADIUS * Math.cos(angle) | |
| )); | |
| }} | |
| const arcCurve = new THREE.CatmullRomCurve3(arcPoints); | |
| const elevationArc = new THREE.Mesh( | |
| new THREE.TubeGeometry(arcCurve, 64, 0.04, 8, false), | |
| new THREE.MeshStandardMaterial({{ color: 0xff69b4, emissive: 0xff69b4, emissiveIntensity: 0.3 }}) | |
| ); | |
| scene.add(elevationArc); | |
| const elevationHandle = new THREE.Mesh( | |
| new THREE.SphereGeometry(0.18, 16, 16), | |
| new THREE.MeshStandardMaterial({{ color: 0xff69b4, emissive: 0xff69b4, emissiveIntensity: 0.5 }}) | |
| ); | |
| elevationHandle.userData.type = 'pitch'; | |
| scene.add(elevationHandle); | |
| const distanceLineGeo = new THREE.BufferGeometry(); | |
| const distanceLine = new THREE.Line( | |
| distanceLineGeo, | |
| new THREE.LineBasicMaterial({{ color: 0xffa500 }}) | |
| ); | |
| scene.add(distanceLine); | |
| const distanceHandle = new THREE.Mesh( | |
| new THREE.SphereGeometry(0.18, 16, 16), | |
| new THREE.MeshStandardMaterial({{ color: 0xffa500, emissive: 0xffa500, emissiveIntensity: 0.5 }}) | |
| ); | |
| distanceHandle.userData.type = 'zoom'; | |
| scene.add(distanceHandle); | |
| function zoomValueToDistanceFactor(z) {{ | |
| return 1.0 + 0.4 * clamp(z, -1, 1); | |
| }} | |
| function updatePositions() {{ | |
| yawAngle = clamp(yawAngle, YAW_MIN, YAW_MAX); | |
| pitchAngle = clamp(pitchAngle, PITCH_MIN, PITCH_MAX); | |
| zoomValue = clamp(zoomValue, -1, 1); | |
| const distanceFactor = zoomValueToDistanceFactor(zoomValue); | |
| const distance = BASE_DISTANCE * distanceFactor; | |
| const yawRad = THREE.MathUtils.degToRad(yawAngle); | |
| const pitchRad = THREE.MathUtils.degToRad(pitchAngle); | |
| const camX = distance * Math.sin(yawRad) * Math.cos(pitchRad); | |
| const camY = distance * Math.sin(pitchRad) + CENTER.y; | |
| const camZ = distance * Math.cos(yawRad) * Math.cos(pitchRad); | |
| cameraGroup.position.set(camX, camY, camZ); | |
| cameraGroup.lookAt(CENTER); | |
| azimuthHandle.position.set( | |
| AZIMUTH_RADIUS * Math.sin(yawRad), | |
| 0.05, | |
| AZIMUTH_RADIUS * Math.cos(yawRad) | |
| ); | |
| elevationHandle.position.set( | |
| -0.8, | |
| ELEVATION_RADIUS * Math.sin(pitchRad) + CENTER.y, | |
| ELEVATION_RADIUS * Math.cos(pitchRad) | |
| ); | |
| const orangeDist = distance - 0.5; | |
| distanceHandle.position.set( | |
| orangeDist * Math.sin(yawRad) * Math.cos(pitchRad), | |
| orangeDist * Math.sin(pitchRad) + CENTER.y, | |
| orangeDist * Math.cos(yawRad) * Math.cos(pitchRad) | |
| ); | |
| distanceLineGeo.setFromPoints([cameraGroup.position.clone(), CENTER.clone()]); | |
| promptOverlay.textContent = buildCompactDisplay(yawAngle, pitchAngle, zoomValue); | |
| }} | |
| function updatePropsAndTrigger() {{ | |
| props.value = {{ | |
| yaw: snapYaw(yawAngle), | |
| pitch: snapPitch(pitchAngle), | |
| zoom_delta: snapZoom(zoomValue), | |
| }}; | |
| trigger('change', props.value); | |
| }} | |
| const raycaster = new THREE.Raycaster(); | |
| const mouse = new THREE.Vector2(); | |
| let isDragging = false; | |
| let dragTarget = null; | |
| let dragStartMouse = new THREE.Vector2(); | |
| let dragStartZoom = 0.0; | |
| const intersection = new THREE.Vector3(); | |
| const canvas = renderer.domElement; | |
| canvas.addEventListener('mousedown', (e) => {{ | |
| const rect = canvas.getBoundingClientRect(); | |
| mouse.x = ((e.clientX - rect.left) / rect.width) * 2 - 1; | |
| mouse.y = -((e.clientY - rect.top) / rect.height) * 2 + 1; | |
| raycaster.setFromCamera(mouse, camera); | |
| const intersects = raycaster.intersectObjects([azimuthHandle, elevationHandle, distanceHandle]); | |
| if (intersects.length > 0) {{ | |
| isDragging = true; | |
| dragTarget = intersects[0].object; | |
| dragTarget.material.emissiveIntensity = 1.0; | |
| dragTarget.scale.setScalar(1.3); | |
| dragStartMouse.copy(mouse); | |
| dragStartZoom = zoomValue; | |
| canvas.style.cursor = 'grabbing'; | |
| }} | |
| }}); | |
| canvas.addEventListener('mousemove', (e) => {{ | |
| const rect = canvas.getBoundingClientRect(); | |
| mouse.x = ((e.clientX - rect.left) / rect.width) * 2 - 1; | |
| mouse.y = -((e.clientY - rect.top) / rect.height) * 2 + 1; | |
| if (isDragging && dragTarget) {{ | |
| raycaster.setFromCamera(mouse, camera); | |
| if (dragTarget.userData.type === 'yaw') {{ | |
| const plane = new THREE.Plane(new THREE.Vector3(0, 1, 0), -0.05); | |
| if (raycaster.ray.intersectPlane(plane, intersection)) {{ | |
| yawAngle = THREE.MathUtils.radToDeg(Math.atan2(intersection.x, intersection.z)); | |
| yawAngle = clamp(yawAngle, YAW_MIN, YAW_MAX); | |
| }} | |
| }} else if (dragTarget.userData.type === 'pitch') {{ | |
| const plane = new THREE.Plane(new THREE.Vector3(1, 0, 0), -0.8); | |
| if (raycaster.ray.intersectPlane(plane, intersection)) {{ | |
| const relY = intersection.y - CENTER.y; | |
| const relZ = intersection.z; | |
| pitchAngle = THREE.MathUtils.radToDeg(Math.atan2(relY, relZ)); | |
| pitchAngle = clamp(pitchAngle, PITCH_MIN, PITCH_MAX); | |
| }} | |
| }} else if (dragTarget.userData.type === 'zoom') {{ | |
| const deltaY = mouse.y - dragStartMouse.y; | |
| zoomValue = clamp(dragStartZoom - deltaY * 1.5, -1, 1); | |
| }} | |
| updatePositions(); | |
| }} else {{ | |
| raycaster.setFromCamera(mouse, camera); | |
| const intersects = raycaster.intersectObjects([azimuthHandle, elevationHandle, distanceHandle]); | |
| [azimuthHandle, elevationHandle, distanceHandle].forEach(h => {{ | |
| h.material.emissiveIntensity = 0.5; | |
| h.scale.setScalar(1); | |
| }}); | |
| if (intersects.length > 0) {{ | |
| intersects[0].object.material.emissiveIntensity = 0.8; | |
| intersects[0].object.scale.setScalar(1.1); | |
| canvas.style.cursor = 'grab'; | |
| }} else {{ | |
| canvas.style.cursor = 'default'; | |
| }} | |
| }} | |
| }}); | |
| const onMouseUp = () => {{ | |
| if (dragTarget) {{ | |
| dragTarget.material.emissiveIntensity = 0.5; | |
| dragTarget.scale.setScalar(1); | |
| const targetYaw = snapYaw(yawAngle); | |
| const targetPitch = snapPitch(pitchAngle); | |
| const targetZoom = snapZoom(zoomValue); | |
| const startYaw = yawAngle; | |
| const startPitch = pitchAngle; | |
| const startZoom = zoomValue; | |
| const startTime = Date.now(); | |
| function animateSnap() {{ | |
| const t = Math.min((Date.now() - startTime) / 200, 1); | |
| const ease = 1 - Math.pow(1 - t, 3); | |
| yawAngle = clamp(startYaw + (targetYaw - startYaw) * ease, YAW_MIN, YAW_MAX); | |
| pitchAngle = clamp(startPitch + (targetPitch - startPitch) * ease, PITCH_MIN, PITCH_MAX); | |
| zoomValue = clamp(startZoom + (targetZoom - startZoom) * ease, -1, 1); | |
| updatePositions(); | |
| if (t < 1) requestAnimationFrame(animateSnap); | |
| else {{ | |
| yawAngle = targetYaw; | |
| pitchAngle = targetPitch; | |
| zoomValue = targetZoom; | |
| updatePositions(); | |
| updatePropsAndTrigger(); | |
| }} | |
| }} | |
| animateSnap(); | |
| }} | |
| isDragging = false; | |
| dragTarget = null; | |
| canvas.style.cursor = 'default'; | |
| }}; | |
| canvas.addEventListener('mouseup', onMouseUp); | |
| canvas.addEventListener('mouseleave', onMouseUp); | |
| canvas.addEventListener('touchstart', (e) => {{ | |
| e.preventDefault(); | |
| const touch = e.touches[0]; | |
| const rect = canvas.getBoundingClientRect(); | |
| mouse.x = ((touch.clientX - rect.left) / rect.width) * 2 - 1; | |
| mouse.y = -((touch.clientY - rect.top) / rect.height) * 2 + 1; | |
| raycaster.setFromCamera(mouse, camera); | |
| const intersects = raycaster.intersectObjects([azimuthHandle, elevationHandle, distanceHandle]); | |
| if (intersects.length > 0) {{ | |
| isDragging = true; | |
| dragTarget = intersects[0].object; | |
| dragTarget.material.emissiveIntensity = 1.0; | |
| dragTarget.scale.setScalar(1.3); | |
| dragStartMouse.copy(mouse); | |
| dragStartZoom = zoomValue; | |
| }} | |
| }}, {{ passive: false }}); | |
| canvas.addEventListener('touchmove', (e) => {{ | |
| e.preventDefault(); | |
| const touch = e.touches[0]; | |
| const rect = canvas.getBoundingClientRect(); | |
| mouse.x = ((touch.clientX - rect.left) / rect.width) * 2 - 1; | |
| mouse.y = -((touch.clientY - rect.top) / rect.height) * 2 + 1; | |
| if (isDragging && dragTarget) {{ | |
| raycaster.setFromCamera(mouse, camera); | |
| if (dragTarget.userData.type === 'yaw') {{ | |
| const plane = new THREE.Plane(new THREE.Vector3(0, 1, 0), -0.05); | |
| if (raycaster.ray.intersectPlane(plane, intersection)) {{ | |
| yawAngle = THREE.MathUtils.radToDeg(Math.atan2(intersection.x, intersection.z)); | |
| yawAngle = clamp(yawAngle, YAW_MIN, YAW_MAX); | |
| }} | |
| }} else if (dragTarget.userData.type === 'pitch') {{ | |
| const plane = new THREE.Plane(new THREE.Vector3(1, 0, 0), -0.8); | |
| if (raycaster.ray.intersectPlane(plane, intersection)) {{ | |
| const relY = intersection.y - CENTER.y; | |
| const relZ = intersection.z; | |
| pitchAngle = THREE.MathUtils.radToDeg(Math.atan2(relY, relZ)); | |
| pitchAngle = clamp(pitchAngle, PITCH_MIN, PITCH_MAX); | |
| }} | |
| }} else if (dragTarget.userData.type === 'zoom') {{ | |
| const deltaY = mouse.y - dragStartMouse.y; | |
| zoomValue = clamp(dragStartZoom - deltaY * 1.5, -1, 1); | |
| }} | |
| updatePositions(); | |
| }} | |
| }}, {{ passive: false }}); | |
| canvas.addEventListener('touchend', (e) => {{ e.preventDefault(); onMouseUp(); }}, {{ passive: false }}); | |
| canvas.addEventListener('touchcancel', (e) => {{ e.preventDefault(); onMouseUp(); }}, {{ passive: false }}); | |
| updatePositions(); | |
| function render() {{ | |
| requestAnimationFrame(render); | |
| renderer.render(scene, camera); | |
| }} | |
| render(); | |
| new ResizeObserver(() => {{ | |
| camera.aspect = wrapper.clientWidth / wrapper.clientHeight; | |
| camera.updateProjectionMatrix(); | |
| renderer.setSize(wrapper.clientWidth, wrapper.clientHeight); | |
| }}).observe(wrapper); | |
| let lastImageUrl = props.imageUrl; | |
| let lastValue = JSON.stringify(props.value); | |
| setInterval(() => {{ | |
| if (props.imageUrl !== lastImageUrl) {{ | |
| lastImageUrl = props.imageUrl; | |
| updateTextureFromUrl(props.imageUrl); | |
| }} | |
| const currentValue = JSON.stringify(props.value); | |
| if (currentValue !== lastValue) {{ | |
| lastValue = currentValue; | |
| if (props.value && typeof props.value === 'object') {{ | |
| yawAngle = props.value.yaw ?? yawAngle; | |
| pitchAngle = props.value.pitch ?? pitchAngle; | |
| zoomValue = props.value.zoom_delta ?? zoomValue; | |
| updatePositions(); | |
| }} | |
| }} | |
| }}, 100); | |
| }} | |
| waitForThreeAndInit(); | |
| }})(); | |
| """ if three_available else "" | |
| super().__init__( | |
| value=value, | |
| html_template=html_template, | |
| js_on_load=js_on_load, | |
| imageUrl=imageUrl, | |
| **kwargs, | |
| ) | |
| class BoxSelector(gr.HTML): | |
| def __init__(self, value=None, imageUrl=None, **kwargs): | |
| html_template = """ | |
| <div class="se-box-root"> | |
| <style> | |
| .se-box-wrap { position: relative; width: 100%; max-width: 520px; overflow: hidden; border-radius: 8px; } | |
| .se-box-img { | |
| width: 100%; | |
| display: none; | |
| border-radius: 8px; | |
| border: 1px solid #444; | |
| cursor: crosshair; | |
| user-select: none; | |
| -webkit-user-drag: none; | |
| touch-action: none; | |
| } | |
| .se-box-empty { | |
| width: 100%; | |
| min-height: 220px; | |
| display: flex; | |
| align-items: center; | |
| justify-content: center; | |
| border-radius: 8px; | |
| border: 1px dashed #666; | |
| color: #9a9a9a; | |
| font-size: 13px; | |
| text-align: center; | |
| box-sizing: border-box; | |
| padding: 16px; | |
| } | |
| .se-box-overlay { | |
| position: absolute; | |
| border: 2px solid #ff2d2d; | |
| background: rgba(255, 45, 45, 0.06); | |
| box-sizing: border-box; | |
| pointer-events: none; | |
| display: none; | |
| } | |
| .se-box-hint { margin-top: 6px; color: #bbb; font-size: 12px; } | |
| </style> | |
| <div class="se-box-wrap"> | |
| <div class="se-box-empty">Upload an image above, then drag here to draw the box.</div> | |
| <img class="se-box-img" alt="box selector" /> | |
| <div class="se-box-overlay"></div> | |
| </div> | |
| <div class="se-box-hint">Drag the mouse to draw the red box (real-time preview).</div> | |
| </div> | |
| """ | |
| js_on_load = f""" | |
| (() => {{ | |
| const root = element.querySelector('.se-box-root'); | |
| if (!root) return; | |
| const img = root.querySelector('.se-box-img'); | |
| const empty = root.querySelector('.se-box-empty'); | |
| const overlay = root.querySelector('.se-box-overlay'); | |
| const DEFAULT_BOX = {{ | |
| cx: {DEFAULT_BOX_CX}, | |
| cy: {DEFAULT_BOX_CY}, | |
| w: {DEFAULT_BOX_W}, | |
| h: {DEFAULT_BOX_H}, | |
| }}; | |
| let dragging = false; | |
| let startX = 0; | |
| let startY = 0; | |
| let currentBox = props.value || null; | |
| let lastHiddenValue = ""; | |
| let lastImageUrl = props.imageUrl || ""; | |
| let lastValue = JSON.stringify(props.value || null); | |
| function clamp(v, lo, hi) {{ | |
| return Math.max(lo, Math.min(hi, v)); | |
| }} | |
| function sanitizeBox(box) {{ | |
| if (!box || typeof box !== 'object') return null; | |
| const w = clamp(Number(box?.w ?? DEFAULT_BOX.w) || DEFAULT_BOX.w, 0.001, 1.0); | |
| const h = clamp(Number(box?.h ?? DEFAULT_BOX.h) || DEFAULT_BOX.h, 0.001, 1.0); | |
| const cx = clamp(Number(box?.cx ?? DEFAULT_BOX.cx) || DEFAULT_BOX.cx, w / 2, 1 - w / 2); | |
| const cy = clamp(Number(box?.cy ?? DEFAULT_BOX.cy) || DEFAULT_BOX.cy, h / 2, 1 - h / 2); | |
| return {{ cx, cy, w, h }}; | |
| }} | |
| function getHiddenInput() {{ | |
| return document.querySelector('#se_box_json_object textarea, #se_box_json_object input'); | |
| }} | |
| function showPlaceholder() {{ | |
| img.style.display = 'none'; | |
| empty.style.display = 'flex'; | |
| overlay.style.display = 'none'; | |
| }} | |
| function showImage() {{ | |
| img.style.display = 'block'; | |
| empty.style.display = 'none'; | |
| }} | |
| function setOverlayFromNorm(box) {{ | |
| currentBox = sanitizeBox(box); | |
| if (!img.src) {{ | |
| showPlaceholder(); | |
| return; | |
| }} | |
| showImage(); | |
| if (!currentBox) {{ | |
| overlay.style.display = 'none'; | |
| return; | |
| }} | |
| const rect = img.getBoundingClientRect(); | |
| if (!rect.width || !rect.height) {{ | |
| overlay.style.display = 'none'; | |
| return; | |
| }} | |
| const x = (currentBox.cx - currentBox.w / 2) * rect.width; | |
| const y = (currentBox.cy - currentBox.h / 2) * rect.height; | |
| const w = currentBox.w * rect.width; | |
| const h = currentBox.h * rect.height; | |
| overlay.style.left = `${{x}}px`; | |
| overlay.style.top = `${{y}}px`; | |
| overlay.style.width = `${{w}}px`; | |
| overlay.style.height = `${{h}}px`; | |
| overlay.style.display = 'block'; | |
| }} | |
| function syncHiddenFromBox(box) {{ | |
| const hidden = getHiddenInput(); | |
| if (!hidden) return; | |
| const sanitized = sanitizeBox(box); | |
| const payload = sanitized ? JSON.stringify(sanitized) : ''; | |
| hidden.value = payload; | |
| hidden.dispatchEvent(new Event('input', {{ bubbles: true }})); | |
| hidden.dispatchEvent(new Event('change', {{ bubbles: true }})); | |
| lastHiddenValue = payload; | |
| }} | |
| function syncBoxFromHidden() {{ | |
| const hidden = getHiddenInput(); | |
| if (!hidden) return; | |
| const value = String(hidden.value || '').trim(); | |
| if (value === lastHiddenValue) return; | |
| if (!value) {{ | |
| lastHiddenValue = value; | |
| currentBox = null; | |
| setOverlayFromNorm(null); | |
| return; | |
| }} | |
| try {{ | |
| currentBox = sanitizeBox(JSON.parse(value)); | |
| lastHiddenValue = value; | |
| setOverlayFromNorm(currentBox); | |
| }} catch (_err) {{}} | |
| }} | |
| function pointToNorm(clientX, clientY) {{ | |
| const rect = img.getBoundingClientRect(); | |
| const x = clamp((clientX - rect.left) / rect.width, 0, 1); | |
| const y = clamp((clientY - rect.top) / rect.height, 0, 1); | |
| return {{ x, y }}; | |
| }} | |
| function beginDrag(clientX, clientY) {{ | |
| if (!img.src) return; | |
| dragging = true; | |
| const p = pointToNorm(clientX, clientY); | |
| startX = p.x; | |
| startY = p.y; | |
| setOverlayFromNorm({{ cx: startX, cy: startY, w: 0.001, h: 0.001 }}); | |
| }} | |
| function updateDrag(clientX, clientY) {{ | |
| if (!dragging) return; | |
| const p = pointToNorm(clientX, clientY); | |
| const x1 = Math.min(startX, p.x); | |
| const x2 = Math.max(startX, p.x); | |
| const y1 = Math.min(startY, p.y); | |
| const y2 = Math.max(startY, p.y); | |
| setOverlayFromNorm({{ | |
| cx: (x1 + x2) / 2, | |
| cy: (y1 + y2) / 2, | |
| w: Math.max(0.001, x2 - x1), | |
| h: Math.max(0.001, y2 - y1), | |
| }}); | |
| }} | |
| function endDrag(clientX, clientY) {{ | |
| if (!dragging) return; | |
| dragging = false; | |
| const p = pointToNorm(clientX, clientY); | |
| const x1 = Math.min(startX, p.x); | |
| const x2 = Math.max(startX, p.x); | |
| const y1 = Math.min(startY, p.y); | |
| const y2 = Math.max(startY, p.y); | |
| const box = sanitizeBox({{ | |
| cx: (x1 + x2) / 2, | |
| cy: (y1 + y2) / 2, | |
| w: Math.max(0.001, x2 - x1), | |
| h: Math.max(0.001, y2 - y1), | |
| }}); | |
| setOverlayFromNorm(box); | |
| syncHiddenFromBox(box); | |
| }} | |
| img.addEventListener('load', () => setOverlayFromNorm(currentBox)); | |
| img.addEventListener('mousedown', (ev) => {{ | |
| if (ev.button !== 0) return; | |
| beginDrag(ev.clientX, ev.clientY); | |
| }}); | |
| window.addEventListener('mousemove', (ev) => {{ | |
| updateDrag(ev.clientX, ev.clientY); | |
| }}); | |
| window.addEventListener('mouseup', (ev) => {{ | |
| endDrag(ev.clientX, ev.clientY); | |
| }}); | |
| img.addEventListener('touchstart', (ev) => {{ | |
| if (!ev.touches || !ev.touches.length) return; | |
| const touch = ev.touches[0]; | |
| beginDrag(touch.clientX, touch.clientY); | |
| ev.preventDefault(); | |
| }}, {{ passive: false }}); | |
| window.addEventListener('touchmove', (ev) => {{ | |
| if (!dragging || !ev.touches || !ev.touches.length) return; | |
| const touch = ev.touches[0]; | |
| updateDrag(touch.clientX, touch.clientY); | |
| ev.preventDefault(); | |
| }}, {{ passive: false }}); | |
| window.addEventListener('touchend', (ev) => {{ | |
| const touch = (ev.changedTouches && ev.changedTouches[0]) || null; | |
| if (!touch) return; | |
| endDrag(touch.clientX, touch.clientY); | |
| ev.preventDefault(); | |
| }}, {{ passive: false }}); | |
| function syncFromProps() {{ | |
| const nextImageUrl = props.imageUrl || ''; | |
| if (nextImageUrl !== lastImageUrl) {{ | |
| lastImageUrl = nextImageUrl; | |
| if (nextImageUrl) {{ | |
| img.src = nextImageUrl; | |
| }} else {{ | |
| img.removeAttribute('src'); | |
| showPlaceholder(); | |
| }} | |
| }} | |
| const currentValueText = JSON.stringify(props.value || null); | |
| if (currentValueText !== lastValue) {{ | |
| lastValue = currentValueText; | |
| currentBox = sanitizeBox(props.value || null); | |
| setOverlayFromNorm(currentBox); | |
| }} | |
| syncBoxFromHidden(); | |
| }} | |
| syncFromProps(); | |
| setInterval(syncFromProps, 120); | |
| window.addEventListener('resize', () => setOverlayFromNorm(currentBox)); | |
| }})(); | |
| """ | |
| super().__init__( | |
| value=value, | |
| html_template=html_template, | |
| js_on_load=js_on_load, | |
| imageUrl=imageUrl, | |
| **kwargs, | |
| ) | |
| class ObjectRotationControl3D(gr.HTML): | |
| def __init__(self, value=None, three_available: bool = False, **kwargs): | |
| if value is None: | |
| value = {"yaw": 0.0, "view": ROTATION_VIEW_ORDER[0]} | |
| offline_msg = """ | |
| <div style="width:100%;height:320px;display:flex;align-items:center;justify-content:center;background:linear-gradient(180deg,#ffffff 0%,#f3f6fb 100%);border:1px solid #d8e1ee;border-radius:12px;color:#334155;padding:24px;box-sizing:border-box;text-align:center;"> | |
| <div> | |
| <div style="font-size:18px;font-weight:600;margin-bottom:10px;">Object rotation preview is disabled</div> | |
| <div style="font-size:13px;line-height:1.6;opacity:0.9;"> | |
| No local three.js was provided.<br/> | |
| Put <code>three.min.js</code> next to this <code>.py</code> file to enable 3D rotation. | |
| </div> | |
| </div> | |
| </div> | |
| """ | |
| html_template = """ | |
| <div id="object-rotation-wrapper" style="width:100%;height:320px;position:relative;background:radial-gradient(circle at top, #ffffff 0%, #fbfdff 52%, #f3f7fb 100%);border:1px solid #d8e1ee;border-radius:12px;overflow:hidden;"> | |
| <div style=" | |
| position:absolute; | |
| top:12px; | |
| left:12px; | |
| padding:7px 10px; | |
| border-radius:999px; | |
| background:rgba(255,255,255,0.96); | |
| border:1px solid rgba(100,116,139,0.24); | |
| box-shadow:0 10px 24px rgba(15,23,42,0.06); | |
| font-size:11px; | |
| font-weight:600; | |
| letter-spacing:0.03em; | |
| color:#334155; | |
| z-index:10; | |
| ">Drag to rotate · 8 fixed views</div> | |
| <div id="object-rotation-overlay" style=" | |
| position:absolute; | |
| bottom:12px; | |
| left:50%; | |
| transform:translateX(-50%); | |
| background:rgba(255,255,255,0.96); | |
| color:#9a3412; | |
| padding:10px 14px; | |
| border-radius:8px; | |
| border:1px solid rgba(100,116,139,0.24); | |
| box-shadow:0 14px 34px rgba(15,23,42,0.10); | |
| font-family:monospace; | |
| font-size:12px; | |
| text-align:center; | |
| z-index:10; | |
| min-width:180px; | |
| "></div> | |
| </div> | |
| """ if three_available else offline_msg | |
| js_on_load = f""" | |
| (() => {{ | |
| function waitForThreeAndInit() {{ | |
| if (typeof THREE === 'undefined') {{ | |
| setTimeout(waitForThreeAndInit, 50); | |
| return; | |
| }} | |
| init3D(); | |
| }} | |
| function init3D() {{ | |
| const VIEW_ORDER = {json.dumps(ROTATION_VIEW_ORDER, ensure_ascii=True)}; | |
| const wrapper = element.querySelector('#object-rotation-wrapper'); | |
| const overlay = element.querySelector('#object-rotation-overlay'); | |
| if (!wrapper || !overlay) return; | |
| function normalizeYaw(yaw) {{ | |
| let value = Number(yaw || 0); | |
| value = ((value % 360) + 360) % 360; | |
| return value; | |
| }} | |
| function snapYaw(yaw) {{ | |
| return (Math.round(normalizeYaw(yaw) / 45) * 45) % 360; | |
| }} | |
| function viewFromYaw(yaw) {{ | |
| const snapped = snapYaw(yaw); | |
| const index = Math.round(snapped / 45) % 8; | |
| return VIEW_ORDER[index]; | |
| }} | |
| let yawAngle = snapYaw(props.value?.yaw ?? 0); | |
| let currentView = props.value?.view || viewFromYaw(yawAngle); | |
| let isDragging = false; | |
| let dragStartX = 0; | |
| let dragStartYaw = yawAngle; | |
| let lastValue = JSON.stringify(props.value || {{ yaw: 0, view: VIEW_ORDER[0] }}); | |
| const scene = new THREE.Scene(); | |
| scene.background = new THREE.Color(0xfcfdff); | |
| const camera = new THREE.PerspectiveCamera(42, wrapper.clientWidth / wrapper.clientHeight, 0.1, 1000); | |
| camera.position.set(0, 1.2, 4.6); | |
| camera.lookAt(0, 0.55, 0); | |
| const renderer = new THREE.WebGLRenderer({{ antialias: true, alpha: true }}); | |
| renderer.setSize(wrapper.clientWidth, wrapper.clientHeight); | |
| renderer.setPixelRatio(Math.min(window.devicePixelRatio, 2)); | |
| wrapper.insertBefore(renderer.domElement, overlay); | |
| scene.add(new THREE.AmbientLight(0xffffff, 1.12)); | |
| const keyLight = new THREE.DirectionalLight(0xffffff, 1.16); | |
| keyLight.position.set(3, 4.5, 5); | |
| scene.add(keyLight); | |
| const rimLight = new THREE.DirectionalLight(0xcbd5e1, 0.42); | |
| rimLight.position.set(-4, 2.2, -3); | |
| scene.add(rimLight); | |
| const warmLight = new THREE.DirectionalLight(0xfed7aa, 0.22); | |
| warmLight.position.set(0, 3, 4); | |
| scene.add(warmLight); | |
| const floor = new THREE.Mesh( | |
| new THREE.CircleGeometry(2.3, 48), | |
| new THREE.MeshStandardMaterial({{ | |
| color: 0xf7fafc, | |
| transparent: true, | |
| opacity: 0.995, | |
| roughness: 0.96, | |
| }}) | |
| ); | |
| floor.rotation.x = -Math.PI / 2; | |
| floor.position.y = -0.55; | |
| scene.add(floor); | |
| const grid = new THREE.GridHelper(3.6, 12, 0xd1d9e6, 0xe6ebf2); | |
| grid.position.y = -0.545; | |
| scene.add(grid); | |
| const ring = new THREE.Mesh( | |
| new THREE.RingGeometry(1.45, 1.62, 64), | |
| new THREE.MeshBasicMaterial({{ | |
| color: 0xd6dee8, | |
| transparent: true, | |
| opacity: 0.78, | |
| side: THREE.DoubleSide, | |
| }}) | |
| ); | |
| ring.rotation.x = -Math.PI / 2; | |
| ring.position.y = -0.545; | |
| scene.add(ring); | |
| const markerGroup = new THREE.Group(); | |
| for (let i = 0; i < 8; i++) {{ | |
| const angle = THREE.MathUtils.degToRad(i * 45); | |
| const marker = new THREE.Mesh( | |
| new THREE.SphereGeometry(0.04, 10, 10), | |
| new THREE.MeshStandardMaterial({{ color: 0x64748b }}) | |
| ); | |
| marker.position.set(Math.sin(angle) * 1.7, -0.5, Math.cos(angle) * 1.7); | |
| markerGroup.add(marker); | |
| }} | |
| scene.add(markerGroup); | |
| const objectGroup = new THREE.Group(); | |
| scene.add(objectGroup); | |
| const body = new THREE.Mesh( | |
| new THREE.BoxGeometry(1.3, 1.05, 0.9), | |
| [ | |
| new THREE.MeshStandardMaterial({{ color: 0xb8c4d3, metalness: 0.16, roughness: 0.5 }}), | |
| new THREE.MeshStandardMaterial({{ color: 0xb8c4d3, metalness: 0.16, roughness: 0.5 }}), | |
| new THREE.MeshStandardMaterial({{ color: 0xea580c, metalness: 0.12, roughness: 0.34 }}), | |
| new THREE.MeshStandardMaterial({{ color: 0x1f2937, metalness: 0.22, roughness: 0.28 }}), | |
| new THREE.MeshStandardMaterial({{ color: 0x64748b, metalness: 0.14, roughness: 0.42 }}), | |
| new THREE.MeshStandardMaterial({{ color: 0x64748b, metalness: 0.14, roughness: 0.42 }}), | |
| ] | |
| ); | |
| body.position.y = 0.15; | |
| objectGroup.add(body); | |
| const topCap = new THREE.Mesh( | |
| new THREE.CylinderGeometry(0.28, 0.28, 1.15, 24), | |
| new THREE.MeshStandardMaterial({{ color: 0x94a3b8, metalness: 0.24, roughness: 0.42 }}) | |
| ); | |
| topCap.rotation.z = Math.PI / 2; | |
| topCap.position.set(0, 0.55, 0); | |
| objectGroup.add(topCap); | |
| const frontArrow = new THREE.Mesh( | |
| new THREE.ConeGeometry(0.14, 0.35, 20), | |
| new THREE.MeshStandardMaterial({{ color: 0xea580c, emissive: 0xfb923c, emissiveIntensity: 0.12 }}) | |
| ); | |
| frontArrow.rotation.x = Math.PI / 2; | |
| frontArrow.position.set(0, 0.18, 0.72); | |
| objectGroup.add(frontArrow); | |
| const frontPlate = new THREE.Mesh( | |
| new THREE.PlaneGeometry(0.62, 0.2), | |
| new THREE.MeshBasicMaterial({{ color: 0xfff7ed }}) | |
| ); | |
| frontPlate.position.set(0, 0.18, 0.455); | |
| objectGroup.add(frontPlate); | |
| const shadow = new THREE.Mesh( | |
| new THREE.CircleGeometry(0.86, 40), | |
| new THREE.MeshBasicMaterial({{ | |
| color: 0x64748b, | |
| transparent: true, | |
| opacity: 0.14, | |
| }}) | |
| ); | |
| shadow.rotation.x = -Math.PI / 2; | |
| shadow.position.y = -0.535; | |
| shadow.scale.set(1.0, 0.72, 1.0); | |
| objectGroup.add(shadow); | |
| function setOverlayText() {{ | |
| overlay.textContent = `View: ${{currentView}}`; | |
| }} | |
| function updateObjectRotation() {{ | |
| objectGroup.rotation.y = THREE.MathUtils.degToRad(yawAngle); | |
| currentView = viewFromYaw(yawAngle); | |
| const activeIndex = Math.round(snapYaw(yawAngle) / 45) % 8; | |
| markerGroup.children.forEach((marker, index) => {{ | |
| marker.material.color.setHex(index === activeIndex ? 0xea580c : 0x64748b); | |
| marker.scale.setScalar(index === activeIndex ? 1.55 : 1.0); | |
| }}); | |
| setOverlayText(); | |
| }} | |
| function emitChange() {{ | |
| props.value = {{ | |
| yaw: snapYaw(yawAngle), | |
| view: viewFromYaw(yawAngle), | |
| }}; | |
| trigger('change', props.value); | |
| }} | |
| const canvas = renderer.domElement; | |
| canvas.addEventListener('mousedown', (e) => {{ | |
| isDragging = true; | |
| dragStartX = e.clientX; | |
| dragStartYaw = yawAngle; | |
| canvas.style.cursor = 'grabbing'; | |
| }}); | |
| window.addEventListener('mousemove', (e) => {{ | |
| if (!isDragging) return; | |
| const deltaX = e.clientX - dragStartX; | |
| yawAngle = normalizeYaw(dragStartYaw + deltaX * 0.45); | |
| updateObjectRotation(); | |
| }}); | |
| function finishDrag() {{ | |
| if (!isDragging) return; | |
| isDragging = false; | |
| yawAngle = snapYaw(yawAngle); | |
| updateObjectRotation(); | |
| emitChange(); | |
| canvas.style.cursor = 'grab'; | |
| }} | |
| window.addEventListener('mouseup', finishDrag); | |
| canvas.addEventListener('mouseleave', () => {{ | |
| if (!isDragging) canvas.style.cursor = 'grab'; | |
| }}); | |
| canvas.addEventListener('touchstart', (e) => {{ | |
| if (!e.touches || !e.touches.length) return; | |
| isDragging = true; | |
| dragStartX = e.touches[0].clientX; | |
| dragStartYaw = yawAngle; | |
| e.preventDefault(); | |
| }}, {{ passive: false }}); | |
| window.addEventListener('touchmove', (e) => {{ | |
| if (!isDragging || !e.touches || !e.touches.length) return; | |
| const deltaX = e.touches[0].clientX - dragStartX; | |
| yawAngle = normalizeYaw(dragStartYaw + deltaX * 0.45); | |
| updateObjectRotation(); | |
| e.preventDefault(); | |
| }}, {{ passive: false }}); | |
| window.addEventListener('touchend', (e) => {{ | |
| if (!isDragging) return; | |
| finishDrag(); | |
| e.preventDefault(); | |
| }}, {{ passive: false }}); | |
| function animate() {{ | |
| requestAnimationFrame(animate); | |
| renderer.render(scene, camera); | |
| }} | |
| function syncFromProps() {{ | |
| const currentValue = JSON.stringify(props.value || {{ yaw: 0, view: VIEW_ORDER[0] }}); | |
| if (currentValue === lastValue) return; | |
| lastValue = currentValue; | |
| yawAngle = snapYaw(props.value?.yaw ?? 0); | |
| currentView = props.value?.view || viewFromYaw(yawAngle); | |
| updateObjectRotation(); | |
| }} | |
| updateObjectRotation(); | |
| animate(); | |
| canvas.style.cursor = 'grab'; | |
| new ResizeObserver(() => {{ | |
| camera.aspect = wrapper.clientWidth / wrapper.clientHeight; | |
| camera.updateProjectionMatrix(); | |
| renderer.setSize(wrapper.clientWidth, wrapper.clientHeight); | |
| }}).observe(wrapper); | |
| setInterval(syncFromProps, 120); | |
| }} | |
| waitForThreeAndInit(); | |
| }})(); | |
| """ if three_available else "" | |
| super().__init__( | |
| value=value, | |
| html_template=html_template, | |
| js_on_load=js_on_load, | |
| **kwargs, | |
| ) | |
| def create_demo( | |
| app: EditorApp, | |
| three_available: bool, | |
| hide_advanced_options: bool = False, | |
| examples_table: Optional[List[List[Any]]] = None, | |
| examples_full: Optional[List[List[Any]]] = None, | |
| auto_pe: bool = False, | |
| default_save_dir: str = "", | |
| ): | |
| examples_table = examples_table or [] | |
| examples_full = examples_full or [] | |
| theme = gr.themes.Soft() | |
| page_css = """ | |
| #col-container { max-width: 1400px; margin: 0 auto; } | |
| .dark .progress-text { color: white !important; } | |
| #camera-3d-control { min-height: 450px; } | |
| .slider-row { display: flex; gap: 10px; align-items: center; } | |
| .se-hidden-json { | |
| position: absolute !important; | |
| left: -10000px !important; | |
| top: auto !important; | |
| width: 1px !important; | |
| height: 1px !important; | |
| overflow: hidden !important; | |
| opacity: 0 !important; | |
| pointer-events: none !important; | |
| } | |
| """ | |
| def default_camera_value() -> Dict[str, float]: | |
| return {"yaw": 0.0, "pitch": 0.0, "zoom_delta": 0} | |
| def preprocess_to_dit_bucket(image: Image.Image) -> Image.Image: | |
| image = ImageOps.exif_transpose(image).convert("RGB") | |
| width, height = image.size | |
| if width <= 0 or height <= 0: | |
| return image | |
| bucket_config = generate_video_image_bucket( | |
| basesize=BASE_SIZE, | |
| min_temporal=56, | |
| max_temporal=56, | |
| bs_img=4, | |
| bs_vid=4, | |
| bs_mimg=8, | |
| min_items=2, | |
| max_items=2, | |
| ) | |
| bucket_group = BucketGroup(bucket_config) | |
| bucket = bucket_group.find_best_bucket((1, 1, height, width)) | |
| target_height, target_width = int(bucket[-2]), int(bucket[-1]) | |
| processed = resize_center_crop(image, (target_height, target_width)) | |
| print( | |
| f"[Input] Preprocessed to DIT bucket: {width}x{height} -> " | |
| f"{target_width}x{target_height}" | |
| ) | |
| return processed | |
| def load_uploaded_image(uploaded_image: Any) -> Optional[Image.Image]: | |
| if uploaded_image is None: | |
| return None | |
| if isinstance(uploaded_image, Image.Image): | |
| return preprocess_to_dit_bucket(uploaded_image) | |
| if isinstance(uploaded_image, dict): | |
| maybe_path = uploaded_image.get("path") or uploaded_image.get("name") | |
| if maybe_path: | |
| return preprocess_to_dit_bucket(Image.open(maybe_path)) | |
| return preprocess_to_dit_bucket(Image.open(uploaded_image)) | |
| def estimate_spatial_edit_duration( | |
| original_image: Image.Image, | |
| boxed_image: Optional[Image.Image] = None, | |
| task_mode: str = TASK_MODE_CAMERA, | |
| rotation_value: Optional[Dict[str, Any]] = None, | |
| rotation_object_desc: str = "", | |
| yaw: float = 0.0, | |
| pitch: float = 0.0, | |
| zoom_delta: int = 0, | |
| prompt_text: str = "", | |
| seed: int = 0, | |
| randomize_seed: bool = False, | |
| guidance_scale: float = 4.0, | |
| num_inference_steps: int = 30, | |
| object_desc: str = "", | |
| ) -> int: | |
| max_duration = int(os.getenv("SPACES_GPU_MAX_DURATION", "300")) | |
| loaded_base = int(os.getenv("SPACES_GPU_LOADED_DURATION", "420")) | |
| cold_extra = int(os.getenv("SPACES_GPU_COLD_START_EXTRA", "900")) | |
| per_step = float(os.getenv("SPACES_GPU_STEP_DURATION", "4.0")) | |
| if boxed_image is not None: | |
| img = boxed_image | |
| else: | |
| img = original_image | |
| pixel_factor = 1.0 | |
| try: | |
| if img is not None: | |
| w, h = img.size | |
| pixel_factor = max(0.7, min(2.5, (float(w) * float(h)) / float(1024 * 1024))) | |
| except Exception: | |
| pixel_factor = 1.0 | |
| duration = loaded_base + int(max(0, int(num_inference_steps)) * per_step * pixel_factor) | |
| if not app.is_model_loaded(): | |
| duration += cold_extra | |
| elif app.is_model_on_cpu() and not app.is_model_on_gpu(): | |
| duration += int(os.getenv("SPACES_GPU_CPU_TO_GPU_EXTRA", "600")) | |
| return max(120, min(max_duration, duration)) | |
| def estimate_model_warmup_duration() -> int: | |
| return max(300, int(os.getenv("SPACES_GPU_WARMUP_DURATION", os.getenv("SPACES_GPU_MAX_DURATION", "300")))) | |
| def warmup_model_on_page_load() -> str: | |
| try: | |
| if app.is_model_loaded() and (app.is_model_on_gpu() or app.is_model_on_cpu()): | |
| if is_rank0(): | |
| print("[Model] Page-load warmup skipped: model already available.") | |
| return "ready" | |
| if is_rank0(): | |
| print("[Model] Page-load warmup started.") | |
| app._ensure_model_loaded(target_device=resolve_device()) | |
| if app._move_back_to_cpu_after_infer: | |
| app._move_model_back_to_cpu() | |
| if is_rank0(): | |
| print("[Model] Page-load warmup finished.") | |
| return "ready" | |
| except Exception as e: | |
| if is_rank0(): | |
| print(f"[Model] Page-load warmup failed: {e}") | |
| return f"error: {e}" | |
| def infer_camera_edit( | |
| original_image: Image.Image, | |
| boxed_image: Optional[Image.Image], | |
| task_mode: str = TASK_MODE_CAMERA, | |
| rotation_value: Optional[Dict[str, Any]] = None, | |
| rotation_object_desc: str = "", | |
| yaw: float = 0.0, | |
| pitch: float = 0.0, | |
| zoom_delta: int = 0, | |
| prompt_text: str = "", | |
| seed: int = 0, | |
| randomize_seed: bool = False, | |
| guidance_scale: float = 4.0, | |
| num_inference_steps: int = 30, | |
| object_desc: str = "", | |
| ): | |
| _ = gr.Progress(track_tqdm=True) | |
| if original_image is None: | |
| raise gr.Error("Please upload an image first.") | |
| use_object = task_mode == TASK_MODE_OBJECT | |
| use_rotation = task_mode == TASK_MODE_ROTATION | |
| use_camera = task_mode == TASK_MODE_CAMERA | |
| if use_object and boxed_image is None: | |
| raise gr.Error("Please draw a red target box first for Object Moving.") | |
| yaw = snap_yaw(float(yaw)) | |
| pitch = snap_pitch(float(pitch)) | |
| zoom_delta = snap_zoom_delta(int(round(zoom_delta))) | |
| object_desc_clean = (object_desc or "").strip() | |
| rotation_object_desc_clean = (rotation_object_desc or "").strip() | |
| object_desc_prompt = normalize_object_description_for_prompt(object_desc_clean) | |
| rotation_object_desc_prompt = normalize_object_description_for_prompt(rotation_object_desc_clean) | |
| object_prompt_auto = build_object_moving_prompt(object_desc_clean) | |
| rotation_prompt_auto = build_object_rotation_prompt(rotation_object_desc_clean, rotation_value) | |
| camera_prompt_auto = build_camera_instruction(yaw, pitch, zoom_delta, seed=0) | |
| if use_object and not object_desc_prompt: | |
| raise gr.Error("Please provide Object Description for Object Moving.") | |
| if use_rotation and not rotation_object_desc_prompt: | |
| raise gr.Error("Please provide Object Description for Object Rotation.") | |
| seed = random.randint(0, int(MAX_SEED)) if bool(randomize_seed) else int(seed) | |
| if use_rotation: | |
| final_prompt = (prompt_text or "").strip() or rotation_prompt_auto | |
| display_text = build_object_rotation_display(rotation_value) | |
| elif use_object: | |
| final_prompt = (prompt_text or "").strip() or object_prompt_auto | |
| display_text = "Object Moving" | |
| else: | |
| final_prompt = (prompt_text or "").strip() or camera_prompt_auto | |
| display_text = build_camera_display(yaw, pitch, zoom_delta) | |
| model_input_image = boxed_image if use_object else original_image | |
| pil_image = model_input_image.convert("RGB") | |
| out_h, out_w = app._choose_size(pil_image, BASE_SIZE, BASE_SIZE) | |
| print(f"[SpatialEdit] Task mode = {task_mode}") | |
| print(f"[SpatialEdit] Base prompt =\n{final_prompt}") | |
| print(f"[SpatialEdit] Chosen bucket resolution = {out_w}x{out_h}") | |
| print(f"[SpatialEdit] Auto PE = {bool(auto_pe)}") | |
| print(f"[SpatialEdit] Seed = {int(seed)}") | |
| print(f"[SpatialEdit] Randomize seed = {bool(randomize_seed)}") | |
| expected_caption = "" | |
| try: | |
| out, rewritten_prompt = app.run( | |
| image=pil_image, | |
| prompt=final_prompt, | |
| steps=int(num_inference_steps), | |
| guidance=float(guidance_scale), | |
| seed=int(seed), | |
| ) | |
| except torch.cuda.OutOfMemoryError as e: | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| raise gr.Error( | |
| "GPU out of memory during inference. On Spaces, use ZeroGPU/H200 or a larger dedicated GPU, " | |
| "or lower BASESIZE / inference steps." | |
| ) from e | |
| finally: | |
| if torch.cuda.is_available(): | |
| torch.cuda.synchronize() | |
| out_png_path = save_png_no_compression(out) | |
| print(f"[Output] Saved uncompressed PNG to: {out_png_path}") | |
| save_dir = str(default_save_dir or "").strip() | |
| if save_dir: | |
| try: | |
| saved_dir = save_generation_bundle( | |
| save_root=save_dir, | |
| input_image=pil_image, | |
| output_image=out, | |
| final_prompt=rewritten_prompt, | |
| base_prompt=final_prompt, | |
| expected_caption=expected_caption, | |
| display_text=display_text, | |
| yaw=yaw, | |
| pitch=pitch, | |
| zoom_delta=zoom_delta, | |
| seed=seed, | |
| guidance_scale=guidance_scale, | |
| num_inference_steps=num_inference_steps, | |
| enable_prompt_rewrite=app.enable_prompt_rewrite, | |
| rewrite_backend=app.rewrite_model, | |
| ) | |
| print(f"[Save] Saved input/prompt/output to: {saved_dir}") | |
| except Exception as e: | |
| print(f"[Save] Failed to save artifacts: {e}") | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| return out_png_path, int(seed), display_text | |
| def infer_object_moving( | |
| original_image: Image.Image, | |
| boxed_image: Image.Image, | |
| object_desc: str, | |
| prompt_text: str, | |
| seed: int, | |
| randomize_seed: bool, | |
| guidance_scale: float, | |
| num_inference_steps: int, | |
| ): | |
| out_png_path, _, _ = infer_camera_edit( | |
| original_image=original_image, | |
| boxed_image=boxed_image, | |
| task_mode=TASK_MODE_OBJECT, | |
| object_desc=object_desc, | |
| prompt_text=prompt_text, | |
| seed=seed, | |
| randomize_seed=randomize_seed, | |
| guidance_scale=guidance_scale, | |
| num_inference_steps=num_inference_steps, | |
| ) | |
| return out_png_path | |
| def infer_object_rotation( | |
| original_image: Image.Image, | |
| rotation_value: Optional[Dict[str, Any]], | |
| rotation_object_desc: str, | |
| prompt_text: str, | |
| seed: int, | |
| randomize_seed: bool, | |
| guidance_scale: float, | |
| num_inference_steps: int, | |
| ): | |
| out_png_path, _, _ = infer_camera_edit( | |
| original_image=original_image, | |
| boxed_image=None, | |
| task_mode=TASK_MODE_ROTATION, | |
| rotation_value=rotation_value, | |
| rotation_object_desc=rotation_object_desc, | |
| prompt_text=prompt_text, | |
| seed=seed, | |
| randomize_seed=randomize_seed, | |
| guidance_scale=guidance_scale, | |
| num_inference_steps=num_inference_steps, | |
| ) | |
| return out_png_path | |
| def infer_camera_only( | |
| original_image: Image.Image, | |
| yaw: float, | |
| pitch: float, | |
| zoom_delta: int, | |
| prompt_text: str, | |
| seed: int, | |
| randomize_seed: bool, | |
| guidance_scale: float, | |
| num_inference_steps: int, | |
| ): | |
| return infer_camera_edit( | |
| original_image=original_image, | |
| boxed_image=None, | |
| task_mode=TASK_MODE_CAMERA, | |
| yaw=yaw, | |
| pitch=pitch, | |
| zoom_delta=zoom_delta, | |
| prompt_text=prompt_text, | |
| seed=seed, | |
| randomize_seed=randomize_seed, | |
| guidance_scale=guidance_scale, | |
| num_inference_steps=num_inference_steps, | |
| ) | |
| def update_object_prompt_from_description(object_desc: str): | |
| return build_object_moving_prompt(object_desc) | |
| def update_rotation_prompt_from_controls(rotation_value, object_desc: str): | |
| return build_object_rotation_prompt(object_desc, rotation_value) | |
| def update_prompt_from_controls(yaw, pitch, zoom_delta): | |
| yaw = snap_yaw(float(yaw)) | |
| pitch = snap_pitch(float(pitch)) | |
| zoom_delta = snap_zoom_delta(int(round(zoom_delta))) | |
| display_text = build_camera_display(yaw, pitch, zoom_delta) | |
| prompt_text = build_camera_instruction(yaw, pitch, zoom_delta, seed=0) | |
| return display_text, prompt_text | |
| def sync_3d_to_sliders(camera_value): | |
| if camera_value and isinstance(camera_value, dict): | |
| yaw = snap_yaw(float(camera_value.get("yaw", 0.0))) | |
| pitch = snap_pitch(float(camera_value.get("pitch", 0.0))) | |
| zoom_delta = snap_zoom_delta(int(round(camera_value.get("zoom_delta", 0)))) | |
| display_text = build_camera_display(yaw, pitch, zoom_delta) | |
| prompt_text = build_camera_instruction(yaw, pitch, zoom_delta, seed=0) | |
| return yaw, pitch, zoom_delta, display_text, prompt_text | |
| return gr.update(), gr.update(), gr.update(), gr.update(), gr.update() | |
| def sync_sliders_to_3d(yaw, pitch, zoom_delta): | |
| return { | |
| "yaw": snap_yaw(float(yaw)), | |
| "pitch": snap_pitch(float(pitch)), | |
| "zoom_delta": snap_zoom_delta(int(round(zoom_delta))), | |
| } | |
| def sanitize_box_values(box_cx, box_cy, box_w, box_h): | |
| box_w = min(max(float(box_w), 0.001), 1.0) | |
| box_h = min(max(float(box_h), 0.001), 1.0) | |
| box_cx = min(max(float(box_cx), box_w / 2.0), 1.0 - box_w / 2.0) | |
| box_cy = min(max(float(box_cy), box_h / 2.0), 1.0 - box_h / 2.0) | |
| return box_cx, box_cy, box_w, box_h | |
| def empty_box_json() -> str: | |
| return "" | |
| def default_box_json() -> str: | |
| return json.dumps( | |
| { | |
| "cx": DEFAULT_BOX_CX, | |
| "cy": DEFAULT_BOX_CY, | |
| "w": DEFAULT_BOX_W, | |
| "h": DEFAULT_BOX_H, | |
| }, | |
| ensure_ascii=True, | |
| ) | |
| def default_box_value() -> Optional[Dict[str, float]]: | |
| return None | |
| def parse_box_json(box_json_value: Any): | |
| if not box_json_value: | |
| return None | |
| try: | |
| if isinstance(box_json_value, str): | |
| box_json_value = json.loads(box_json_value) | |
| if not isinstance(box_json_value, dict): | |
| return None | |
| return sanitize_box_values( | |
| box_json_value.get("cx", DEFAULT_BOX_CX), | |
| box_json_value.get("cy", DEFAULT_BOX_CY), | |
| box_json_value.get("w", DEFAULT_BOX_W), | |
| box_json_value.get("h", DEFAULT_BOX_H), | |
| ) | |
| except Exception: | |
| return None | |
| def selector_source_data_url(image: Optional[Image.Image]) -> str: | |
| if image is None: | |
| return "" | |
| return pil_to_data_url(image) | |
| def make_box_preview(original_image, box_cx, box_cy, box_w, box_h): | |
| if original_image is None: | |
| return None | |
| boxed_image = draw_red_box_on_image(original_image, box_cx, box_cy, box_w, box_h) | |
| return boxed_image | |
| def update_boxed_image_state(original_image, box_json_value): | |
| if original_image is None: | |
| return None | |
| parsed_box = parse_box_json(box_json_value) | |
| if parsed_box is None: | |
| return None | |
| box_cx, box_cy, box_w, box_h = parsed_box | |
| return make_box_preview(original_image, box_cx, box_cy, box_w, box_h) | |
| def on_object_upload(uploaded_image): | |
| pil_image = load_uploaded_image(uploaded_image) | |
| if pil_image is None: | |
| return None, None, gr.update(value=default_box_value(), imageUrl=None), empty_box_json() | |
| return ( | |
| pil_image, | |
| None, | |
| gr.update(value=default_box_value(), imageUrl=selector_source_data_url(pil_image)), | |
| empty_box_json(), | |
| ) | |
| def on_rotation_upload(uploaded_image): | |
| pil_image = load_uploaded_image(uploaded_image) | |
| if pil_image is None: | |
| return None, default_rotation_value() | |
| return pil_image, default_rotation_value() | |
| def on_camera_upload(uploaded_image): | |
| pil_image = load_uploaded_image(uploaded_image) | |
| if pil_image is None: | |
| return ( | |
| None, | |
| build_camera_preview_html(None), | |
| gr.update(value=default_camera_value(), imageUrl=None), | |
| 0.0, | |
| 0.0, | |
| 0, | |
| build_camera_display(0.0, 0.0, 0), | |
| build_camera_instruction(0.0, 0.0, 0, seed=0), | |
| ) | |
| data_url = selector_source_data_url(pil_image) | |
| return ( | |
| pil_image, | |
| build_camera_preview_html(pil_image), | |
| gr.update(value=default_camera_value(), imageUrl=data_url), | |
| 0.0, | |
| 0.0, | |
| 0, | |
| build_camera_display(0.0, 0.0, 0), | |
| build_camera_instruction(0.0, 0.0, 0, seed=0), | |
| ) | |
| def reset_object_upload(): | |
| return None, None, gr.update(value=default_box_value(), imageUrl=None), empty_box_json() | |
| def reset_rotation_upload(): | |
| return None, default_rotation_value() | |
| def reset_camera_upload(): | |
| return ( | |
| None, | |
| build_camera_preview_html(None), | |
| gr.update(value=default_camera_value(), imageUrl=None), | |
| 0.0, | |
| 0.0, | |
| 0, | |
| build_camera_display(0.0, 0.0, 0), | |
| build_camera_instruction(0.0, 0.0, 0, seed=0), | |
| ) | |
| def load_camera_example_by_index(example_index): | |
| try: | |
| idx = int(example_index) | |
| except Exception: | |
| idx = -1 | |
| if idx < 0 or idx >= len(examples_full): | |
| return ( | |
| None, | |
| build_camera_preview_html(None), | |
| gr.update(value=default_camera_value(), imageUrl=None), | |
| 0.0, | |
| 0.0, | |
| 0, | |
| build_camera_display(0.0, 0.0, 0), | |
| build_camera_instruction(0.0, 0.0, 0, seed=0), | |
| ) | |
| image_path, yaw, pitch, zoom_delta, prompt = examples_full[idx] | |
| pil_image = load_uploaded_image(image_path) | |
| if pil_image is None: | |
| return ( | |
| None, | |
| build_camera_preview_html(None), | |
| gr.update(value=default_camera_value(), imageUrl=None), | |
| 0.0, | |
| 0.0, | |
| 0, | |
| build_camera_display(0.0, 0.0, 0), | |
| build_camera_instruction(0.0, 0.0, 0, seed=0), | |
| ) | |
| yaw = snap_yaw(float(yaw)) | |
| pitch = snap_pitch(float(pitch)) | |
| zoom_delta = snap_zoom_delta(int(round(zoom_delta))) | |
| prompt = str(prompt or "").strip() or build_camera_instruction(yaw, pitch, zoom_delta, seed=0) | |
| data_url = selector_source_data_url(pil_image) | |
| return ( | |
| pil_image, | |
| build_camera_preview_html(pil_image), | |
| gr.update(value={"yaw": yaw, "pitch": pitch, "zoom_delta": zoom_delta}, imageUrl=data_url), | |
| yaw, | |
| pitch, | |
| zoom_delta, | |
| build_camera_display(yaw, pitch, zoom_delta), | |
| prompt, | |
| ) | |
| def set_run_button_busy(): | |
| return gr.update(value=RUN_BUTTON_BUSY_TEXT, interactive=False) | |
| def set_run_button_idle(): | |
| return gr.update(value=RUN_BUTTON_IDLE_TEXT, interactive=True) | |
| with gr.Blocks(title="JoyAI-Image-Edit Spatial Editing Demo", theme=theme) as demo: | |
| gr.Markdown("# JoyAI-Image-Edit Spatial Editing Demo") | |
| with gr.Tabs(): | |
| with gr.Tab(TASK_MODE_OBJECT): | |
| object_original_image_state = gr.State(value=None) | |
| object_boxed_image_state = gr.State(value=None) | |
| gr.Markdown("### 📦 Object Moving") | |
| gr.Markdown("Use this mode to move a target object to a new location. First upload an image, then drag on the canvas to draw the destination red box, enter the description of the object you want to move, review or edit the generated editing prompt, and click Generate.") | |
| object_uploader = gr.File( | |
| label="Upload Image", | |
| file_types=["image"], | |
| type="filepath", | |
| ) | |
| object_box_json = gr.Textbox( | |
| value=empty_box_json(), | |
| show_label=False, | |
| elem_id="se_box_json_object", | |
| elem_classes=["se-hidden-json"], | |
| ) | |
| object_box_selector = BoxSelector( | |
| value=default_box_value(), | |
| imageUrl=None, | |
| ) | |
| object_desc = gr.Textbox( | |
| label="Object Description (describe the object you want to move)", | |
| placeholder="e.g., 'the cat on the sofa', 'the golden trophy on the table'", | |
| lines=2, | |
| info="Describe the object that should be moved into the red box. Use a clear visual description so the model can identify the correct object.", | |
| ) | |
| object_result = gr.Image(label="Output Image", height=500, format="png") | |
| object_run_btn = gr.Button(RUN_BUTTON_IDLE_TEXT, variant="primary", size="lg") | |
| object_prompt_box = gr.Textbox( | |
| label="Editing Prompt", | |
| value="", | |
| interactive=True, | |
| lines=3, | |
| placeholder="The editing prompt will be automatically generated from the controls above. You can further edit it before clicking Generate.", | |
| info="This prompt is auto-generated from the current controls. You can further modify it before generation.", | |
| ) | |
| default_steps = 30 | |
| with gr.Accordion("⚙️ Advanced Settings", open=False): | |
| if hide_advanced_options: | |
| object_seed = gr.Slider(visible=False, minimum=0, maximum=int(MAX_SEED), step=1, value=0) | |
| object_randomize_seed = gr.Checkbox(visible=False, value=True) | |
| object_guidance_scale = gr.Slider(visible=False, minimum=0.0, maximum=20.0, step=0.1, value=4.0) | |
| object_num_inference_steps = gr.Slider(visible=False, minimum=1, maximum=100, step=1, value=default_steps) | |
| else: | |
| object_seed = gr.Slider(label="Seed", minimum=0, maximum=int(MAX_SEED), step=1, value=0) | |
| object_randomize_seed = gr.Checkbox(label="Randomize Seed", value=True, interactive=True) | |
| object_guidance_scale = gr.Slider(label="Guidance Scale", minimum=0.0, maximum=20.0, step=0.1, value=4.0) | |
| object_num_inference_steps = gr.Slider(label="Inference Steps", minimum=1, maximum=100, step=1, value=default_steps) | |
| object_uploader.upload( | |
| fn=on_object_upload, | |
| inputs=[object_uploader], | |
| outputs=[object_original_image_state, object_boxed_image_state, object_box_selector, object_box_json], | |
| ) | |
| object_uploader.change( | |
| fn=on_object_upload, | |
| inputs=[object_uploader], | |
| outputs=[object_original_image_state, object_boxed_image_state, object_box_selector, object_box_json], | |
| ) | |
| object_uploader.clear( | |
| fn=reset_object_upload, | |
| outputs=[object_original_image_state, object_boxed_image_state, object_box_selector, object_box_json], | |
| ) | |
| object_box_json.change( | |
| fn=update_boxed_image_state, | |
| inputs=[object_original_image_state, object_box_json], | |
| outputs=[object_boxed_image_state], | |
| ) | |
| object_desc.change( | |
| fn=update_object_prompt_from_description, | |
| inputs=[object_desc], | |
| outputs=[object_prompt_box], | |
| ) | |
| object_run_btn.click( | |
| fn=set_run_button_busy, | |
| outputs=[object_run_btn], | |
| queue=False, | |
| ).then( | |
| fn=infer_object_moving, | |
| inputs=[object_original_image_state, object_boxed_image_state, object_desc, object_prompt_box, object_seed, object_randomize_seed, object_guidance_scale, object_num_inference_steps], | |
| outputs=[object_result], | |
| ).then( | |
| fn=set_run_button_idle, | |
| outputs=[object_run_btn], | |
| queue=False, | |
| ) | |
| with gr.Tab(TASK_MODE_ROTATION): | |
| rotation_original_image_state = gr.State(value=None) | |
| gr.Markdown("### ↻ Object Rotation") | |
| gr.Markdown("Use this mode to rotate a target object to a canonical viewpoint. Upload an image, specify which object should be rotated, drag the control below to choose the desired view, review or edit the generated editing prompt, and click Generate.") | |
| rotation_uploader = gr.Image( | |
| label="Upload Image", | |
| type="filepath", | |
| height=320, | |
| ) | |
| object_rotation_control = ObjectRotationControl3D( | |
| value=default_rotation_value(), | |
| three_available=three_available, | |
| ) | |
| rotation_object_desc = gr.Textbox( | |
| label="Object Description (describe the object you want to rotate)", | |
| placeholder="e.g., 'the chair in the center', 'the red car', 'the left shoe'", | |
| lines=2, | |
| info="Describe the object whose viewpoint you want to change. Use a precise visual description so the model can identify the correct object.", | |
| ) | |
| rotation_result = gr.Image(label="Output Image", height=500, format="png") | |
| rotation_run_btn = gr.Button(RUN_BUTTON_IDLE_TEXT, variant="primary", size="lg") | |
| rotation_prompt_box = gr.Textbox( | |
| label="Editing Prompt", | |
| value="", | |
| interactive=True, | |
| lines=3, | |
| placeholder="The editing prompt will be automatically generated from the controls above. You can further edit it before clicking Generate.", | |
| info="This prompt is auto-generated from the current controls. You can further modify it before generation.", | |
| ) | |
| default_steps = 30 | |
| with gr.Accordion("⚙️ Advanced Settings", open=False): | |
| if hide_advanced_options: | |
| rotation_seed = gr.Slider(visible=False, minimum=0, maximum=int(MAX_SEED), step=1, value=0) | |
| rotation_randomize_seed = gr.Checkbox(visible=False, value=True) | |
| rotation_guidance_scale = gr.Slider(visible=False, minimum=0.0, maximum=20.0, step=0.1, value=4.0) | |
| rotation_num_inference_steps = gr.Slider(visible=False, minimum=1, maximum=100, step=1, value=default_steps) | |
| else: | |
| rotation_seed = gr.Slider(label="Seed", minimum=0, maximum=int(MAX_SEED), step=1, value=0) | |
| rotation_randomize_seed = gr.Checkbox(label="Randomize Seed", value=True, interactive=True) | |
| rotation_guidance_scale = gr.Slider(label="Guidance Scale", minimum=0.0, maximum=20.0, step=0.1, value=4.0) | |
| rotation_num_inference_steps = gr.Slider(label="Inference Steps", minimum=1, maximum=100, step=1, value=default_steps) | |
| rotation_uploader.upload( | |
| fn=on_rotation_upload, | |
| inputs=[rotation_uploader], | |
| outputs=[rotation_original_image_state, object_rotation_control], | |
| ) | |
| rotation_uploader.change( | |
| fn=on_rotation_upload, | |
| inputs=[rotation_uploader], | |
| outputs=[rotation_original_image_state, object_rotation_control], | |
| ) | |
| rotation_uploader.clear( | |
| fn=reset_rotation_upload, | |
| outputs=[rotation_original_image_state, object_rotation_control], | |
| ) | |
| rotation_object_desc.change( | |
| fn=update_rotation_prompt_from_controls, | |
| inputs=[object_rotation_control, rotation_object_desc], | |
| outputs=[rotation_prompt_box], | |
| ) | |
| object_rotation_control.change( | |
| fn=update_rotation_prompt_from_controls, | |
| inputs=[object_rotation_control, rotation_object_desc], | |
| outputs=[rotation_prompt_box], | |
| ) | |
| rotation_run_btn.click( | |
| fn=set_run_button_busy, | |
| outputs=[rotation_run_btn], | |
| queue=False, | |
| ).then( | |
| fn=infer_object_rotation, | |
| inputs=[rotation_original_image_state, object_rotation_control, rotation_object_desc, rotation_prompt_box, rotation_seed, rotation_randomize_seed, rotation_guidance_scale, rotation_num_inference_steps], | |
| outputs=[rotation_result], | |
| ).then( | |
| fn=set_run_button_idle, | |
| outputs=[rotation_run_btn], | |
| queue=False, | |
| ) | |
| with gr.Tab(TASK_MODE_CAMERA): | |
| camera_original_image_state = gr.State(value=None) | |
| gr.Markdown("### 🎮 Camera Control") | |
| gr.Markdown("Use this mode to change only the camera viewpoint while keeping the scene itself unchanged. Upload an image, adjust yaw/pitch/zoom with the 3D controller or sliders, review or edit the generated editing prompt, and click Generate.") | |
| with gr.Row(): | |
| camera_uploader = gr.File( | |
| label="Upload Image", | |
| file_types=["image"], | |
| type="filepath", | |
| elem_id="se_upload_image_camera_file", | |
| ) | |
| camera_clear_btn = gr.Button("Clear Input", size="sm") | |
| gr.Markdown("### Input Image") | |
| camera_input_preview = gr.HTML( | |
| value=build_camera_preview_html(None), | |
| elem_id="se_upload_image_camera_preview", | |
| ) | |
| camera_prompt_box = gr.Textbox( | |
| label="Editing Prompt", | |
| value=build_camera_instruction(0.0, 0.0, 0, seed=0), | |
| interactive=True, | |
| lines=5, | |
| placeholder="The editing prompt will be automatically generated from the controls below. You can further edit it before clicking Generate.", | |
| info="This prompt is auto-generated from the current controls. You can further modify it before generation.", | |
| render=False, | |
| ) | |
| if three_available: | |
| gr.Markdown("*3D control usage: drag the colored handles — 🟢 Yaw, 🩷 Pitch, 🟠 Zoom.*") | |
| else: | |
| gr.Markdown("*Fallback: use the sliders below to control yaw, pitch, and zoom.*") | |
| camera_3d = CameraControl3D( | |
| value=default_camera_value(), | |
| elem_id="camera-3d-control", | |
| three_available=three_available, | |
| ) | |
| camera_result = gr.Image(label="Output Image", height=500, format="png") | |
| camera_run_btn = gr.Button(RUN_BUTTON_IDLE_TEXT, variant="primary", size="lg") | |
| gr.Markdown("### 🎚️ Slider Controls") | |
| yaw_slider = gr.Slider(label="Yaw", minimum=YAW_MIN, maximum=YAW_MAX, step=YAW_STEP, value=0.0) | |
| pitch_slider = gr.Slider(label="Pitch", minimum=PITCH_MIN, maximum=PITCH_MAX, step=PITCH_STEP, value=0.0) | |
| zoom_slider = gr.Slider(label="Zoom Delta", minimum=-1, maximum=1, step=1, value=0) | |
| display_line = gr.Textbox( | |
| label="Display", | |
| value=build_camera_display(0.0, 0.0, 0), | |
| interactive=False, | |
| lines=1, | |
| ) | |
| if examples_full: | |
| gr.HTML( | |
| """ | |
| <div style=" | |
| display:flex; | |
| align-items:center; | |
| gap:10px; | |
| margin: 6px 0 10px 0; | |
| font-size: 1.25rem; | |
| font-weight: 700; | |
| line-height: 1.2; | |
| color: var(--body-text-color); | |
| "> | |
| <span style="font-size: 1.05em; line-height: 1;">🖼️</span> | |
| <span>Camera Control Examples</span> | |
| </div> | |
| """ | |
| ) | |
| with gr.Row(): | |
| for idx, example_item in enumerate(examples_full): | |
| example_name = Path(str(example_item[0])).name | |
| example_btn = gr.Button(f"Example {idx + 1}: {example_name}", size="sm") | |
| example_btn.click( | |
| fn=lambda i=idx: load_camera_example_by_index(i), | |
| inputs=None, | |
| outputs=[camera_original_image_state, camera_input_preview, camera_3d, yaw_slider, pitch_slider, zoom_slider, display_line, camera_prompt_box], | |
| ) | |
| camera_prompt_box.render() | |
| default_steps = 30 | |
| with gr.Accordion("⚙️ Advanced Settings", open=False): | |
| if hide_advanced_options: | |
| seed = gr.Slider(visible=False, minimum=0, maximum=int(MAX_SEED), step=1, value=0) | |
| randomize_seed = gr.Checkbox(visible=False, value=True) | |
| guidance_scale = gr.Slider(visible=False, minimum=0.0, maximum=20.0, step=0.1, value=4.0) | |
| num_inference_steps = gr.Slider(visible=False, minimum=1, maximum=100, step=1, value=default_steps) | |
| else: | |
| seed = gr.Slider(label="Seed", minimum=0, maximum=int(MAX_SEED), step=1, value=0) | |
| randomize_seed = gr.Checkbox(label="Randomize Seed", value=True, interactive=True) | |
| guidance_scale = gr.Slider(label="Guidance Scale", minimum=0.0, maximum=20.0, step=0.1, value=4.0) | |
| num_inference_steps = gr.Slider(label="Inference Steps", minimum=1, maximum=100, step=1, value=default_steps) | |
| if auto_pe: | |
| gr.Markdown( | |
| f"**Auto PE rule:** when yaw is within [{AUTO_PE_DISABLE_YAW_MIN:g}, {AUTO_PE_DISABLE_YAW_MAX:g}], PE is disabled; otherwise PE is enabled automatically." | |
| ) | |
| for slider in [yaw_slider, pitch_slider, zoom_slider]: | |
| slider.change( | |
| fn=update_prompt_from_controls, | |
| inputs=[yaw_slider, pitch_slider, zoom_slider], | |
| outputs=[display_line, camera_prompt_box], | |
| ) | |
| if three_available: | |
| camera_3d.change( | |
| fn=sync_3d_to_sliders, | |
| inputs=[camera_3d], | |
| outputs=[yaw_slider, pitch_slider, zoom_slider, display_line, camera_prompt_box], | |
| ) | |
| for slider in [yaw_slider, pitch_slider, zoom_slider]: | |
| slider.release( | |
| fn=sync_sliders_to_3d, | |
| inputs=[yaw_slider, pitch_slider, zoom_slider], | |
| outputs=[camera_3d], | |
| ) | |
| camera_uploader.change( | |
| fn=on_camera_upload, | |
| inputs=[camera_uploader], | |
| outputs=[camera_original_image_state, camera_input_preview, camera_3d, yaw_slider, pitch_slider, zoom_slider, display_line, camera_prompt_box], | |
| ) | |
| camera_clear_btn.click( | |
| fn=reset_camera_upload, | |
| outputs=[camera_original_image_state, camera_input_preview, camera_3d, yaw_slider, pitch_slider, zoom_slider, display_line, camera_prompt_box], | |
| queue=False, | |
| ).then( | |
| fn=lambda: None, | |
| outputs=[camera_uploader], | |
| queue=False, | |
| ) | |
| camera_run_btn.click( | |
| fn=set_run_button_busy, | |
| outputs=[camera_run_btn], | |
| queue=False, | |
| ).then( | |
| fn=infer_camera_only, | |
| inputs=[ | |
| camera_original_image_state, | |
| yaw_slider, | |
| pitch_slider, | |
| zoom_slider, | |
| camera_prompt_box, | |
| seed, | |
| randomize_seed, | |
| guidance_scale, | |
| num_inference_steps | |
| ], | |
| outputs=[camera_result, seed, display_line], | |
| ).then( | |
| fn=set_run_button_idle, | |
| outputs=[camera_run_btn], | |
| queue=False, | |
| ) | |
| if app.model_load_mode in {"page_warmup", "warmup", "load_on_page", "onload"}: | |
| warmup_status = gr.Textbox(value="", visible=False, elem_id="se-model-warmup-status") | |
| demo.load( | |
| fn=warmup_model_on_page_load, | |
| inputs=None, | |
| outputs=[warmup_status], | |
| api_name=False, | |
| ) | |
| return demo, theme, page_css | |
| def main(): | |
| parser = argparse.ArgumentParser(description="JoyAi-Image-Edit (3D Camera Control)") | |
| parser.add_argument("--ckpt-root", type=str, required=True, help="Checkpoint root used by the new infer_runtime framework.") | |
| parser.add_argument("--config", type=str, default=None, help="Optional config path. Defaults to <ckpt-root>/infer_config.py.") | |
| parser.add_argument("--rewrite-prompt", action="store_true", help="Enable runtime LLM-based prompt rewriting via model.maybe_rewrite_prompt().") | |
| parser.add_argument("--rewrite-model", type=str, default="gpt-5", help="Rewrite model name passed into load_settings().") | |
| parser.add_argument("--hsdp-shard-dim", type=int, default=None, help="Optional hsdp_shard_dim override for multi-GPU FSDP inference.") | |
| parser.add_argument("--basesize", type=int, default=1024, help="Resize bucket base size passed to InferenceParams for image editing.") | |
| parser.add_argument( | |
| "--three-js-path", | |
| type=str, | |
| default="three.min.js", | |
| help="Optional local path to three.min.js. If omitted, will try ./three.min.js next to this .py file.", | |
| ) | |
| parser.add_argument( | |
| "--hide-advanced-options", | |
| action="store_true", | |
| help="Hide Advanced Settings in the Gradio UI.", | |
| ) | |
| parser.add_argument( | |
| "--auto-pe", | |
| action="store_true", | |
| help=f"Enable automatic PE: disable PE when yaw is within [{AUTO_PE_DISABLE_YAW_MIN:g}, {AUTO_PE_DISABLE_YAW_MAX:g}], otherwise enable PE.", | |
| ) | |
| parser.add_argument( | |
| "--default-save-dir", | |
| type=str, | |
| default="", | |
| help="Default directory for saving input image, prompt, and output image. Each run will create a timestamp-named subfolder.", | |
| ) | |
| parser.add_argument("--server-name", type=str, default="0.0.0.0", help="Gradio server name") | |
| parser.add_argument("--server-port", type=int, default=7773, help="Gradio server port") | |
| parser.add_argument("--share", action="store_true", help="Enable Gradio share") | |
| args = parser.parse_args() | |
| from modules.utils import maybe_init_distributed, clean_dist_env | |
| from modules.models.attention import describe_attention_backend | |
| three_js_file = resolve_local_three_js(args.three_js_path or None) | |
| three_available = three_js_file is not None | |
| inline_js = read_local_js_inline(three_js_file) | |
| if three_available: | |
| print(f"[Info] Using local three.js: {three_js_file}") | |
| else: | |
| print("[Info] No local three.min.js found. Falling back to slider-only mode.") | |
| print(f"[Info] Hide advanced options: {bool(args.hide_advanced_options)}") | |
| print(f"[Info] Camera precision: yaw={YAW_STEP:g}°, pitch={PITCH_STEP:g}°") | |
| print(f"[Info] Yaw range: [{YAW_MIN:g}, {YAW_MAX:g}]°") | |
| print(f"[Info] Pitch range: [{PITCH_MIN:g}, {PITCH_MAX:g}]°") | |
| print(f"[Info] Auto PE enabled: {bool(args.auto_pe)}") | |
| if args.auto_pe: | |
| print(f"[Info] Auto PE rule: disable PE for yaw in [{AUTO_PE_DISABLE_YAW_MIN:g}, {AUTO_PE_DISABLE_YAW_MAX:g}], enable PE otherwise.") | |
| print(f"[Info] Runtime prompt rewrite: {bool(args.rewrite_prompt)}") | |
| print(f"[Info] Rewrite model: {args.rewrite_model}") | |
| print(f"[Info] Basesize: {int(args.basesize)}") | |
| print(f"[Info] Default save dir: {args.default_save_dir or '(empty)'}") | |
| print(f"[Info] Default server name: {args.server_name}") | |
| print("[Info] Seed behavior: all three tasks use a randomized seed for each run") | |
| print("[Info] Output image format: PNG (compress_level=0)") | |
| dist_initialized = False | |
| try: | |
| device = resolve_device() | |
| dist_initialized = maybe_init_distributed() | |
| if is_rank0(): | |
| print(f"[Info] Chosen device: {device}") | |
| print(f"[Info] Attention backend: {describe_attention_backend()}") | |
| if args.hsdp_shard_dim is not None: | |
| print(f"[Info] Override hsdp_shard_dim: {args.hsdp_shard_dim}") | |
| app = EditorApp( | |
| ckpt_root=args.ckpt_root, | |
| config_path=args.config, | |
| rewrite_model=args.rewrite_model, | |
| hsdp_shard_dim=args.hsdp_shard_dim, | |
| enable_prompt_rewrite=bool(args.rewrite_prompt), | |
| basesize=int(args.basesize), | |
| device=device, | |
| ) | |
| examples_table, examples_full = build_demo_examples_from_config() | |
| demo, theme, page_css = create_demo( | |
| app, | |
| three_available=three_available, | |
| hide_advanced_options=bool(args.hide_advanced_options), | |
| examples_table=examples_table, | |
| examples_full=examples_full, | |
| auto_pe=bool(args.auto_pe), | |
| default_save_dir=str(args.default_save_dir or ""), | |
| ) | |
| launch_css = page_css + "\n" + ".fillable{max-width: 1400px !important}" | |
| demo.queue().launch( | |
| server_name=args.server_name, | |
| server_port=args.server_port, | |
| share=bool(args.share), | |
| head=inline_js, | |
| css=launch_css, | |
| ) | |
| finally: | |
| if dist_initialized: | |
| clean_dist_env() | |
| if __name__ == "__main__": | |
| main() | |