akhaliq's picture
akhaliq HF Staff
Prefill reference values via node data, wire condition image, handle canvas file URLs in edit_image
4eac3f3
Raw History Blame
4.03 kB
"""Qwen-Image 2.1 — Gradio Workflow on ZeroGPU.
A node-based canvas (gr.Workflow) exposing the diffusers QwenImage21Pipeline:
- text_to_image: prompt -> image
- edit_image: condition image + instruction -> edited image (chained after
text_to_image, or fed from an uploaded image reference)
The pipeline is placed on `cuda` at module level: outside @spaces.GPU functions
PyTorch runs in CUDA emulation mode, and a real ZeroGPU is attached only while
a decorated function executes.
"""
import base64
import os
import urllib.parse
import urllib.request
import gradio as gr
from gradio_client import utils as client_utils
from gradio.utils import get_upload_folder
import spaces
import torch
from PIL import Image
from diffusers import QwenImage21Pipeline
MODEL_ID = os.environ.get("QWEN_IMAGE_MODEL", "Qwen/Qwen-Image-2.1")
# The model repo is private/gated. Per the Hub auth docs, the HF_TOKEN
# environment variable is used implicitly for all Hub requests and takes
# priority over any stored token — so set HF_TOKEN as a Space secret
# (Settings -> Secrets) with a token whose account has access to the repo.
if not os.environ.get("HF_TOKEN"):
raise RuntimeError(
"HF_TOKEN is not set. Add it as a Space secret (Settings -> Secrets) "
f"with read access to {MODEL_ID}."
)
pipe = QwenImage21Pipeline.from_pretrained(MODEL_ID, dtype=torch.bfloat16)
pipe.to("cuda")
def _to_pil(image) -> Image.Image:
"""Accept whatever the canvas hands a bound function for an image port:
a PIL image, a local path, a /gradio_api/file= reference, an http(s) or
data: URL, or a file dict carrying any of those. Mirrors _file_ref in
gradio.workflow — canvas file values carry only `url`, no `path`."""
import io
if isinstance(image, Image.Image):
return image
if isinstance(image, dict):
image = image.get("path") or image.get("url") or ""
if not isinstance(image, str) or not image:
raise ValueError(f"Unsupported image input: {type(image)!r}")
if image.startswith("/gradio_api/file="):
image = urllib.parse.unquote(image.removeprefix("/gradio_api/file="))
if image.startswith("data:"):
return Image.open(io.BytesIO(base64.b64decode(image.split(",", 1)[1])))
if image.startswith(("http://", "https://")):
return Image.open(io.BytesIO(urllib.request.urlopen(image).read()))
return Image.open(image)
def _save(image: Image.Image) -> dict:
# Mirror gradio.workflow._save_tmp: the canvas renders media values only
# from {path, url, is_file} dicts whose url is a /gradio_api/file= link,
# and the file must live under the upload folder to be servable.
directory = get_upload_folder()
os.makedirs(directory, exist_ok=True)
path = os.path.join(directory, f"workflow_{os.urandom(8).hex()}.png")
image.save(path)
url = f"/gradio_api/file={client_utils.encode_file_path(path)}"
return {"path": path, "url": url, "is_file": True}
def _steps(value, default: int = 40) -> int:
# Unconnected number ports arrive as None; function defaults are not
# applied because the workflow executor passes arguments positionally.
return default if value is None else int(value)
@spaces.GPU(duration=120)
def text_to_image(prompt: str, steps: int = 40) -> dict:
"""Generate an image from a text prompt with Qwen-Image 2.1."""
image = pipe(prompt, num_inference_steps=_steps(steps)).images[0]
return _save(image)
@spaces.GPU(duration=120)
def edit_image(image, instruction: str, steps: int = 40) -> dict:
"""Edit a condition image following an instruction (image-conditioned
generation with Qwen-Image 2.1)."""
edited = pipe(instruction, image=_to_pil(image), num_inference_steps=_steps(steps)).images[0]
return _save(edited)
demo = gr.Workflow(
graph=os.path.join(os.path.dirname(__file__), "workflow.json"),
bind={"text_to_image": text_to_image, "edit_image": edit_image},
)
if __name__ == "__main__":
demo.launch()