JoyAI-Image-Edit-Space / demo_release.py
stevengrove's picture
Update demo_release.py (#1)
3906621
Raw History Blame Contribute Delete
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:
@staticmethod
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
@torch.inference_mode()
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"))))
@spaces.GPU(duration=estimate_model_warmup_duration)
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}"
@spaces.GPU(duration=estimate_spatial_edit_duration)
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()