"""Prompt-enhancement service for Qwen-Image-2.1. Hosts the two official prompt-rewriting models that ship alongside Qwen-Image-2.1: * ``Qwen/Qwen-Image-2.1-PE-T2I`` — turns a short text-to-image brief into a long, detailed English prompt plus a recommended aspect ratio. * ``Qwen/Qwen-Image-2.1-PE-I2I`` — turns a vague editing instruction plus the input image(s) into a precise, actionable editing directive. Both are loaded once at module scope and moved to CUDA eagerly, so the ZeroGPU runtime packs them at startup and streams them straight into VRAM on the first request. They stay resident on the worker between calls; nothing is cold-loaded per request. """ import os os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") import spaces # noqa: E402 — must precede torch / any CUDA-touching import import json # noqa: E402 import random # noqa: E402 import re # noqa: E402 import time # noqa: E402 import gradio as gr # noqa: E402 import torch # noqa: E402 from huggingface_hub import hf_hub_download # noqa: E402 from PIL import Image # noqa: E402 from transformers import AutoModelForImageTextToText, AutoProcessor # noqa: E402 T2I_MODEL_ID = "Qwen/Qwen-Image-2.1-PE-T2I" I2I_MODEL_ID = "Qwen/Qwen-Image-2.1-PE-I2I" MAX_SEED = 2**31 - 1 MAX_INPUT_IMAGES = 10 # ~23 tokens/s of decode on this hardware, so the cap sets the latency. 2048 is enough for # a full non-thinking answer from either rewriter; thinking needs considerably more. DEFAULT_MAX_NEW_TOKENS = 2048 # Aspect ratios the rewriter is allowed to recommend, per the PE-T2I system prompt. KNOWN_RATIOS = { "1:1", "3:2", "2:3", "16:9", "9:16", "4:3", "3:4", "1:2", "2:1", "21:9", "9:21", "4:5", "5:4", "3:1", "1:3", } def _load(model_id: str): """Load one rewriter and place it on the GPU eagerly (ZeroGPU packs it at startup).""" print(f"[load] {model_id} …", flush=True) processor = AutoProcessor.from_pretrained(model_id) model = AutoModelForImageTextToText.from_pretrained( model_id, dtype=torch.bfloat16, attn_implementation="sdpa" ) model = model.eval().to("cuda") system_prompt = open(hf_hub_download(model_id, "system_prompt.txt")).read().strip() print(f"[load] {model_id} ready ({len(system_prompt)} char system prompt)", flush=True) return processor, model, system_prompt T2I_PROCESSOR, T2I_MODEL, T2I_SYSTEM_PROMPT = _load(T2I_MODEL_ID) I2I_PROCESSOR, I2I_MODEL, I2I_SYSTEM_PROMPT = _load(I2I_MODEL_ID) def _salvage_truncated(answer: str) -> dict | None: match = re.search(r'"rewritten_prompt"\s*:\s*"', answer) if not match: return None body = answer[match.end() :] closing = re.search(r'(? dict: """Pull the JSON answer out of a rewriter's ``…{json}`` output.""" _thinking, _sep, answer = generated.partition("") answer = (answer or generated).strip() # Strip a markdown fence if the model wrapped the JSON in one. fence = re.match(r"^```(?:json)?\s*(.*?)\s*```$", answer, re.DOTALL) if fence: answer = fence.group(1).strip() try: parsed = json.loads(answer) except json.JSONDecodeError: start, end = answer.find("{"), answer.rfind("}") parsed = None if start != -1 and end > start: try: parsed = json.loads(answer[start : end + 1]) except json.JSONDecodeError: parsed = None if parsed is None: parsed = _salvage_truncated(answer) if parsed is None: raise ValueError( "the rewriter did not return JSON — its answer was almost certainly cut off. " "Raise 'Max new tokens', or turn 'Enable thinking' off." ) if not isinstance(parsed, dict): raise ValueError("the rewriter returned JSON that is not an object") return parsed def _normalize_ratio(value) -> str: value = (value or "").strip() return value if value in KNOWN_RATIOS else "" def _to_pil(item) -> Image.Image: if isinstance(item, Image.Image): return item.convert("RGB") path = getattr(item, "name", None) or (item.get("path") if isinstance(item, dict) else None) or item return Image.open(path).convert("RGB") _T2I_PREFIX: dict = {} def _t2i_prefix() -> dict: if "ids" not in _T2I_PREFIX: from transformers import DynamicCache def render(text): return T2I_PROCESSOR.apply_chat_template( [ {"role": "system", "content": [{"type": "text", "text": T2I_SYSTEM_PROMPT}]}, {"role": "user", "content": [{"type": "text", "text": text}]}, ], add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt", enable_thinking=False, )["input_ids"][0] a, b = render("a corgi"), render("zebra stripes") n = 0 while n < min(len(a), len(b)) and a[n] == b[n]: n += 1 ids = a[:n][None].to(T2I_MODEL.device) cache = DynamicCache(config=T2I_MODEL.config) with torch.no_grad(): T2I_MODEL(input_ids=ids, past_key_values=cache, use_cache=True) _T2I_PREFIX.update(ids=ids, cache=cache) print(f"[prefix] cached {n} system-prompt tokens for T2I", flush=True) return _T2I_PREFIX def _prefix_kwargs(inputs, prefix: dict) -> dict: import copy ids = prefix["ids"] n = ids.shape[1] if inputs["input_ids"].shape[1] <= n or not torch.equal(inputs["input_ids"][:, :n], ids): print("[prefix] tokenization does not share the cached prefix, running the full prompt", flush=True) return {} return {"past_key_values": copy.deepcopy(prefix["cache"])} def _estimate_duration(prompt, image_paths=None, max_new_tokens=DEFAULT_MAX_NEW_TOKENS, *args, **kwargs): """Budget GPU seconds from the token cap — decoding dominates the call. Measured on this Space: ~23 tokens/s of decode, plus a few seconds of prefill per input image. The cap is the only hard bound on the call, so budget against it. """ try: tokens = int(max_new_tokens) except (TypeError, ValueError): tokens = DEFAULT_MAX_NEW_TOKENS n_images = len(image_paths) if image_paths else 0 return int(min(400, 8 + tokens * 0.047 + n_images * 4)) @spaces.GPU(duration=_estimate_duration) def enhance( prompt: str, image_paths: list | None = None, max_new_tokens: int = DEFAULT_MAX_NEW_TOKENS, enable_thinking: bool = False, seed: int = 0, randomize_seed: bool = True, progress=gr.Progress(track_tqdm=True), ): """Rewrite an image prompt with the official Qwen-Image-2.1 prompt-enhancement models. Routes to Qwen-Image-2.1-PE-I2I when input images are supplied and to Qwen-Image-2.1-PE-T2I otherwise. Both models stay resident on the GPU worker. Args: prompt: The user's short or vague image request, in any language. image_paths: Optional input images for an editing request (up to 10). max_new_tokens: Decoding cap, including the model's reasoning block. enable_thinking: Let the model reason before emitting its JSON answer. seed: Sampling seed. randomize_seed: Draw a fresh seed instead of using ``seed``. Returns: A tuple of (rewritten prompt, recommended aspect ratio or "", raw model output, resolved seed, elapsed seconds). """ prompt = (prompt or "").strip() if not prompt: raise gr.Error("Provide a prompt to rewrite.") image_paths = [p for p in (image_paths or []) if p is not None] if len(image_paths) > MAX_INPUT_IMAGES: raise gr.Error(f"Up to {MAX_INPUT_IMAGES} input images are supported.") if randomize_seed: seed = random.randint(0, MAX_SEED) seed = int(seed) % (MAX_SEED + 1) torch.manual_seed(seed) is_edit = len(image_paths) > 0 processor = I2I_PROCESSOR if is_edit else T2I_PROCESSOR model = I2I_MODEL if is_edit else T2I_MODEL system_prompt = I2I_SYSTEM_PROMPT if is_edit else T2I_SYSTEM_PROMPT user_content = [] if is_edit: user_content += [{"type": "image", "image": _to_pil(p)} for p in image_paths] user_content.append({"type": "text", "text": prompt}) messages = [ {"role": "system", "content": [{"type": "text", "text": system_prompt}]}, {"role": "user", "content": user_content}, ] inputs = processor.apply_chat_template( messages, add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt", enable_thinking=enable_thinking, ).to(model.device) generate_kwargs = {} if is_edit else _prefix_kwargs(inputs, _t2i_prefix()) started = time.perf_counter() with torch.no_grad(): out = model.generate( **inputs, **generate_kwargs, max_new_tokens=int(max_new_tokens), do_sample=True, temperature=1.0, top_p=0.95, top_k=20, ) elapsed = time.perf_counter() - started generated = processor.tokenizer.decode( out[0, inputs["input_ids"].shape[1] :], skip_special_tokens=True ) n_new = int(out.shape[1] - inputs["input_ids"].shape[1]) try: result = _parse_result(generated) except (ValueError, json.JSONDecodeError) as exc: raise gr.Error(f"Prompt enhancement failed: {exc}") from exc print( f"[enhance] mode={'i2i' if is_edit else 't2i'} images={len(image_paths)} " f"new_tokens={n_new} elapsed={elapsed:.1f}s prefix={'past_key_values' in generate_kwargs} " f"truncated={bool(result.get('truncated'))}", flush=True, ) rewritten = (result.get("rewritten_prompt") or "").strip() if not rewritten: raise gr.Error("Prompt enhancement returned an empty prompt.") return rewritten, _normalize_ratio(result.get("wh_ratio")), generated, seed, round(elapsed, 2) @spaces.GPU(duration=600) def bench(new_tokens: int = 128): import traceback lines = [] def log(msg): lines.append(msg) print(f"[bench] {msg}", flush=True) return "\n".join(lines) model, processor, system_prompt = T2I_MODEL, T2I_PROCESSOR, T2I_SYSTEM_PROMPT messages = [ {"role": "system", "content": [{"type": "text", "text": system_prompt}]}, {"role": "user", "content": [{"type": "text", "text": "a corgi playing guitar in the rain"}]}, ] inputs = processor.apply_chat_template( messages, add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt", enable_thinking=False ).to(model.device) in_len = inputs["input_ids"].shape[1] def run(label, tokens, **kw): try: torch.cuda.synchronize() t0 = time.perf_counter() with torch.no_grad(): out = model.generate(**inputs, max_new_tokens=tokens, min_new_tokens=tokens, do_sample=False, **kw) torch.cuda.synchronize() dt = time.perf_counter() - t0 n = int(out.shape[1] - in_len) return f"{label}: {n} tokens in {dt:.2f}s = {n / dt:.1f} tok/s (prompt {in_len} tokens)" except Exception: return f"{label}: FAILED\n{traceback.format_exc()[-1500:]}" yield log(f"torch {torch.__version__}; T2I model {type(model).__name__}") for i in (1, 2): yield log(run(f"full prompt #{i}", new_tokens)) yield log(run("full prompt, prefill only", 1)) prefix = _t2i_prefix() for i in (1, 2): yield log(run(f"cached system prompt #{i}", new_tokens, **_prefix_kwargs(inputs, prefix))) yield log(run("cached system prompt, prefill only", 1, **_prefix_kwargs(inputs, prefix))) yield log("done") @spaces.GPU(duration=600) def bench_i2i(new_tokens: int = 64): import traceback from pathlib import Path from huggingface_hub import hf_hub_download lines = [] def log(msg): lines.append(msg) print(f"[bench_i2i] {msg}", flush=True) return "\n".join(lines) def triton_files(): root = Path.home() / ".triton" / "cache" return sum(1 for _ in root.rglob("*")) if root.exists() else 0 def timed(fn): torch.cuda.synchronize() t0 = time.perf_counter() with torch.no_grad(): fn() torch.cuda.synchronize() return time.perf_counter() - t0 model, processor, system_prompt = I2I_MODEL, I2I_PROCESSOR, I2I_SYSTEM_PROMPT names = ["examples/49fd6d29-231a-4254-a507-304c1b6de119.webp", "examples/8ffd565d-7915-4353-b78d-cc11111c71da.webp"] paths = [hf_hub_download("hugging-apps/qwen-image-2-1", n, repo_type="space") for n in names] base = [_to_pil(p) for p in paths] cases = [ ("1 image", [base[0]]), ("1 image resized 1000x1400", [base[0].resize((1000, 1400))]), ("1 image resized 1000x1400 again", [base[0].resize((1000, 1400))]), ("2 images", base), ("1 image resized 1200x900", [base[0].resize((1200, 900))]), ] for label, images in cases: try: messages = [ {"role": "system", "content": [{"type": "text", "text": system_prompt}]}, {"role": "user", "content": [{"type": "image", "image": im} for im in images] + [{"type": "text", "text": "make it a poster"}]}, ] inputs = processor.apply_chat_template( messages, add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt", enable_thinking=False ).to(model.device) in_len = inputs["input_ids"].shape[1] grid = inputs["image_grid_thw"] before = triton_files() t_vis1 = timed(lambda: model.model.visual(inputs["pixel_values"], grid_thw=grid)) t_vis2 = timed(lambda: model.model.visual(inputs["pixel_values"], grid_thw=grid)) mid = triton_files() t_pre1 = timed(lambda: model.generate(**inputs, max_new_tokens=1, min_new_tokens=1, do_sample=False)) after1 = triton_files() t_pre2 = timed(lambda: model.generate(**inputs, max_new_tokens=1, min_new_tokens=1, do_sample=False)) t_gen = timed(lambda: model.generate(**inputs, max_new_tokens=new_tokens, min_new_tokens=new_tokens, do_sample=False)) yield log( f"{label}: prompt {in_len} tok, grid {grid.tolist()}; visual {t_vis1:.2f}s/{t_vis2:.2f}s " f"(triton files +{mid - before}); prefill {t_pre1:.2f}s/{t_pre2:.2f}s (triton files +{after1 - mid}); " f"{new_tokens} tokens {t_gen:.2f}s" ) except Exception: yield log(f"{label}: FAILED\n" + traceback.format_exc()[-1500:]) yield log("done") with gr.Blocks(title="Qwen-Image-2.1 Prompt Enhancer") as demo: gr.Markdown( "# Qwen-Image-2.1 Prompt Enhancer\n" "The official prompt-rewriting models for **Qwen-Image-2.1**, kept warm on one GPU worker. " "Leave the image list empty for text-to-image rewriting " "(`Qwen-Image-2.1-PE-T2I`); add images for editing rewriting (`Qwen-Image-2.1-PE-I2I`)." ) with gr.Row(): with gr.Column(): prompt = gr.Textbox( label="Prompt", lines=4, placeholder="A short brief in any language — the model expands it.", ) image_paths = gr.File( label="Input images (optional, up to 10)", file_count="multiple", file_types=["image"], type="filepath", ) run_button = gr.Button("Enhance prompt", variant="primary") with gr.Column(): rewritten_out = gr.Textbox(label="Rewritten prompt", lines=12) ratio_out = gr.Textbox(label="Recommended aspect ratio") seed_out = gr.Number(label="Seed used") elapsed_out = gr.Number(label="Elapsed (s)") with gr.Accordion("Advanced settings", open=False): max_new_tokens = gr.Slider( label="Max new tokens", minimum=256, maximum=8192, step=64, value=DEFAULT_MAX_NEW_TOKENS, info="Decoding cap. Raise it well above 4096 before enabling thinking.", ) enable_thinking = gr.Checkbox( label="Enable thinking", value=False, info="Higher quality, several times slower — raise the token cap with it.", ) seed = gr.Slider(label="Seed", minimum=0, maximum=MAX_SEED, step=1, value=0) randomize_seed = gr.Checkbox(label="Randomize seed", value=True) raw_out = gr.Textbox(label="Raw model output", lines=8) gr.Examples( examples=[ ["一只在雨中弹吉他的柯基"], ["a neon shop sign that reads QWEN IMAGE 2.1, rainy night"], ["minimalist poster for a jazz festival"], ], inputs=[prompt], fn=enhance, outputs=[rewritten_out, ratio_out, raw_out, seed_out, elapsed_out], cache_examples=True, cache_mode="lazy", ) with gr.Accordion("Benchmark (operator)", open=False): bench_tokens = gr.Slider(label="Tokens per run", minimum=32, maximum=512, step=32, value=128) bench_button = gr.Button("Run decode benchmark") bench_out = gr.Textbox(label="Benchmark log", lines=16) bench_button.click(fn=bench, inputs=[bench_tokens], outputs=[bench_out], api_name="bench") bench_i2i_button = gr.Button("Run i2i benchmark") bench_i2i_button.click(fn=bench_i2i, inputs=[bench_tokens], outputs=[bench_out], api_name="bench_i2i") run_button.click( fn=enhance, inputs=[prompt, image_paths, max_new_tokens, enable_thinking, seed, randomize_seed], outputs=[rewritten_out, ratio_out, raw_out, seed_out, elapsed_out], api_name="enhance", ) demo.queue(default_concurrency_limit=1, max_size=16) if __name__ == "__main__": demo.launch(mcp_server=True, theme=gr.themes.Citrus())