Spaces:
Running on Zero
Running on Zero
Download app.py from hugging-apps/qwen-image-2-1-prompt-enhancer: direct link, hf CLI and curl.
- Browser
- Download file 18.6 kB
-
https://huggingface.co/spaces/hugging-apps/qwen-image-2-1-prompt-enhancer/resolve/main/app.py
- Command line
-
hf download hf://spaces/hugging-apps/qwen-image-2-1-prompt-enhancer/app.py
-
curl -L -o app.py https://huggingface.co/spaces/hugging-apps/qwen-image-2-1-prompt-enhancer/resolve/main/app.py
18.6 kB
| """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'(?<!\\)"\s*(,|\})', body) | |
| if closing: | |
| body = body[: closing.start()] | |
| try: | |
| rewritten = json.loads(f'"{body}"') | |
| except json.JSONDecodeError: | |
| rewritten = body.replace('\\"', '"').replace("\\n", " ") | |
| rewritten = rewritten.strip() | |
| return {"rewritten_prompt": rewritten, "wh_ratio": "", "truncated": True} if rewritten else None | |
| def _parse_result(generated: str) -> dict: | |
| """Pull the JSON answer out of a rewriter's ``<think>…</think>{json}`` output.""" | |
| _thinking, _sep, answer = generated.partition("</think>") | |
| 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)) | |
| 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) | |
| 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") | |
| 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()) | |