Philippe Potvin
Fri, 03 Jul 2026 00:43 - Use mounted model cache bucket
c8494c8
Raw
History Blame Contribute Delete
53.2 kB
#!/usr/bin/env python3
"""
Pro Realism Edit Studio - Enhanced Edition
=========================================
Advanced image editing and enhancement studio powered by:
- Qwen-Image-Edit-2511 with Phr00t's Rapid-AIO v23 accelerated transformer
- Nomos-family tiled upscaling
- Qwen-aware masked cleanup and final photographic detail
Author: Enhanced with Hugging Face CLI and image generation expertise
Version: 1.0.4
"""
import gradio as gr
import numpy as np
import random
import torch
import spaces
import os
import time
import tempfile
from pathlib import Path
def _configure_model_cache_environment():
bucket_root = os.environ.get("MODEL_BUCKET_DIR", "/models").strip()
fallback_root = os.environ.get("MODEL_CACHE_FALLBACK_DIR", "/data/models").strip()
candidates = [bucket_root, fallback_root, "/tmp/pro-realism-models"]
cache_root = None
for candidate in candidates:
if not candidate:
continue
path = Path(candidate)
if path.exists() or candidate.startswith(("/data", "/tmp")):
try:
path.mkdir(parents=True, exist_ok=True)
probe = path / ".write-test"
probe.write_text("ok", encoding="utf-8")
probe.unlink(missing_ok=True)
cache_root = path
break
except Exception:
continue
if cache_root is None:
return None
hf_home = cache_root / "huggingface"
hf_hub_cache = hf_home / "hub"
torch_home = cache_root / "torch"
ultralytics_cache = cache_root / "ultralytics"
for path in (hf_home, hf_hub_cache, torch_home, ultralytics_cache):
path.mkdir(parents=True, exist_ok=True)
os.environ.setdefault("HF_HOME", str(hf_home))
os.environ.setdefault("HF_HUB_CACHE", str(hf_hub_cache))
os.environ.setdefault("TRANSFORMERS_CACHE", str(hf_hub_cache))
os.environ.setdefault("DIFFUSERS_CACHE", str(hf_hub_cache))
os.environ.setdefault("TORCH_HOME", str(torch_home))
os.environ.setdefault("ULTRALYTICS_CACHE_DIR", str(ultralytics_cache))
os.environ.setdefault("YOLO_CONFIG_DIR", str(ultralytics_cache))
print(f"Model cache root: {cache_root}")
return cache_root
MODEL_CACHE_ROOT = _configure_model_cache_environment()
# Advanced imports
from accelerate import init_empty_weights
from collections import OrderedDict
from PIL import Image, ImageChops, ImageEnhance, ImageFilter, ImageOps
from diffusers.models import QwenImageTransformer2DModel as DiffusersQwenImageTransformer2DModel
from diffusers.models.model_loading_utils import load_model_dict_into_meta
from huggingface_hub import hf_hub_download, HfApi, login, whoami
from qwenimage.pipeline_qwenimage_edit_plus import QwenImageEditPlusPipeline
from qwenimage.transformer_qwenimage import QwenImageTransformer2DModel
from qwenimage.qwen_fa3_processor import QwenDoubleStreamAttnProcessorFA3
from safetensors import safe_open
from gradio_client import Client, handle_file
# ============================================================================
# CONFIGURATION - Model IDs and Parameters
# ============================================================================
# Base model configuration
BASE_MODEL_ID = "Qwen/Qwen-Image-Edit-2511"
APP_VERSION = "1.0.4"
PHR00T_REPO_ID = os.environ.get("PHR00T_REPO_ID", "Phr00t/Qwen-Image-Edit-Rapid-AIO").strip()
RAPID_TRANSFORMER_FILENAME = os.environ.get(
"RAPID_TRANSFORMER_FILENAME",
"v23/Qwen-Rapid-AIO-NSFW-v23.safetensors",
).strip()
PHR00T_TRANSFORMER_PREFIX = "model.diffusion_model."
VIDEO_SPACE_ID = os.environ.get("VIDEO_SPACE_ID", "").strip()
# Qwen-optimized enhancement configuration
UPSCALER_MODEL_ID = os.environ.get("UPSCALER_MODEL_ID", "Phips/4xNomos8k_atd_jpg").strip()
UPSCALER_MODEL_FILENAME = os.environ.get("UPSCALER_MODEL_FILENAME", "4xNomos8k_atd_jpg.safetensors").strip()
UPSCALER_TILE_SIZE = int(os.environ.get("UPSCALER_TILE_SIZE", "512"))
UPSCALER_TILE_OVERLAP = int(os.environ.get("UPSCALER_TILE_OVERLAP", "64"))
ENHANCE_MAX_INPUT_EDGE = int(os.environ.get("ENHANCE_MAX_INPUT_EDGE", "2048"))
ENHANCE_QWEN_DOWNSCALE_TRIGGER_EDGE = int(os.environ.get("ENHANCE_QWEN_DOWNSCALE_TRIGGER_EDGE", "1536"))
ENHANCE_QWEN_DOWNSCALE_FACTOR = float(os.environ.get("ENHANCE_QWEN_DOWNSCALE_FACTOR", "0.75"))
ENHANCE_GRAIN_STRENGTH = float(os.environ.get("ENHANCE_GRAIN_STRENGTH", "0.010"))
ENHANCE_DETAILER_ENABLED = os.environ.get("ENHANCE_DETAILER_ENABLED", "true").lower() == "true"
ENHANCE_DETAILER_MAX_REGIONS = int(os.environ.get("ENHANCE_DETAILER_MAX_REGIONS", "8"))
ENHANCE_DETAILER_MIN_REGION_AREA = float(os.environ.get("ENHANCE_DETAILER_MIN_REGION_AREA", "0.003"))
ENHANCE_YOLO_SAM_ENABLED = os.environ.get("ENHANCE_YOLO_SAM_ENABLED", "true").lower() == "true"
ENHANCE_YOLO_MODEL = os.environ.get("ENHANCE_YOLO_MODEL", "yolo26n-seg.pt").strip()
ENHANCE_SAM_MODEL = os.environ.get("ENHANCE_SAM_MODEL", "sam_l.pt").strip()
ENHANCE_YOLO_CONF = float(os.environ.get("ENHANCE_YOLO_CONF", "0.18"))
ENHANCE_DA3_PRESERVE_ENABLED = os.environ.get("ENHANCE_DA3_PRESERVE_ENABLED", "false").lower() == "true"
ENHANCE_DA3_MODEL_ID = os.environ.get("ENHANCE_DA3_MODEL_ID", "depth-anything/DA3-BASE").strip()
# Advanced Detail Enhancement Configuration
DETAIL_ENHANCEMENT_ENABLED = os.environ.get("DETAIL_ENHANCEMENT_ENABLED", "true").lower() == "true"
SMART_SHARPENING_STRENGTH = float(os.environ.get("SMART_SHARPENING_STRENGTH", "1.15"))
# ============================================================================
# ENHANCEMENT MODES
# ============================================================================
ENHANCE_MODE_OFF = "Off"
ENHANCE_MODE_UPSCALE = "Upscale"
ENHANCE_MODE_CLEAN = "Clean"
ENHANCE_MODE_MAX_DETAIL = "Max Detail"
ENHANCE_MODE_CHOICES = [ENHANCE_MODE_OFF, ENHANCE_MODE_UPSCALE, ENHANCE_MODE_CLEAN, ENHANCE_MODE_MAX_DETAIL]
# ============================================================================
# GLOBAL MODEL CACHE
# ============================================================================
_upscaler_model = None
_detail_enhancement_model = None
_yolo_model = None
_sam_model = None
_da3_model = None
# ============================================================================
# HUGGING FACE CLI EXPERT FUNCTIONS
# ============================================================================
def download_model_with_retry(repo_id, filename, max_retries=3):
"""Download model with retry logic and error handling"""
for attempt in range(max_retries):
try:
return hf_hub_download(
repo_id=repo_id,
filename=filename,
cache_dir=os.environ.get("HF_HUB_CACHE"),
)
except Exception as e:
if attempt == max_retries - 1:
raise RuntimeError(f"Failed to download {filename} from {repo_id} after {max_retries} attempts: {e}")
time.sleep(2 ** attempt)
return None
# ============================================================================
# VIDEO GENERATION (Preserved from original)
# ============================================================================
def turn_into_video(input_image, output_images, prompt, progress=gr.Progress(track_tqdm=True)):
"""Convert image edit into video transition"""
if not VIDEO_SPACE_ID:
raise gr.Error("Video generation is not configured for this Space.")
if not input_image or not output_images:
raise gr.Error("Please generate an output image first.")
progress(0.02, desc="Preparing images...")
def extract_pil(img_entry):
if isinstance(img_entry, tuple) and isinstance(img_entry[0], Image.Image):
return img_entry[0]
elif isinstance(img_entry, Image.Image):
return img_entry
elif isinstance(img_entry, str):
return Image.open(img_entry)
else:
raise gr.Error(f"Unsupported image format: {type(img_entry)}")
start_img = extract_pil(input_image)
end_img = extract_pil(output_images[0])
progress(0.10, desc="Saving temp files...")
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as tmp_start, \
tempfile.NamedTemporaryFile(suffix=".png", delete=False) as tmp_end:
start_img.save(tmp_start.name)
end_img.save(tmp_end.name)
progress(0.20, desc="Connecting to video Space...")
client = Client(VIDEO_SPACE_ID)
progress(0.35, desc="Generating video...")
video_path, seed = client.predict(
start_image_pil=handle_file(tmp_start.name),
end_image_pil=handle_file(tmp_end.name),
prompt=prompt or "smooth cinematic transition",
api_name="/generate_video"
)
progress(0.95, desc="Finalizing...")
return video_path['video']
# ============================================================================
# HISTORY MANAGEMENT (Enhanced)
# ============================================================================
def update_history(new_images, history):
"""Updates the history gallery with the new images."""
time.sleep(0.3)
if history is None:
history = []
if new_images is not None and len(new_images) > 0:
if not isinstance(history, list):
history = list(history) if history else []
for img in new_images:
history.insert(0, img)
history = history[:50]
return history
def use_history_as_input(evt: gr.SelectData):
"""Sets the selected history image into the Image 1 slot."""
if evt.value is not None:
return gr.update(value=evt.value)
return gr.update()
# ============================================================================
# MODEL LOADING (Enhanced with better error handling)
# ============================================================================
dtype = torch.bfloat16
device = "cuda" if torch.cuda.is_available() else "cpu"
def load_phr00t_rapid_transformer(torch_dtype):
"""Load Phr00t's Rapid-AIO v23 transformer with enhanced error handling"""
checkpoint_path = download_model_with_retry(PHR00T_REPO_ID, RAPID_TRANSFORMER_FILENAME)
try:
config = DiffusersQwenImageTransformer2DModel.load_config(
BASE_MODEL_ID,
subfolder="transformer",
cache_dir=os.environ.get("HF_HUB_CACHE"),
)
except Exception as e:
raise RuntimeError(f"Failed to load config for {BASE_MODEL_ID}: {e}")
with init_empty_weights():
transformer = DiffusersQwenImageTransformer2DModel.from_config(config)
expected_keys = set(transformer.state_dict().keys())
state_dict = OrderedDict()
try:
with safe_open(checkpoint_path, framework="pt", device="cpu") as checkpoint:
for key in checkpoint.keys():
if not key.startswith(PHR00T_TRANSFORMER_PREFIX):
continue
mapped_key = key.removeprefix(PHR00T_TRANSFORMER_PREFIX)
if mapped_key in expected_keys:
state_dict[mapped_key] = checkpoint.get_tensor(key)
except Exception as e:
raise RuntimeError(f"Failed to load checkpoint from {checkpoint_path}: {e}")
missing_keys = sorted(expected_keys.difference(state_dict.keys()))
if missing_keys:
sample = ", ".join(missing_keys[:20])
raise RuntimeError(
f"Phr00t Rapid-AIO transformer checkpoint is missing {len(missing_keys)} "
f"required diffusers keys after prefix conversion. First missing keys: {sample}"
)
try:
load_model_dict_into_meta(transformer, state_dict, dtype=torch_dtype)
except Exception as e:
raise RuntimeError(f"Failed to load state dict into meta: {e}")
meta_parameters = [name for name, parameter in transformer.named_parameters() if parameter.is_meta]
if meta_parameters:
sample = ", ".join(meta_parameters[:20])
raise RuntimeError(
f"Phr00t Rapid-AIO transformer still has {len(meta_parameters)} meta parameters "
f"after loading. First meta parameters: {sample}"
)
transformer.eval()
return transformer
# Load main pipeline
try:
pipe = QwenImageEditPlusPipeline.from_pretrained(
BASE_MODEL_ID,
transformer=load_phr00t_rapid_transformer(dtype),
torch_dtype=dtype,
cache_dir=os.environ.get("HF_HUB_CACHE"),
).to(device)
print("Successfully loaded Qwen-Image-Edit-2511 with Rapid-AIO v23 transformer")
except Exception as e:
print(f"Failed to load main pipeline: {e}")
raise
# Apply optimizations
pipe.transformer.__class__ = QwenImageTransformer2DModel
pipe.transformer.set_attn_processor(QwenDoubleStreamAttnProcessorFA3())
print("Applied FA3 attention processor optimization")
# ============================================================================
# QWEN-OPTIMIZED UPSCALER
# ============================================================================
def load_upscaler_model():
"""Load a Spandrel upscaler, defaulting to the Nomos model used by this Space."""
global _upscaler_model
if _upscaler_model is not None:
return _upscaler_model
try:
import spandrel
import spandrel_extra_arches
spandrel_extra_arches.install()
except ImportError as exc:
raise gr.Error("Enhance mode requires spandrel and spandrel_extra_arches to be installed. "
"Install with: pip install spandrel spandrel_extra_arches") from exc
candidates = [
(UPSCALER_MODEL_ID, UPSCALER_MODEL_FILENAME),
("Phips/4xNomos8k_atd_jpg", "4xNomos8k_atd_jpg.safetensors"),
]
seen = set()
errors = []
for repo_id, filename in candidates:
key = (repo_id, filename)
if key in seen:
continue
seen.add(key)
try:
model_path = download_model_with_retry(repo_id, filename)
model = spandrel.ModelLoader().load_from_file(model_path)
model.eval().to(device)
_upscaler_model = model
print(f"Successfully loaded upscaler: {repo_id}/{filename}")
return _upscaler_model
except Exception as exc:
errors.append(f"{repo_id}/{filename}: {exc}")
print(f"Failed to load upscaler {repo_id}/{filename}: {exc}")
raise gr.Error(f"Failed to load all upscaler models: {' | '.join(errors)}")
def image_to_tensor(image):
"""Convert PIL Image to tensor"""
array = np.asarray(image.convert("RGB")).astype(np.float32) / 255.0
tensor = torch.from_numpy(array).permute(2, 0, 1).unsqueeze(0)
return tensor.to(device)
def tensor_to_image(tensor):
"""Convert tensor to PIL Image"""
array = tensor.squeeze(0).detach().float().cpu().clamp(0, 1).permute(1, 2, 0).numpy()
return Image.fromarray((array * 255.0).round().astype(np.uint8), mode="RGB")
def validate_enhance_input_size(image):
"""Validate image size for enhancement"""
max_edge = max(image.size)
if max_edge > ENHANCE_MAX_INPUT_EDGE:
raise gr.Error(
f"Enhance mode accepts images up to {ENHANCE_MAX_INPUT_EDGE}px on the longest edge. "
f"Current image is {image.width}x{image.height}. "
f"Consider resizing your image first."
)
def load_yolo_model():
global _yolo_model
if _yolo_model is not None:
return _yolo_model
try:
from ultralytics import YOLO
_yolo_model = YOLO(ENHANCE_YOLO_MODEL)
return _yolo_model
except Exception as exc:
print(f"YOLO detailer unavailable: {exc}")
return None
def load_sam_model():
global _sam_model
if _sam_model is not None:
return _sam_model
try:
from ultralytics import SAM
_sam_model = SAM(ENHANCE_SAM_MODEL)
return _sam_model
except Exception as exc:
print(f"SAM detailer unavailable: {exc}")
return None
def load_da3_model():
global _da3_model
if _da3_model is not None:
return _da3_model
if not ENHANCE_DA3_PRESERVE_ENABLED:
return None
try:
from depth_anything_3.api import DepthAnything3
_da3_model = DepthAnything3.from_pretrained(ENHANCE_DA3_MODEL_ID).to(device)
return _da3_model
except Exception as exc:
print(f"DA-3 preservation unavailable: {exc}")
return None
def _tile_weight(height, width, overlap_y, overlap_x, touches_top, touches_bottom, touches_left, touches_right, dtype, device):
weight = torch.ones((1, 1, height, width), dtype=dtype, device=device)
if overlap_y > 1 and not touches_top:
ramp = torch.linspace(0.0, 1.0, overlap_y, dtype=dtype, device=device).view(1, 1, overlap_y, 1)
weight[:, :, :overlap_y, :] *= ramp
if overlap_y > 1 and not touches_bottom:
ramp = torch.linspace(1.0, 0.0, overlap_y, dtype=dtype, device=device).view(1, 1, overlap_y, 1)
weight[:, :, -overlap_y:, :] *= ramp
if overlap_x > 1 and not touches_left:
ramp = torch.linspace(0.0, 1.0, overlap_x, dtype=dtype, device=device).view(1, 1, 1, overlap_x)
weight[:, :, :, :overlap_x] *= ramp
if overlap_x > 1 and not touches_right:
ramp = torch.linspace(1.0, 0.0, overlap_x, dtype=dtype, device=device).view(1, 1, 1, overlap_x)
weight[:, :, :, -overlap_x:] *= ramp
return weight
def tile_upscale(image, scale=4):
"""Tiled Spandrel upscaling with feathered overlaps to avoid visible seams."""
validate_enhance_input_size(image)
model = load_upscaler_model()
tensor = image_to_tensor(image)
_, _, height, width = tensor.shape
base_tile_size = UPSCALER_TILE_SIZE
optimal_tile_size = min(base_tile_size, max(height, width) // 2)
tile_size = max(64, optimal_tile_size)
overlap = max(0, min(UPSCALER_TILE_OVERLAP, tile_size // 2))
step = max(1, tile_size - overlap)
y_positions = list(range(0, height, step))
if y_positions[-1] + tile_size < height:
y_positions.append(max(0, height - tile_size))
x_positions = list(range(0, width, step))
if x_positions[-1] + tile_size < width:
x_positions.append(max(0, width - tile_size))
y_positions = sorted(set(y_positions))
x_positions = sorted(set(x_positions))
output = None
weights = None
with torch.inference_mode():
for y in y_positions:
for x in x_positions:
y1 = min(y + tile_size, height)
x1 = min(x + tile_size, width)
tile = tensor[:, :, y:y1, x:x1]
upscaled_tile = model(tile)
if isinstance(upscaled_tile, (tuple, list)):
upscaled_tile = upscaled_tile[0]
upscaled_tile = upscaled_tile.clamp(0, 1)
scale_y = upscaled_tile.shape[-2] // tile.shape[-2]
scale_x = upscaled_tile.shape[-1] // tile.shape[-1]
if output is None:
output = torch.zeros(
(1, 3, height * scale_y, width * scale_x),
dtype=upscaled_tile.dtype,
device=upscaled_tile.device,
)
weights = torch.zeros_like(output)
oy0, oy1 = y * scale_y, y1 * scale_y
ox0, ox1 = x * scale_x, x1 * scale_x
blend = _tile_weight(
upscaled_tile.shape[-2],
upscaled_tile.shape[-1],
min(overlap * scale_y, max(1, upscaled_tile.shape[-2] // 2)),
min(overlap * scale_x, max(1, upscaled_tile.shape[-1] // 2)),
y == 0,
y1 == height,
x == 0,
x1 == width,
upscaled_tile.dtype,
upscaled_tile.device,
)
output[:, :, oy0:oy1, ox0:ox1] += upscaled_tile * blend
weights[:, :, oy0:oy1, ox0:ox1] += blend
output = output / weights.clamp_min(1)
return tensor_to_image(output)
def advanced_tile_upscale(image, scale=4):
return tile_upscale(image, scale=scale)
# ============================================================================
# ENHANCED DETAILER (Multi-stage processing)
# ============================================================================
def smart_sharpen(image, strength=1.15):
"""
Smart sharpening with edge detection to avoid oversharpening smooth areas
"""
if strength <= 0:
return image
img_array = np.array(image.convert("RGB"))
if strength > 1.0:
enhanced = ImageEnhance.Sharpness(image).enhance(strength)
gray = image.convert("L")
edges = gray.filter(ImageFilter.FIND_EDGES)
edge_mask = edges.filter(ImageFilter.GaussianBlur(radius=1))
edge_mask = edge_mask.point(lambda x: min(x * 0.3, 255))
sharpened_array = np.array(enhanced)
original_array = img_array
edge_array = np.array(edge_mask).astype(float) / 255.0
for c in range(3):
sharpened_array[:, :, c] = (
edge_array * sharpened_array[:, :, c] +
(1 - edge_array) * original_array[:, :, c]
)
image = Image.fromarray(np.clip(sharpened_array, 0, 255).astype(np.uint8))
return image
def add_ultra_detail(image, strength=0.8):
"""
Add ultra-fine details using high-frequency enhancement
"""
if strength <= 0:
return image
original = image.convert("RGB")
blurred = original.filter(ImageFilter.GaussianBlur(radius=2))
high_freq = ImageChops.subtract(original, blurred)
high_freq_enhanced = ImageEnhance.Contrast(high_freq).enhance(1.0 + strength)
result = ImageChops.add(original, high_freq_enhanced)
return result
def apply_high_frequency_details(image, amount=0.6):
"""
Apply high-frequency detail enhancement for crisp textures
"""
if amount <= 0:
return image
scales = [1, 2, 4]
result = image.convert("RGB")
for scale in scales:
blurred = result.filter(ImageFilter.GaussianBlur(radius=scale))
high_freq = ImageChops.subtract(result, blurred)
enhanced_hf = ImageEnhance.Contrast(high_freq).enhance(1.0 + amount * 0.3)
result = ImageChops.add(result, enhanced_hf)
return result
# ============================================================================
# LEGACY CLEANUP HELPERS
# ============================================================================
def remove_artifacts(image):
"""Remove compression artifacts and noise"""
denoised = image.filter(ImageFilter.MedianFilter(size=3))
smoothed = denoised.filter(ImageFilter.GaussianBlur(radius=0.5))
result = Image.blend(image, smoothed, alpha=0.3)
return result
def enhanced_skin_repair(image):
"""Enhanced skin repair with better color detection and blending"""
base = image.convert("RGB")
ycbcr = np.asarray(base.convert("YCbCr"))
y, cb, cr = ycbcr[:, :, 0], ycbcr[:, :, 1], ycbcr[:, :, 2]
skin_mask = (
(cr > 130) & (cr < 170) &
(cb > 70) & (cb < 140) &
(y > 80)
).astype(np.uint8) * 255
try:
import cv2
kernel = np.ones((5, 5), np.uint8)
skin_mask = cv2.morphologyEx(skin_mask, cv2.MORPH_OPEN, kernel)
skin_mask = cv2.morphologyEx(skin_mask, cv2.MORPH_CLOSE, kernel)
skin_mask = cv2.GaussianBlur(skin_mask, (7, 7), 0)
except ImportError:
from scipy import ndimage
skin_mask = ndimage.binary_opening(skin_mask > 128, structure=np.ones((3, 3))).astype(np.uint8) * 255
skin_mask = ndimage.gaussian_filter(skin_mask, sigma=3)
mask_image = Image.fromarray(skin_mask, mode="L")
repaired = base.filter(ImageFilter.MedianFilter(size=3))
repaired = repaired.filter(ImageFilter.GaussianBlur(radius=0.4))
non_skin = ImageOps.invert(mask_image)
sharpened = ImageEnhance.Sharpness(base).enhance(1.15)
blended = Image.composite(repaired, sharpened, mask_image)
result = ImageEnhance.Sharpness(blended).enhance(1.05)
return result
# Original skin repair functions (preserved for compatibility)
def skin_repair_mask(image):
ycbcr = np.asarray(image.convert("YCbCr"))
cb = ycbcr[:, :, 1]
cr = ycbcr[:, :, 2]
mask = (
(cr >= 135)
& (cr <= 180)
& (cb >= 75)
& (cb <= 135)
).astype(np.uint8) * 255
mask_image = Image.fromarray(mask, mode="L")
return mask_image.filter(ImageFilter.GaussianBlur(radius=1.2))
def repair_skin_texture(image):
base = image.convert("RGB")
mask = skin_repair_mask(base)
repaired = base.filter(ImageFilter.MedianFilter(size=3)).filter(ImageFilter.GaussianBlur(radius=0.35))
blended = Image.composite(repaired, base, mask)
return ImageEnhance.Sharpness(blended).enhance(1.08)
def add_film_grain(image, seed):
base = image.convert("RGB")
array = np.asarray(base).astype(np.float32)
rng = np.random.default_rng(seed)
grain = rng.normal(0.0, 255.0 * ENHANCE_GRAIN_STRENGTH, size=(array.shape[0], array.shape[1], 1))
array = np.clip(array + grain, 0, 255)
return Image.fromarray(array.astype(np.uint8), mode="RGB")
def _resize_to_even_dimensions(width, height):
return max(2, width - (width % 2)), max(2, height - (height % 2))
def qwen_artifact_precondition(image):
"""
Qwen edit outputs can show halftone/plastic texture at larger dimensions.
A small Lanczos downsample before the SR pass suppresses that pattern while
preserving prompt structure for the upscaler.
"""
base = image.convert("RGB")
long_edge = max(base.size)
if long_edge <= ENHANCE_QWEN_DOWNSCALE_TRIGGER_EDGE:
return base
factor = min(0.95, max(0.50, ENHANCE_QWEN_DOWNSCALE_FACTOR))
new_width, new_height = _resize_to_even_dimensions(
int(round(base.width * factor)),
int(round(base.height * factor)),
)
if new_width >= base.width or new_height >= base.height:
return base
return base.resize((new_width, new_height), Image.Resampling.LANCZOS)
def qwen_skin_mask(image):
base = image.convert("RGB")
ycbcr = np.asarray(base.convert("YCbCr"))
y = ycbcr[:, :, 0]
cb = ycbcr[:, :, 1]
cr = ycbcr[:, :, 2]
mask = (
(y > 45)
& (cb >= 72)
& (cb <= 145)
& (cr >= 128)
& (cr <= 182)
).astype(np.uint8) * 255
try:
import cv2
kernel = np.ones((3, 3), np.uint8)
mask = cv2.morphologyEx(mask, cv2.MORPH_OPEN, kernel)
mask = cv2.morphologyEx(mask, cv2.MORPH_CLOSE, kernel)
mask = cv2.GaussianBlur(mask, (0, 0), 1.6)
except ImportError:
mask_image = Image.fromarray(mask, mode="L")
mask_image = mask_image.filter(ImageFilter.MinFilter(size=3))
mask_image = mask_image.filter(ImageFilter.MaxFilter(size=5))
mask_image = mask_image.filter(ImageFilter.GaussianBlur(radius=1.4))
mask = np.asarray(mask_image)
return Image.fromarray(mask.astype(np.uint8), mode="L")
def qwen_defect_mask(image):
base = image.convert("RGB")
skin = np.asarray(qwen_skin_mask(base)).astype(np.float32) / 255.0
gray = np.asarray(base.convert("L")).astype(np.float32)
local = np.asarray(base.convert("L").filter(ImageFilter.MedianFilter(size=5))).astype(np.float32)
pits = np.maximum(local - gray, 0.0)
smears = np.abs(gray - local)
mask = ((pits > 14.0) | (smears > 22.0)) & (skin > 0.12)
mask = (mask.astype(np.uint8) * 255)
mask_image = Image.fromarray(mask, mode="L")
mask_image = mask_image.filter(ImageFilter.MaxFilter(size=3))
return mask_image.filter(ImageFilter.GaussianBlur(radius=1.1))
def qwen_hair_mask(image):
base = image.convert("RGB")
rgb = np.asarray(base).astype(np.float32)
gray_image = base.convert("L")
gray = np.asarray(gray_image).astype(np.float32)
skin = np.asarray(qwen_skin_mask(base)).astype(np.float32) / 255.0
edges = np.asarray(gray_image.filter(ImageFilter.FIND_EDGES).filter(ImageFilter.GaussianBlur(radius=0.7))).astype(np.float32)
chroma = rgb.max(axis=2) - rgb.min(axis=2)
dark_strands = (gray < 122) & (edges > 8) & (skin < 0.45)
light_strands = (gray < 235) & (edges > 18) & (chroma > 8) & (skin < 0.28)
mask = (dark_strands | light_strands).astype(np.uint8) * 255
try:
import cv2
kernel = np.ones((3, 3), np.uint8)
mask = cv2.morphologyEx(mask, cv2.MORPH_CLOSE, kernel)
mask = cv2.GaussianBlur(mask, (0, 0), 1.0)
except ImportError:
mask_image = Image.fromarray(mask, mode="L")
mask_image = mask_image.filter(ImageFilter.MaxFilter(size=3))
mask_image = mask_image.filter(ImageFilter.GaussianBlur(radius=1.0))
mask = np.asarray(mask_image)
return Image.fromarray(mask.astype(np.uint8), mode="L")
def qwen_face_feature_mask(image):
base = image.convert("RGB")
skin = np.asarray(qwen_skin_mask(base)).astype(np.float32) / 255.0
gray_image = base.convert("L")
edges = np.asarray(gray_image.filter(ImageFilter.FIND_EDGES)).astype(np.float32)
features = ((skin > 0.08) & (edges > 10)).astype(np.uint8) * 255
defects = np.asarray(qwen_defect_mask(base)).astype(np.uint8)
mask = np.maximum(features, defects)
mask_image = Image.fromarray(mask, mode="L")
mask_image = mask_image.filter(ImageFilter.MaxFilter(size=5))
return mask_image.filter(ImageFilter.GaussianBlur(radius=1.6))
def _results_to_person_boxes(results, image_size):
boxes = []
for result in results or []:
if getattr(result, "boxes", None) is None:
continue
xyxy = result.boxes.xyxy.detach().cpu().numpy() if result.boxes.xyxy is not None else []
cls = result.boxes.cls.detach().cpu().numpy() if result.boxes.cls is not None else []
conf = result.boxes.conf.detach().cpu().numpy() if result.boxes.conf is not None else []
for idx, box in enumerate(xyxy):
class_id = int(cls[idx]) if idx < len(cls) else -1
score = float(conf[idx]) if idx < len(conf) else 1.0
if class_id == 0 and score >= ENHANCE_YOLO_CONF:
x0, y0, x1, y1 = [int(round(value)) for value in box]
x0 = max(0, min(image_size[0] - 1, x0))
x1 = max(1, min(image_size[0], x1))
y0 = max(0, min(image_size[1] - 1, y0))
y1 = max(1, min(image_size[1], y1))
if x1 > x0 and y1 > y0:
boxes.append((x0, y0, x1, y1))
boxes.sort(key=lambda item: (item[2] - item[0]) * (item[3] - item[1]), reverse=True)
return boxes[:ENHANCE_DETAILER_MAX_REGIONS]
def yolo_person_boxes(image):
if not ENHANCE_YOLO_SAM_ENABLED:
return []
model = load_yolo_model()
if model is None:
return []
try:
results = model.predict(
source=np.asarray(image.convert("RGB")),
classes=[0],
conf=ENHANCE_YOLO_CONF,
verbose=False,
)
return _results_to_person_boxes(results, image.size)
except Exception as exc:
print(f"YOLO person detection failed: {exc}")
return []
def sam_mask_from_boxes(image, boxes):
if not boxes or not ENHANCE_YOLO_SAM_ENABLED:
return Image.new("L", image.size, 0)
model = load_sam_model()
if model is None:
return Image.new("L", image.size, 0)
try:
union = np.zeros((image.height, image.width), dtype=np.uint8)
for box in boxes:
results = model(np.asarray(image.convert("RGB")), bboxes=list(box), verbose=False)
for result in results or []:
masks = getattr(result, "masks", None)
if masks is None or masks.data is None:
continue
data = masks.data.detach().float().cpu().numpy()
for mask in data:
mask_image = Image.fromarray((mask > 0.5).astype(np.uint8) * 255, mode="L")
mask_image = mask_image.resize(image.size, Image.Resampling.BILINEAR)
union = np.maximum(union, np.asarray(mask_image, dtype=np.uint8))
return Image.fromarray(union, mode="L").filter(ImageFilter.GaussianBlur(radius=1.2))
except Exception as exc:
print(f"SAM box segmentation failed: {exc}")
return Image.new("L", image.size, 0)
def da3_preservation_mask(image, person_mask):
if not ENHANCE_DA3_PRESERVE_ENABLED:
return person_mask
model = load_da3_model()
if model is None:
return person_mask
try:
with tempfile.TemporaryDirectory() as tmpdir:
prediction = model.inference(
[np.asarray(image.convert("RGB"))],
export_dir=tmpdir,
export_format="npz",
)
if isinstance(prediction, dict):
depth = prediction.get("depth")
if depth is None:
depth = prediction.get("depths")
else:
depth = getattr(prediction, "depth", None)
if depth is None:
depth = getattr(prediction, "depths", None)
if depth is None:
return person_mask
depth_array = np.asarray(depth[0] if isinstance(depth, (list, tuple)) else depth).astype(np.float32)
if depth_array.ndim > 2:
depth_array = depth_array.squeeze()
depth_array -= depth_array.min()
depth_array /= max(float(depth_array.max()), 1e-6)
depth_image = Image.fromarray((depth_array * 255).astype(np.uint8), mode="L").resize(image.size, Image.Resampling.BILINEAR)
foreground = depth_image.point(lambda value: 255 if value >= 32 else 0).filter(ImageFilter.GaussianBlur(radius=1.4))
return ImageChops.multiply(person_mask.convert("L"), foreground)
except Exception as exc:
print(f"DA-3 preservation mask failed: {exc}")
return person_mask
def yolo_sam_person_mask_and_boxes(image):
boxes = yolo_person_boxes(image)
person_mask = sam_mask_from_boxes(image, boxes)
if person_mask.getbbox() is None:
return person_mask, boxes
return da3_preservation_mask(image, person_mask), boxes
def _mask_to_boxes(mask_image, max_regions=ENHANCE_DETAILER_MAX_REGIONS):
mask = np.asarray(mask_image.convert("L"))
binary = (mask > 24).astype(np.uint8)
min_area = max(32, int(binary.shape[0] * binary.shape[1] * ENHANCE_DETAILER_MIN_REGION_AREA))
boxes = []
try:
import cv2
count, labels, stats, _ = cv2.connectedComponentsWithStats(binary, connectivity=8)
for label in range(1, count):
x, y, width, height, area = stats[label]
if area >= min_area:
boxes.append((int(x), int(y), int(x + width), int(y + height), int(area)))
except ImportError:
bbox = mask_image.point(lambda value: 255 if value > 24 else 0).getbbox()
if bbox:
x0, y0, x1, y1 = bbox
boxes.append((x0, y0, x1, y1, (x1 - x0) * (y1 - y0)))
boxes.sort(key=lambda item: item[4], reverse=True)
return [box[:4] for box in boxes[:max_regions]]
def _expand_box(box, image_size, pad_ratio=0.18, min_size=192):
x0, y0, x1, y1 = box
width = x1 - x0
height = y1 - y0
pad = int(max(width, height) * pad_ratio)
if width < min_size:
extra = (min_size - width) // 2
x0 -= extra
x1 += extra
if height < min_size:
extra = (min_size - height) // 2
y0 -= extra
y1 += extra
return (
max(0, x0 - pad),
max(0, y0 - pad),
min(image_size[0], x1 + pad),
min(image_size[1], y1 + pad),
)
def _local_detail_crop(crop, mask_crop, strength=0.35, hair=False):
base = crop.convert("RGB")
mask = mask_crop.convert("L").filter(ImageFilter.GaussianBlur(radius=1.8))
strength = max(0.0, min(1.0, strength))
if hair:
detailed = base.filter(ImageFilter.UnsharpMask(radius=0.55, percent=int(130 + 120 * strength), threshold=2))
detailed = ImageEnhance.Contrast(detailed).enhance(1.0 + 0.08 * strength)
edge_mask = base.convert("L").filter(ImageFilter.FIND_EDGES).filter(ImageFilter.GaussianBlur(radius=0.7))
mask = ImageChops.multiply(mask, edge_mask.point(lambda value: min(255, int(value * 2.2))))
else:
repaired = masked_texture_repair(base, strength=0.35 + 0.30 * strength)
detailed = repaired.filter(ImageFilter.UnsharpMask(radius=0.75, percent=int(80 + 90 * strength), threshold=3))
detailed = ImageEnhance.Contrast(detailed).enhance(1.0 + 0.045 * strength)
mask = mask.point(lambda value: int(value * strength))
return Image.composite(detailed, base, mask)
def qwen_tiled_detailer_pass(image, strength=0.35, include_hair=True):
if not ENHANCE_DETAILER_ENABLED or strength <= 0:
return image.convert("RGB")
result = image.convert("RGB")
person_mask, person_boxes = yolo_sam_person_mask_and_boxes(result)
person_mask = person_mask.convert("L")
face_mask = qwen_face_feature_mask(result)
if person_mask.getbbox() is not None:
face_mask = ImageChops.multiply(face_mask, person_mask)
region_specs = [(face_mask, strength, False, 0.22, person_boxes)]
if include_hair:
hair_mask = qwen_hair_mask(result)
if person_mask.getbbox() is not None:
hair_mask = ImageChops.multiply(hair_mask, person_mask)
region_specs.append((hair_mask, strength * 0.85, True, 0.12, person_boxes))
for mask, region_strength, hair, pad_ratio, boxes in region_specs:
region_boxes = boxes or _mask_to_boxes(mask)
for box in region_boxes:
expanded = _expand_box(box, result.size, pad_ratio=pad_ratio, min_size=192)
crop = result.crop(expanded)
mask_crop = mask.crop(expanded)
if mask_crop.getbbox() is None:
mask_crop = Image.new("L", crop.size, 255)
detailed = _local_detail_crop(crop, mask_crop, strength=region_strength, hair=hair)
result.paste(detailed, expanded, mask_crop.filter(ImageFilter.GaussianBlur(radius=2.5)))
return result
def masked_texture_repair(image, strength=0.55):
base = image.convert("RGB")
mask = qwen_defect_mask(base)
repaired = base.filter(ImageFilter.MedianFilter(size=3)).filter(ImageFilter.GaussianBlur(radius=0.22))
repaired = ImageEnhance.Sharpness(repaired).enhance(1.06)
mask = mask.point(lambda value: int(value * max(0.0, min(1.0, strength))))
return Image.composite(repaired, base, mask)
def qwen_micro_detail(image, amount=0.35):
if amount <= 0:
return image.convert("RGB")
base = image.convert("RGB")
sharpened = base.filter(ImageFilter.UnsharpMask(radius=0.9, percent=int(95 * amount), threshold=3))
contrast = ImageEnhance.Contrast(sharpened).enhance(1.0 + amount * 0.05)
return Image.blend(base, contrast, alpha=min(0.65, amount))
def add_photographic_grain(image, seed):
return add_film_grain(image, seed)
# ============================================================================
# ENHANCED APPLY ENHANCEMENT (Main enhancement pipeline)
# ============================================================================
def apply_enhancement(image, enhance_mode, seed=0, progress=None):
"""Apply the Qwen 2511 / Rapid-AIO v23 post-process pipeline."""
mode = enhance_mode or ENHANCE_MODE_OFF
if mode not in ENHANCE_MODE_CHOICES:
raise gr.Error(f"Unknown enhance mode: {mode}")
if mode == ENHANCE_MODE_OFF:
return image
enhanced = image.convert("RGB")
total_steps = {
ENHANCE_MODE_UPSCALE: 2,
ENHANCE_MODE_CLEAN: 4,
ENHANCE_MODE_MAX_DETAIL: 7,
}[mode]
step = 0
if progress:
step += 1
progress(0.85 * step / total_steps, desc="Preparing Qwen output...")
enhanced = qwen_artifact_precondition(enhanced)
if mode in (ENHANCE_MODE_CLEAN, ENHANCE_MODE_MAX_DETAIL):
if progress:
step += 1
progress(0.85 * step / total_steps, desc="Repairing Qwen skin artifacts...")
enhanced = masked_texture_repair(enhanced, strength=0.55 if mode == ENHANCE_MODE_CLEAN else 0.68)
enhanced = qwen_micro_detail(enhanced, amount=0.18)
if progress:
step += 1
progress(0.85 * step / total_steps, desc="Detailing face regions...")
enhanced = qwen_tiled_detailer_pass(
enhanced,
strength=0.32 if mode == ENHANCE_MODE_CLEAN else 0.48,
include_hair=(mode == ENHANCE_MODE_MAX_DETAIL),
)
if mode == ENHANCE_MODE_MAX_DETAIL:
if progress:
step += 1
progress(0.85 * step / total_steps, desc="Building micro detail...")
enhanced = qwen_micro_detail(enhanced, amount=0.36)
if progress:
step += 1
progress(0.85 * step / total_steps, desc="Upscaling with Nomos tiles...")
enhanced = tile_upscale(enhanced)
if mode == ENHANCE_MODE_MAX_DETAIL:
if progress:
step += 1
progress(0.94, desc="Detailing final face and hair tiles...")
enhanced = qwen_tiled_detailer_pass(enhanced, strength=0.30, include_hair=True)
if progress:
step += 1
progress(0.98, desc="Adding final photographic grain...")
enhanced = qwen_micro_detail(enhanced, amount=0.22)
enhanced = add_photographic_grain(enhanced, seed)
return enhanced
# ============================================================================
# UTILITY FUNCTIONS
# ============================================================================
def use_output_as_input(output_images):
"""Move the first output image into the Image 1 slot."""
if not output_images:
return gr.update()
first = output_images[0]
path = first[0] if isinstance(first, (list, tuple)) else first
return gr.update(value=path)
def check_gpu_memory():
"""Check available GPU memory"""
if device == "cuda":
try:
total = torch.cuda.get_device_properties(0).total_memory
reserved = torch.cuda.memory_reserved(0)
allocated = torch.cuda.memory_allocated(0)
free = total - reserved
print(f"GPU Memory: Total={total/1024**3:.2f}GB, Reserved={reserved/1024**3:.2f}GB, Allocated={allocated/1024**3:.2f}GB, Free={free/1024**3:.2f}GB")
return free > 1024**3
except Exception as e:
print(f"Failed to check GPU memory: {e}")
return True
return True
def clear_gpu_cache():
"""Clear GPU cache to free up memory"""
if device == "cuda":
try:
torch.cuda.empty_cache()
import gc
gc.collect()
print("GPU cache cleared")
except Exception as e:
print(f"Failed to clear GPU cache: {e}")
# ============================================================================
# MAIN INFERENCE FUNCTION (Enhanced)
# ============================================================================
MAX_SEED = np.iinfo(np.int32).max
@spaces.GPU(duration=60)
def infer(
image_1,
image_2,
prompt,
seed=42,
randomize_seed=False,
true_guidance_scale=1.0,
num_inference_steps=4,
height=None,
width=None,
enhance_mode=ENHANCE_MODE_OFF,
num_images_per_prompt=1,
progress=gr.Progress(track_tqdm=True),
):
"""
Enhanced image generation with advanced editing and enhancement options
"""
negative_prompt = " "
if randomize_seed:
seed = random.randint(0, MAX_SEED)
generator = torch.Generator(device=device).manual_seed(seed)
pil_images = []
for img in (image_1, image_2):
if img is None:
continue
try:
if isinstance(img, str):
pil_images.append(Image.open(img).convert("RGB"))
elif isinstance(img, Image.Image):
pil_images.append(img.convert("RGB"))
elif hasattr(img, "name"):
pil_images.append(Image.open(img.name).convert("RGB"))
except Exception:
continue
if height==256 and width==256:
height, width = None, None
print(f"Generation Parameters:")
print(f" Prompt: '{prompt}'")
print(f" Seed: {seed}, Steps: {num_inference_steps}, Guidance: {true_guidance_scale}")
print(f" Size: {width}x{height}, Enhance Mode: {enhance_mode}")
if not check_gpu_memory():
clear_gpu_cache()
if not check_gpu_memory():
raise gr.Error("Insufficient GPU memory. Please reduce image size or close other applications.")
try:
images_pil = pipe(
image=pil_images if len(pil_images) > 0 else None,
prompt=prompt,
height=height,
width=width,
negative_prompt=negative_prompt,
num_inference_steps=num_inference_steps,
generator=generator,
true_cfg_scale=true_guidance_scale,
num_images_per_prompt=num_images_per_prompt,
).images
except Exception as e:
clear_gpu_cache()
raise gr.Error(f"Image generation failed: {e}")
if enhance_mode != ENHANCE_MODE_OFF:
images_pil = [
apply_enhancement(img, enhance_mode, seed=seed + idx, progress=progress)
for idx, img in enumerate(images_pil)
]
output_paths = []
os.makedirs("outputs", exist_ok=True)
for idx, img in enumerate(images_pil):
output_path = f"outputs/output_{seed}_{idx}_{int(time.time()*1000)}.png"
img.save(output_path)
output_paths.append(output_path)
clear_gpu_cache()
return output_paths, seed, gr.update(visible=True), gr.update(visible=bool(VIDEO_SPACE_ID))
# ============================================================================
# UI LAYOUT (Enhanced)
# ============================================================================
css = """
#col-container {
margin: 0 auto;
max-width: 1024px;
}
#logo-title {
text-align: center;
}
#logo-title h1 {
margin-bottom: 0;
}
#logo-title h2 {
color: #5b47d1;
font-style: italic;
margin-top: 0;
}
#edit_text{margin-top: -62px !important}
.enhance-info {
font-size: 0.9em;
color: #666;
margin-top: 5px;
}
"""
with gr.Blocks(css=css) as demo:
with gr.Column(elem_id="col-container"):
gr.HTML(f"""
<!-- v{APP_VERSION} -->
<div id="logo-title">
<h1>Pro Realism Edit Studio - Enhanced</h1>
<h2>Rapid Edit with Qwen-aware Nomos enhancement</h2>
</div>
""")
gr.Markdown("""
Powered by:
- [Qwen-Image-Edit-2511](https://huggingface.co/Qwen/Qwen-Image-Edit-2511)
- [Phr00t's Rapid-AIO v23](https://huggingface.co/Phr00t/Qwen-Image-Edit-Rapid-AIO) accelerated transformer
- YOLO-boxed SAM person masks with Nomos-family 4x tiled upscaling
Upload an image and enter your prompt to edit it.
Pro Tips:
- Use Clean for portraits with smudged or pitted skin
- Use Max Detail for hair, eyes, lashes, texture, grain, and sharper final output
""")
with gr.Row():
with gr.Column():
with gr.Row():
image_1 = gr.Image(label="Image 1", type="filepath", interactive=True)
image_2 = gr.Image(label="Image 2 (optional)", type="filepath", interactive=True)
prompt = gr.Text(
label="Prompt",
show_label=True,
placeholder="Enter your prompt here...",
)
enhance_mode = gr.Radio(
label="Enhance Mode",
choices=ENHANCE_MODE_CHOICES,
value=ENHANCE_MODE_OFF,
interactive=True,
info="Choose enhancement level for your output"
)
enhance_info = gr.Markdown("""
Enhancement Options:
- Off: No post-processing
- Upscale: Qwen precondition + 4x Nomos tiled upscale
- Clean: YOLO/SAM person mask + tiled face detail + upscaling
- Max Detail: YOLO/SAM person mask + tiled face/hair detail + micro detail + upscaling + final grain
""", visible=False)
run_button = gr.Button("Generate!", variant="primary")
with gr.Accordion("Advanced Settings", open=False):
seed = gr.Slider(
label="Seed",
minimum=0,
maximum=MAX_SEED,
step=1,
value=0,
)
randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
with gr.Row():
true_guidance_scale = gr.Slider(
label="True guidance scale",
minimum=1.0,
maximum=10.0,
step=0.1,
value=1.0
)
num_inference_steps = gr.Slider(
label="Number of inference steps",
minimum=1,
maximum=40,
step=1,
value=4,
)
with gr.Row():
height = gr.Slider(
label="Height",
minimum=256,
maximum=2048,
step=8,
value=None,
)
width = gr.Slider(
label="Width",
minimum=256,
maximum=2048,
step=8,
value=None,
)
gr.Markdown("""
Performance Tips:
- Use 4 steps for fastest results
- Increase steps (8-20) for better quality
- Lower guidance scale for more creative freedom
""")
with gr.Column():
result = gr.Gallery(label="Result", show_label=False, type="filepath")
with gr.Row():
use_output_btn = gr.Button("Use as input", variant="secondary", size="sm", visible=False)
turn_video_btn = gr.Button("Turn into Video", variant="secondary", size="sm", visible=False)
output_video = gr.Video(label="Generated Video", autoplay=True, visible=False)
with gr.Row():
gr.Markdown("### History")
clear_history_button = gr.Button("Clear History", size="sm", variant="stop")
history_gallery = gr.Gallery(
label="Click any image to use as input",
interactive=False,
show_label=True,
visible=True
)
gr.on(
triggers=[run_button.click, prompt.submit],
fn=infer,
inputs=[
image_1,
image_2,
prompt,
seed,
randomize_seed,
true_guidance_scale,
num_inference_steps,
height,
width,
enhance_mode,
],
outputs=[result, seed, use_output_btn, turn_video_btn],
).then(
fn=update_history,
inputs=[result, history_gallery],
outputs=history_gallery,
)
enhance_mode.change(
fn=lambda mode: gr.update(visible=mode != ENHANCE_MODE_OFF),
inputs=[enhance_mode],
outputs=[enhance_info]
)
use_output_btn.click(
fn=use_output_as_input,
inputs=[result],
outputs=[image_1]
)
history_gallery.select(
fn=use_history_as_input,
inputs=None,
outputs=[image_1],
)
clear_history_button.click(
fn=lambda: [],
inputs=None,
outputs=history_gallery,
)
turn_video_btn.click(
fn=lambda: gr.update(visible=True),
inputs=None,
outputs=[output_video],
).then(
fn=turn_into_video,
inputs=[image_1, result, prompt],
outputs=[output_video],
)
if __name__ == "__main__":
print(f"Device: {device}")
if device == "cuda":
print(f"GPU: {torch.cuda.get_device_name(0)}")
check_gpu_memory()
print(f"Starting Pro Realism Edit Studio v{APP_VERSION}")
demo.launch()