| import os |
| import gc |
| import time |
| import threading |
| import traceback |
| import types |
|
|
| |
| os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "backend:cudaMallocAsync") |
|
|
| import gradio as gr |
| import numpy as np |
| import spaces |
| import torch |
| import random |
| import base64 |
| import json |
| import html as html_lib |
| from io import BytesIO |
| from PIL import Image |
| from logging_utils import LogUploader |
|
|
| _log_uploader = LogUploader( |
| token=os.environ.get("HF_TOKEN"), |
| repo_id=os.environ.get("LOG_DATASET_REPO"), |
| max_files=int(os.environ.get("LOG_MAX_FILES", "5000")), |
| batch_interval=int(os.environ.get("LOG_BATCH_INTERVAL", "60")), |
| ) |
|
|
| MAX_SEED = np.iinfo(np.int32).max |
| LANCZOS = getattr(Image, "Resampling", Image).LANCZOS |
|
|
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
|
|
| print("CUDA_VISIBLE_DEVICES=", os.environ.get("CUDA_VISIBLE_DEVICES"), flush=True) |
| print("torch.__version__ =", torch.__version__, flush=True) |
| print("Using device:", device, flush=True) |
| print(f"CUDA device_count={torch.cuda.device_count()}, is_available={torch.cuda.is_available()}", flush=True) |
|
|
|
|
| def _log_env(): |
| import importlib.metadata as _meta |
| if torch.cuda.is_available(): |
| p = torch.cuda.get_device_properties(0) |
| print(f"[env] GPU: {p.name}, VRAM={p.total_memory/1024**3:.1f}GB, cap={p.major}.{p.minor}", flush=True) |
| print(f"[env] CUDA (torch build): {torch.version.cuda}", flush=True) |
| print(f"[env] cuDNN: {torch.backends.cudnn.version()}", flush=True) |
| for pkg in ["spaces", "diffusers", "transformers", "gradio", "accelerate", "peft", "torchvision"]: |
| try: |
| print(f"[env] {pkg}=={_meta.version(pkg)}", flush=True) |
| except Exception as e: |
| print(f"[env] {pkg}==? ({e})", flush=True) |
| try: |
| mem = {} |
| with open("/proc/meminfo") as f: |
| for line in f: |
| k, v = line.split(":", 1) |
| mem[k.strip()] = v.strip() |
| total_gb = int(mem["MemTotal"].split()[0]) / 1024**2 |
| avail_gb = int(mem["MemAvailable"].split()[0]) / 1024**2 |
| print(f"[env] RAM: {total_gb:.0f}GB total, {avail_gb:.0f}GB available", flush=True) |
| except Exception as e: |
| print(f"[env] RAM: unavailable ({e})", flush=True) |
|
|
| _log_env() |
|
|
| |
| torch.backends.cuda.matmul.allow_tf32 = True |
| torch.backends.cudnn.allow_tf32 = True |
| print("[startup] TF32 enabled", flush=True) |
|
|
| print("[startup] importing dimensions...", flush=True) |
| from dimensions import compute_output_dimensions |
| from mode import Mode |
| print("[startup] importing diffusers...", flush=True) |
| from diffusers import FlowMatchEulerDiscreteScheduler |
| from diffusers.models.normalization import RMSNorm |
| from transformers import Qwen2_5_VLForConditionalGeneration |
| print("[startup] importing QwenImageEditPlusPipeline...", flush=True) |
| from qwenimage.pipeline_qwenimage_edit_plus import QwenImageEditPlusPipeline |
| print("[startup] importing QwenImageTransformer2DModel...", flush=True) |
| from qwenimage.transformer_qwenimage import QwenImageTransformer2DModel |
| print("[startup] importing QwenDoubleStreamAttnProcessorFA3...", flush=True) |
| from qwenimage.qwen_fa3_processor import QwenDoubleStreamAttnProcessorFA3 |
| print("[startup] all imports done", flush=True) |
|
|
| dtype = torch.bfloat16 |
|
|
|
|
| def _start_heartbeat(label: str) -> threading.Event: |
| done = threading.Event() |
| t0 = time.perf_counter() |
| def _beat(): |
| while not done.wait(timeout=15): |
| print(f"[startup] {label} still loading... ({time.perf_counter()-t0:.0f}s)", flush=True) |
| threading.Thread(target=_beat, daemon=True).start() |
| return done |
|
|
|
|
| _FP8_DTYPES = (torch.float8_e4m3fn, torch.float8_e5m2) |
|
|
|
|
| def _fp8_upcast_linear_forward(self, input): |
| weight = self.weight.to(input.dtype) if self.weight.dtype in _FP8_DTYPES else self.weight |
| bias = self.bias.to(input.dtype) if (self.bias is not None and self.bias.dtype in _FP8_DTYPES) else self.bias |
| return torch.nn.functional.linear(input, weight, bias) |
|
|
|
|
| def _fp8_upcast_rmsnorm_forward(self, hidden_states): |
| |
| |
| |
| |
| |
| input_dtype = hidden_states.dtype |
| variance = hidden_states.to(torch.float32).pow(2).mean(-1, keepdim=True) |
| hidden_states = hidden_states * torch.rsqrt(variance + self.eps) |
|
|
| if self.weight is not None: |
| weight = self.weight.to(input_dtype) if self.weight.dtype in _FP8_DTYPES else self.weight |
| if weight.dtype in (torch.float16, torch.bfloat16): |
| hidden_states = hidden_states.to(weight.dtype) |
| hidden_states = hidden_states * weight |
| if self.bias is not None: |
| bias = self.bias.to(hidden_states.dtype) if self.bias.dtype in _FP8_DTYPES else self.bias |
| hidden_states = hidden_states + bias |
| else: |
| hidden_states = hidden_states.to(input_dtype) |
|
|
| return hidden_states |
|
|
|
|
| def _patch_fp8_modules(model) -> int: |
| |
| |
| |
| |
| |
| |
| |
| |
| count = 0 |
| for module in model.modules(): |
| if isinstance(module, torch.nn.Linear) and module.weight.dtype in _FP8_DTYPES: |
| module.forward = types.MethodType(_fp8_upcast_linear_forward, module) |
| count += 1 |
| elif isinstance(module, RMSNorm) and module.weight is not None and module.weight.dtype in _FP8_DTYPES: |
| module.forward = types.MethodType(_fp8_upcast_rmsnorm_forward, module) |
| count += 1 |
|
|
| |
| |
| |
| patched_types = (torch.nn.Linear, RMSNorm) |
| for name, module in model.named_modules(): |
| if isinstance(module, patched_types): |
| continue |
| for pname, param in module.named_parameters(recurse=False): |
| if param.dtype in _FP8_DTYPES: |
| print( |
| f"[startup] WARNING: unpatched fp8 parameter {name}.{pname} " |
| f"({type(module).__name__}) — will likely error at inference", |
| flush=True, |
| ) |
| return count |
|
|
|
|
| _t0_load = time.perf_counter() |
| print("[startup] loading transformer from_pretrained (prithivMLmods/Qwen-Image-Edit-Rapid-AIO-V23)...", flush=True) |
| _hb = _start_heartbeat("transformer") |
| _transformer = QwenImageTransformer2DModel.from_pretrained( |
| "prithivMLmods/Qwen-Image-Edit-Rapid-AIO-V23", |
| torch_dtype=torch.float8_e4m3fn, |
| device_map="cpu", |
| ) |
| _hb.set() |
| print(f"[startup] transformer loaded in {time.perf_counter()-_t0_load:.1f}s", flush=True) |
| _n_fp8_patched = _patch_fp8_modules(_transformer) |
| print(f"[startup] patched {_n_fp8_patched} fp8-resident nn.Linear/RMSNorm modules for just-in-time upcast", flush=True) |
| try: |
| print(f"[startup] transformer memory footprint: {_transformer.get_memory_footprint()/1024**3:.2f}GB", flush=True) |
| except Exception as e: |
| print(f"[startup] transformer memory footprint: unavailable ({e})", flush=True) |
|
|
| _t1_load = time.perf_counter() |
| print("[startup] loading pipeline from_pretrained (FireRedTeam/FireRed-Image-Edit-1.1)...", flush=True) |
| _hb = _start_heartbeat("pipeline") |
| pipe = QwenImageEditPlusPipeline.from_pretrained( |
| "FireRedTeam/FireRed-Image-Edit-1.1", |
| transformer=_transformer, |
| torch_dtype=dtype, |
| ) |
| _hb.set() |
| pipe.vae.enable_tiling(tile_sample_min_height=Mode.HIGH_DETAIL.max_dim, tile_sample_min_width=Mode.HIGH_DETAIL.max_dim) |
| print(f"[startup] VAE tiling: threshold={pipe.vae.tile_sample_min_height}x{pipe.vae.tile_sample_min_width}px use_tiling={pipe.vae.use_tiling}", flush=True) |
| print(f"[startup] pipeline loaded in {time.perf_counter()-_t1_load:.1f}s", flush=True) |
|
|
| print("[startup] setting cuDNN SDPA attention processor...", flush=True) |
| pipe.transformer.set_attn_processor(QwenDoubleStreamAttnProcessorFA3()) |
| print("[startup] cuDNN SDPA attention processor set.", flush=True) |
|
|
| with open("examples.json") as _f: |
| EXAMPLES_CONFIG = json.load(_f) |
|
|
| with open("suggestions.json") as _f: |
| SUGGESTIONS_CONFIG = json.load(_f) |
|
|
|
|
| def make_thumb_b64(path, max_dim=220): |
| if not os.path.exists(path): |
| return "" |
| try: |
| img = Image.open(path).convert("RGB") |
| img.thumbnail((max_dim, max_dim), LANCZOS) |
| buf = BytesIO() |
| img.save(buf, format="JPEG", quality=65) |
| return f"data:image/jpeg;base64,{base64.b64encode(buf.getvalue()).decode()}" |
| except Exception as e: |
| print(f"Thumbnail error for {path}: {e}") |
| return "" |
|
|
|
|
| def encode_full_image(path): |
| if not os.path.exists(path): |
| return "" |
| try: |
| with open(path, "rb") as f: |
| data = f.read() |
| ext = path.rsplit(".", 1)[-1].lower() |
| mime = {"jpg": "image/jpeg", "jpeg": "image/jpeg", "png": "image/png", "webp": "image/webp"}.get(ext, "image/jpeg") |
| return f"data:{mime};base64,{base64.b64encode(data).decode()}" |
| except Exception as e: |
| print(f"Encode error for {path}: {e}") |
| return "" |
|
|
|
|
| def _example_thumbs_html(images): |
| html = "" |
| for path in images: |
| thumb = make_thumb_b64(path) |
| if thumb: |
| html += f'<img src="{thumb}" alt="">' |
| else: |
| html += '<div class="example-thumb-placeholder">Preview</div>' |
| return html |
|
|
|
|
| def _example_card_html(idx, ex): |
| thumbs_html = _example_thumbs_html(ex["images"]) |
| n = len(ex["images"]) |
| badge = f'{n} image{"s" if n > 1 else ""}' |
| prompt_short = html_lib.escape(ex["prompt"][:90]) |
| if len(ex["prompt"]) > 90: |
| prompt_short += "..." |
| return f'''<div class="example-card" data-idx="{idx}"> |
| <div class="example-thumbs">{thumbs_html}</div> |
| <div class="example-meta"><span class="example-badge">{badge}</span></div> |
| <div class="example-prompt-text">{prompt_short}</div> |
| </div>''' |
|
|
|
|
| def build_example_cards_html(): |
| return "".join(_example_card_html(i, ex) for i, ex in enumerate(EXAMPLES_CONFIG)) |
|
|
|
|
| def _parse_example_idx(idx_str): |
| try: |
| return int(float(idx_str)) if idx_str and idx_str.strip() else -1 |
| except (ValueError, TypeError): |
| return -1 |
|
|
|
|
| def load_example_data(idx_str): |
| idx = _parse_example_idx(idx_str) |
| if idx < 0 or idx >= len(EXAMPLES_CONFIG): |
| return json.dumps({"images": [], "prompt": "", "names": [], "status": "error"}) |
| ex = EXAMPLES_CONFIG[idx] |
| b64_list, names = [], [] |
| for path in ex["images"]: |
| b64 = encode_full_image(path) |
| if b64: |
| b64_list.append(b64) |
| names.append(os.path.basename(path)) |
| return json.dumps({"images": b64_list, "prompt": ex["prompt"], "names": names, "status": "ok"}) |
|
|
|
|
| def build_suggestion_chips_html(): |
| chips = [] |
| for s in SUGGESTIONS_CONFIG: |
| prompt_json = html_lib.escape(json.dumps(s["prompt"])) |
| label = html_lib.escape(s["label"]) |
| chips.append(f'<button class="suggestion-chip" onclick="window.__setPrompt({prompt_json})">{label}</button>') |
| return "".join(chips) |
|
|
|
|
| print("Building example thumbnails...") |
| EXAMPLE_CARDS_HTML = build_example_cards_html() |
| print(f"Built {len(EXAMPLES_CONFIG)} example cards.") |
| SUGGESTION_CHIPS_HTML = build_suggestion_chips_html() |
| print(f"Built {len(SUGGESTIONS_CONFIG)} suggestion chips.") |
|
|
|
|
| def b64_to_pil_list(b64_json_str): |
| if not b64_json_str or b64_json_str.strip() in ("", "[]"): |
| return [] |
| try: |
| b64_list = json.loads(b64_json_str) |
| except Exception: |
| return [] |
| pil_images = [] |
| for b64_str in b64_list: |
| if not b64_str or not isinstance(b64_str, str): |
| continue |
| try: |
| if b64_str.startswith("data:image"): |
| _, data = b64_str.split(",", 1) |
| else: |
| data = b64_str |
| image_data = base64.b64decode(data) |
| pil_images.append(Image.open(BytesIO(image_data)).convert("RGB")) |
| except Exception as e: |
| print(f"Error decoding image: {e}") |
| return pil_images |
|
|
|
|
| def update_dimensions_on_upload(image, max_dim): |
| if image is None: |
| return max_dim, max_dim |
| w, h = image.size |
| return compute_output_dimensions(w, h, max_dim) |
|
|
|
|
| class _InferTimer: |
| def __init__(self, cuda_ok: bool) -> None: |
| self._cuda_ok = cuda_ok |
| self._marks: dict = {} |
|
|
| def mark(self, name: str) -> None: |
| ev = None |
| if self._cuda_ok: |
| ev = torch.cuda.Event(enable_timing=True) |
| ev.record() |
| self._marks[name] = (ev, time.perf_counter()) |
|
|
| def elapsed_ms(self, a: str, b: str) -> float: |
| ev_a, t_a = self._marks[a] |
| ev_b, t_b = self._marks[b] |
| if ev_a and ev_b: |
| return ev_a.elapsed_time(ev_b) |
| return (t_b - t_a) * 1000.0 |
|
|
| def wall_start(self, name: str) -> float: |
| return self._marks[name][1] |
|
|
| def __contains__(self, name: str) -> bool: |
| return name in self._marks |
|
|
| def print_timings(self) -> None: |
| if self._cuda_ok: |
| try: |
| torch.cuda.synchronize() |
| except Exception: |
| pass |
| rows = [ |
| ("image_load", "load_start", "load_end"), |
| ("preprocess", "pipe_start", "first_step"), |
| ("inference", "first_step", "last_step"), |
| ("vae_decode", "last_step", "pipe_end"), |
| ] |
| total_ms = 0.0 |
| lines = [] |
| for label, a, b in rows: |
| if a in self._marks and b in self._marks: |
| ms = self.elapsed_ms(a, b) |
| total_ms += ms |
| lines.append(f"[timing] {label:<14} {ms:8.1f} ms") |
| if "load_start" in self._marks and "pipe_end" in self._marks: |
| overall_ms = self.elapsed_ms("load_start", "pipe_end") |
| lines.append(f"[timing] {'overhead':<14} {overall_ms - total_ms:8.1f} ms") |
| lines.append(f"[timing] {'── total ──':<14} {overall_ms:8.1f} ms") |
| print("[timing] ─────────────────────────────────────") |
| print("\n".join(lines)) |
| print("[timing] ─────────────────────────────────────") |
|
|
|
|
| def _gpu_mem_str(cuda_ok: bool, sync: bool = False) -> str: |
| if not cuda_ok: |
| return "CUDA not available" |
| if sync: |
| try: |
| torch.cuda.synchronize() |
| except Exception as se: |
| return f"CUDA sync failed: {se}" |
| alloc = torch.cuda.memory_allocated() / 1024**3 |
| reserved = torch.cuda.memory_reserved() / 1024**3 |
| peak = torch.cuda.max_memory_allocated() / 1024**3 |
| return f"alloc={alloc:.2f}GB reserved={reserved:.2f}GB peak={peak:.2f}GB" |
|
|
|
|
| def _validate_infer_inputs(pil_images: list, prompt: str) -> None: |
| if not pil_images: |
| raise gr.Error("Please upload at least one image to edit.") |
| if not prompt or prompt.strip() == "": |
| raise gr.Error("Please enter an edit prompt.") |
|
|
|
|
| def _resolve_seed(seed: int, randomize_seed: bool) -> int: |
| return random.randint(0, MAX_SEED) if randomize_seed else seed |
|
|
|
|
| def _spawn_log(pil_images, result_image, prompt, seed, steps, guidance_scale, |
| width, height, duration, success, error=""): |
| threading.Thread( |
| target=_log_uploader.log_inference, |
| args=(pil_images, result_image, prompt, seed, steps, guidance_scale, |
| width, height, duration, success, error), |
| daemon=True, |
| ).start() |
|
|
|
|
| |
|
|
| with open("static/app.css") as _f: |
| css = _f.read() |
|
|
| with open("static/gallery.js") as _f: |
| gallery_js = _f.read() |
|
|
| with open("static/wire_outputs.js") as _f: |
| wire_outputs_js = _f.read() |
|
|
| with open("static/run_preprocess.js") as _f: |
| run_preprocess_js = _f.read() |
|
|
| with open("static/mode_toggle.js") as _f: |
| mode_toggle_js = _f.read() |
|
|
| with open("static/negative_prompt.txt") as _f: |
| negative_prompt = _f.read().strip() |
|
|
| |
|
|
| with open("templates/app.html") as _f: |
| app_html = _f.read().format( |
| example_cards_html=EXAMPLE_CARDS_HTML, |
| suggestion_chips_html=SUGGESTION_CHIPS_HTML, |
| ) |
|
|
| |
|
|
| def infer(images_b64_json, prompt, seed, randomize_seed, guidance_scale, steps, mode, gpu_duration=20, progress=gr.Progress(track_tqdm=True)): |
| |
| gc.collect() |
| mode = Mode.from_value(mode) |
| pil_images = b64_to_pil_list(images_b64_json) |
| _validate_infer_inputs(pil_images, prompt) |
| seed = _resolve_seed(seed, randomize_seed) |
| width, height = update_dimensions_on_upload(pil_images[0], mode.max_dim) |
| t0 = time.perf_counter() |
| try: |
| result_image, seed, duration = _infer_gpu(pil_images, prompt, seed, guidance_scale, steps, width, height, mode, int(gpu_duration)) |
| |
| |
| _spawn_log(pil_images, result_image, prompt, seed, steps, guidance_scale, width, height, duration, True) |
| return result_image, seed |
| except Exception as e: |
| duration = time.perf_counter() - t0 |
| |
| |
| |
| |
| |
| print(f"[infer] EXCEPTION type={type(e).__module__}.{type(e).__qualname__} repr={e!r}") |
| traceback.print_exc() |
| _spawn_log(pil_images, None, prompt, seed, steps, guidance_scale, width, height, duration, False, str(e)) |
| raise |
|
|
|
|
| def _log_infer_start(prompt, steps, guidance_scale, seed, gpu_duration, mode: Mode): |
| print(f"[infer] ===== START =====") |
| print(f"[infer] steps={steps}, guidance={guidance_scale}, seed={seed}, gpu_duration={gpu_duration}s, mode={mode.value}") |
| print(f"[infer] prompt={repr(prompt[:120])}") |
|
|
|
|
| def _log_gpu_properties(cuda_ok): |
| if not cuda_ok: |
| return None |
| p = torch.cuda.get_device_properties(0) |
| print(f"[infer] GPU: {p.name}, total={p.total_memory/1024**3:.1f}GB, cap={p.major}.{p.minor}") |
| torch.cuda.reset_peak_memory_stats() |
| return p |
|
|
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| _FAST_PATH_MIN_GB = 40 |
|
|
| |
| |
| |
| _TEXT_ENCODER_INT8_REPO = os.environ.get("TEXT_ENCODER_INT8_REPO") |
|
|
|
|
| def _ensure_int8_text_encoder(cuda_ok, t0): |
| |
| |
| |
| |
| |
| |
| if not cuda_ok or not _TEXT_ENCODER_INT8_REPO or getattr(pipe, "_text_encoder_is_int8", False): |
| return |
| try: |
| _t_load = time.perf_counter() |
| quantized = Qwen2_5_VLForConditionalGeneration.from_pretrained( |
| _TEXT_ENCODER_INT8_REPO, |
| device_map={"": device}, |
| dtype=torch.bfloat16, |
| ) |
| pipe.text_encoder = quantized |
| pipe._text_encoder_is_int8 = True |
| print( |
| f"[infer] loaded int8 text_encoder from {_TEXT_ENCODER_INT8_REPO} — " |
| f"{(time.perf_counter()-_t_load)*1000:.0f}ms | t={time.perf_counter()-t0:.1f}s" |
| ) |
| except Exception as e: |
| print(f"[infer] WARNING: int8 text_encoder load failed, keeping bf16: {type(e).__name__}: {e}") |
|
|
|
|
| def _place_pipe_on_device(cuda_ok, gpu_props, t0): |
| if getattr(pipe.transformer, "_hf_hook", None) is not None: |
| return |
| if cuda_ok and gpu_props.total_memory / 1024**3 >= _FAST_PATH_MIN_GB: |
| try: |
| pipe.to(device) |
| print(f"[infer] moved full pipe to {device} — t={time.perf_counter()-t0:.1f}s") |
| return |
| except torch.cuda.OutOfMemoryError: |
| print(f"[infer] OOM moving full pipe to {device}, falling back to cpu offload") |
| pipe.to("cpu") |
| torch.cuda.empty_cache() |
| pipe.enable_model_cpu_offload(device=device) |
| print(f"[infer] enabled cpu offload on {device} (fallback)") |
| return |
| pipe.enable_model_cpu_offload(device=device) |
| print(f"[infer] enabled cpu offload on {device} (slice too small for fast path)") |
|
|
|
|
| def _instrument_first_touch(modules_with_names, t0): |
| """Install self-removing forward-pre-hooks that log the moment each module is first entered.""" |
| def _make_hook(name, handle_box): |
| def _hook(mod, inputs): |
| print(f"[infer] first call into {name} — {_gpu_mem_str(True, sync=True)} | t={time.perf_counter()-t0:.1f}s") |
| handle_box["h"].remove() |
| return _hook |
| for module, name in modules_with_names: |
| handle_box = {} |
| handle_box["h"] = module.register_forward_pre_hook(_make_hook(name, handle_box)) |
|
|
|
|
| def _make_step_callback(steps, timer, t0, mode: Mode, cuda_ok: bool = False): |
| """Build the diffusers step callback that logs per-step timing and marks timer checkpoints.""" |
| step_times = [] |
| def _step_cb(pipeline, step_idx, timestep, cb_kwargs): |
| now = time.perf_counter() |
| step_times.append(now) |
| if step_idx == 0: |
| timer.mark("first_step") |
| timer.mark("last_step") |
| delta_ms = (now - (step_times[-2] if len(step_times) > 1 else t0)) * 1000 |
| tag = " ← includes cold-start (offload hook install + first weight transfer)" if step_idx == 0 else "" |
| print(f"[infer] step {step_idx+1}/{steps} done — {delta_ms:.0f}ms{tag} | t={now-t0:.1f}s") |
| |
| |
| |
| |
| |
| |
| |
| if step_idx == steps - 1 and getattr(pipeline.text_encoder, "_hf_hook", None) is None: |
| if mode.offloads_text_encoder_before_decode and not getattr(pipeline, "_text_encoder_is_int8", False): |
| _offload_t0 = time.perf_counter() |
| pipeline.text_encoder.to("cpu") |
| torch.cuda.empty_cache() |
| _offload_ms = (time.perf_counter() - _offload_t0) * 1000 |
| print(f"[infer] text_encoder offload to cpu — {_offload_ms:.0f}ms | t={time.perf_counter()-t0:.1f}s") |
| elif getattr(pipeline, "_text_encoder_is_int8", False): |
| print("[infer] skipping text_encoder offload (int8, .to() unsupported / smaller footprint)") |
| else: |
| print(f"[infer] skipping text_encoder offload for mode={mode.value} (ample headroom at this resolution)") |
| if step_idx == steps - 1 and cuda_ok: |
| print(f"[infer] pre-VAE-decode — {_gpu_mem_str(True, sync=True)} | t={time.perf_counter()-t0:.1f}s") |
| torch.cuda.reset_peak_memory_stats() |
| return cb_kwargs |
| return _step_cb |
|
|
|
|
| def _log_infer_error(e, t0, timer): |
| print(f"[infer] ERROR: {type(e).__name__}: {e} | t={time.perf_counter()-t0:.1f}s") |
| print(traceback.format_exc()) |
| try: |
| torch.cuda.synchronize() |
| except Exception as cuda_err: |
| print(f"[infer] CUDA synchronize after error: {cuda_err}") |
| timer.print_timings() |
|
|
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| _COLD_START_BUFFER_S = 45 |
| _MAX_GPU_DURATION_S = 120 |
|
|
|
|
| @spaces.GPU(duration=lambda *a, **kw: min(int(a[8]) + _COLD_START_BUFFER_S, _MAX_GPU_DURATION_S) if len(a) > 8 else 60) |
| def _infer_gpu(pil_images, prompt, seed, guidance_scale, steps, width, height, mode: Mode, gpu_duration=20): |
| _cuda_ok = torch.cuda.is_available() |
| timer = _InferTimer(_cuda_ok) |
| t0 = time.perf_counter() |
|
|
| _log_infer_start(prompt, steps, guidance_scale, seed, gpu_duration, mode) |
| gpu_props = _log_gpu_properties(_cuda_ok) |
|
|
| _ensure_int8_text_encoder(_cuda_ok, t0) |
| _place_pipe_on_device(_cuda_ok, gpu_props, t0) |
| print(f"[infer] {_gpu_mem_str(_cuda_ok)} — t={time.perf_counter()-t0:.1f}s") |
|
|
| if _cuda_ok: |
| _instrument_first_touch( |
| [(pipe.text_encoder, "text_encoder"), (pipe.transformer, "transformer"), (pipe.vae, "vae")], |
| t0, |
| ) |
|
|
| print(f"[infer] {len(pil_images)} image(s) pre-decoded, output={width}x{height}, seed={seed}") |
| if _cuda_ok: |
| _will_tile = pipe.vae.use_tiling and ( |
| width > pipe.vae.tile_sample_min_width or height > pipe.vae.tile_sample_min_height |
| ) |
| print(f"[infer] VAE tiling will {'activate' if _will_tile else 'NOT activate'} " |
| f"(threshold={pipe.vae.tile_sample_min_height}x{pipe.vae.tile_sample_min_width}px)") |
|
|
| generator = torch.Generator(device=device).manual_seed(seed) |
| step_cb = _make_step_callback(steps, timer, t0, mode, _cuda_ok) |
|
|
| timer.mark("pipe_start") |
| print(f"[infer] calling pipe... t={time.perf_counter()-t0:.1f}s") |
| try: |
| result_image = pipe( |
| image=pil_images, prompt=prompt, negative_prompt=negative_prompt, |
| height=height, width=width, num_inference_steps=steps, |
| generator=generator, true_cfg_scale=guidance_scale, |
| callback_on_step_end=step_cb, |
| callback_on_step_end_tensor_inputs=["latents"], |
| ).images[0] |
| timer.mark("pipe_end") |
| print(f"[infer] VAE decode + postprocess done — {_gpu_mem_str(_cuda_ok, sync=True)} | t={time.perf_counter()-t0:.1f}s") |
| timer.print_timings() |
| duration = timer.elapsed_ms("pipe_start", "pipe_end") / 1000.0 |
| return result_image, seed, duration |
| except Exception as e: |
| _log_infer_error(e, t0, timer) |
| raise |
| finally: |
| |
| |
| |
| |
| |
| gc.collect() |
| torch.cuda.empty_cache() |
| print(f"[infer] ===== END t={time.perf_counter()-t0:.1f}s =====") |
|
|
|
|
| with gr.Blocks() as demo: |
|
|
| hidden_images_b64 = gr.Textbox(value="[]", elem_id="hidden-images-b64", elem_classes="hidden-input", container=False) |
| prompt = gr.Textbox(value="", elem_id="prompt-gradio-input", elem_classes="hidden-input", container=False) |
| seed = gr.Slider(minimum=0, maximum=MAX_SEED, step=1, value=0, elem_id="gradio-seed", elem_classes="hidden-input", container=False) |
| randomize_seed = gr.Checkbox(value=True, elem_id="gradio-randomize", elem_classes="hidden-input", container=False) |
| guidance_scale = gr.Slider(minimum=1.0, maximum=10.0, step=0.1, value=1.0, elem_id="gradio-guidance", elem_classes="hidden-input", container=False) |
| steps = gr.Slider(minimum=1, maximum=50, step=1, value=3, elem_id="gradio-steps", elem_classes="hidden-input", container=False) |
| mode = gr.Textbox(value="fast", elem_id="gradio-mode", elem_classes="hidden-input", container=False) |
| gpu_duration = gr.Slider(minimum=10, maximum=120, step=5, value=15, elem_id="gradio-gpu-duration", elem_classes="hidden-input", container=False) |
| result = gr.Image(elem_id="gradio-result", elem_classes="hidden-input", container=False, format="png") |
|
|
| example_idx = gr.Textbox(value="", elem_id="example-idx-input", elem_classes="hidden-input", container=False) |
| example_result = gr.Textbox(value="", elem_id="example-result-data", elem_classes="hidden-input", container=False) |
| example_load_btn = gr.Button("Load Example", elem_id="example-load-btn") |
|
|
| gr.HTML(app_html) |
|
|
| run_btn = gr.Button("Run", elem_id="gradio-run-btn") |
|
|
| demo.load(fn=None, js=gallery_js) |
| demo.load(fn=None, js=wire_outputs_js) |
| demo.load(fn=None, js=mode_toggle_js) |
|
|
| run_btn.click( |
| fn=infer, |
| inputs=[hidden_images_b64, prompt, seed, randomize_seed, guidance_scale, steps, mode, gpu_duration], |
| outputs=[result, seed], |
| js=run_preprocess_js, |
| ) |
|
|
| example_load_btn.click( |
| fn=load_example_data, |
| inputs=[example_idx], |
| outputs=[example_result], |
| queue=False, |
| ) |
|
|
| if __name__ == "__main__": |
| demo.queue(max_size=30).launch( |
| css=css, |
| mcp_server=True, |
| ssr_mode=False, |
| show_error=True, |
| allowed_paths=["examples"], |
| ) |
|
|