multimodalart's picture
multimodalart HF Staff
i2i benchmark: vision vs prefill, triton cache growth
19b6fb4 verified
Raw History Blame Contribute Delete
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))
@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())