omni-diffusion / app.py
multimodalart's picture
multimodalart HF Staff
Upload app.py with huggingface_hub
28831b8 verified
Raw
History Blame Contribute Delete
14.8 kB
import os
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
import spaces # noqa: E402 must precede torch / CUDA-touching imports
import re # noqa: E402
import random # noqa: E402
import tempfile # noqa: E402
import numpy as np # noqa: E402
import torch # noqa: E402
import gradio as gr # noqa: E402
from PIL import Image # noqa: E402
from torchvision import transforms # noqa: E402
from transformers import AutoConfig, AutoModel, AutoTokenizer # noqa: E402
from transformers.generation import GenerationConfig # noqa: E402
from load_magvit import load_magvit # noqa: E402
MODEL_ID = "lijiang/Omni-Diffusion"
IMAGE_TOKENIZER_ID = "showlab/magvitv2"
DTYPE = torch.bfloat16
DEVICE = "cuda"
# Qwen2 chat template used by the authors' reference inference.
QWEN2_CHAT_TEMPLATE = (
"{%- if messages[0]['role'] == 'system' %}"
"{{- '<|im_start|>system\\n' + messages[0]['content'] + '<|im_end|>\\n' }}"
"{%- endif %}"
"{%- for message in messages %}"
"{%- if (message.role == \"user\") or (message.role == \"system\" and not loop.first) "
"or (message.role == \"assistant\") %}"
"{{- '<|im_start|>' + message.role + '\\n' + message.content + '<|im_end|>' + '\\n' }}"
"{%- endif %}"
"{%- endfor %}"
"{%- if add_generation_prompt %}"
"{{- '<|im_start|>assistant\\n' }}"
"{%- endif %}"
)
IMAGENET_MEAN = [0.485, 0.456, 0.406]
IMAGENET_STD = [0.229, 0.224, 0.225]
# ---------------------------------------------------------------------------
# Model loading (module scope, eager .to("cuda") per ZeroGPU rules)
# ---------------------------------------------------------------------------
print("Loading tokenizer ...", flush=True)
tokenizer = AutoTokenizer.from_pretrained(
MODEL_ID, trust_remote_code=True, chat_template=QWEN2_CHAT_TEMPLATE
)
print("Loading Omni-Diffusion (Dream) model ...", flush=True)
config = AutoConfig.from_pretrained(MODEL_ID, trust_remote_code=True)
model = AutoModel.from_pretrained(
MODEL_ID,
trust_remote_code=True,
torch_dtype=DTYPE,
attn_implementation="sdpa",
).eval().to(DEVICE)
gen_cfg = GenerationConfig.from_pretrained(MODEL_ID, trust_remote_code=True)
gen_cfg.pad_token_id = tokenizer.pad_token_id
model.generation_config = gen_cfg
print("Loading MagViT-v2 image detokenizer ...", flush=True)
image_tokenizer = load_magvit(IMAGE_TOKENIZER_ID).to(DEVICE)
AUDIO_OFFSET = tokenizer.convert_tokens_to_ids("<|audio_0|>")
IMAGE_OFFSET = tokenizer.convert_tokens_to_ids("<|image_0|>")
print(f"AUDIO_OFFSET={AUDIO_OFFSET} IMAGE_OFFSET={IMAGE_OFFSET}", flush=True)
def _image_transform(image: Image.Image, resolution: int = 512) -> torch.Tensor:
image = transforms.Resize(
resolution, interpolation=transforms.InterpolationMode.BICUBIC
)(image)
image = transforms.CenterCrop((resolution, resolution))(image)
image = transforms.ToTensor()(image)
image = transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])(image)
return image
def _encode_image_to_tokens(pil_image: Image.Image) -> str:
"""Encode a PIL image into the model's <|image_i|> token string."""
pixel = _image_transform(pil_image.convert("RGB"), 512).unsqueeze(0)
pixel = pixel.to(DEVICE, dtype=next(image_tokenizer.parameters()).dtype)
codes = image_tokenizer.get_code(pixel)[0].tolist()
return "".join(f"<|image_{i}|>" for i in codes)
def _decode_tokens_to_image(image_tokens):
if len(image_tokens) == 0:
return None
if len(image_tokens) < 256:
image_tokens = image_tokens + [image_tokens[-1]] * (256 - len(image_tokens))
gen_token_ids = torch.tensor(image_tokens[:256], device=DEVICE).unsqueeze(0)
gen_token_ids = torch.clamp(gen_token_ids, max=8192 - 1, min=0)
with torch.no_grad():
image = image_tokenizer.decode_code(gen_token_ids)
image = torch.clamp((image + 1.0) / 2.0, min=0.0, max=1.0) * 255.0
image = image.permute(0, 2, 3, 1).cpu().float().numpy().astype(np.uint8)[0]
return Image.fromarray(image)
def _build_input_ids(message, image_pil=None, system=None):
messages = []
if system:
messages.append({"role": "system", "content": system})
content = message
if image_pil is not None:
img_tokens = _encode_image_to_tokens(image_pil)
content = content + "\n<|begin_of_image|>" + img_tokens + "<|end_of_image|>"
messages.append({"role": "user", "content": content})
input_ids = tokenizer.apply_chat_template(
messages, tokenize=True, add_generation_prompt=True
)
return torch.tensor([input_ids], dtype=torch.long, device=DEVICE)
def _split_output_tokens(out_ids):
audio_tokens, image_tokens, text_tokens = [], [], []
for tid in out_ids:
tid = int(tid)
if AUDIO_OFFSET <= tid < AUDIO_OFFSET + 16384:
audio_tokens.append(tid - AUDIO_OFFSET)
elif tid >= IMAGE_OFFSET:
image_tokens.append(tid - IMAGE_OFFSET)
else:
text_tokens.append(tid)
return audio_tokens, image_tokens, text_tokens
def _clean_text(text_tokens):
text = tokenizer.decode(text_tokens, skip_special_tokens=True)
# Strip any residual chat-template markers that survive decoding.
text = re.sub(r"<\|.*?\|>", "", text)
return text.strip()
# ---------------------------------------------------------------------------
# Inference handlers
# ---------------------------------------------------------------------------
@spaces.GPU(duration=180)
def text_to_image(prompt: str, steps: int = 260, cfg: float = 0.0, seed: int = 42):
"""Generate an image from a text prompt using masked discrete diffusion.
Args:
prompt: text description of the image to generate.
steps: number of diffusion denoising steps.
cfg: classifier-free guidance scale (0 disables CFG).
seed: RNG seed for reproducibility.
"""
if not prompt or not prompt.strip():
raise gr.Error("Please enter a text prompt.")
steps = int(steps)
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
message = "Generate an image based on the provided text description.\n" + prompt
input_ids = _build_input_ids(message)
outputs, _ = model.generate(
input_ids,
max_new_tokens=260,
steps=steps,
temperature=0.0,
top_p=0.9,
alg="entropy-penalty",
cfg=float(cfg),
tokenizer=tokenizer,
max_position_penalty=2.0,
repeat_penalty=1.2,
task="T2I",
)
_, image_tokens, _ = _split_output_tokens(outputs[0][input_ids.shape[1]:])
image = _decode_tokens_to_image(image_tokens)
if image is None:
raise gr.Error("The model did not produce image tokens. Try a different prompt or more steps.")
return image
@spaces.GPU(duration=120)
def visual_qa(image, question: str, steps: int = 64, max_tokens: int = 128):
"""Answer a question about an image (visual understanding).
Args:
image: input image.
question: question or instruction about the image.
steps: number of diffusion denoising steps.
max_tokens: maximum number of new tokens to generate.
"""
if image is None:
raise gr.Error("Please upload an image.")
if not question or not question.strip():
question = "Describe this image in detail."
steps = int(steps)
max_tokens = int(max_tokens)
input_ids = _build_input_ids(question, image_pil=image)
outputs, _ = model.generate(
input_ids,
max_new_tokens=max_tokens,
steps=steps,
temperature=0.0,
top_p=0.9,
alg="entropy",
cfg=0.0,
tokenizer=tokenizer,
max_position_penalty=1.0,
repeat_penalty=1.0,
task="VQA",
)
_, _, text_tokens = _split_output_tokens(outputs[0][input_ids.shape[1]:])
text = _clean_text(text_tokens)
return text or "(no answer produced)"
@spaces.GPU(duration=120)
def text_chat(message: str, steps: int = 64, max_tokens: int = 256):
"""Generate a text response to a prompt.
Args:
message: the user prompt / instruction.
steps: number of diffusion denoising steps.
max_tokens: maximum number of new tokens to generate.
"""
if not message or not message.strip():
raise gr.Error("Please enter a prompt.")
steps = int(steps)
max_tokens = int(max_tokens)
input_ids = _build_input_ids(message)
outputs, _ = model.generate(
input_ids,
max_new_tokens=max_tokens,
steps=steps,
temperature=0.0,
top_p=0.9,
alg="entropy",
cfg=0.0,
tokenizer=tokenizer,
max_position_penalty=1.0,
repeat_penalty=1.0,
task="chat",
)
_, _, text_tokens = _split_output_tokens(outputs[0][input_ids.shape[1]:])
text = _clean_text(text_tokens)
return text or "(no response produced)"
# ---------------------------------------------------------------------------
# UI
# ---------------------------------------------------------------------------
CSS = """
#col-container { max-width: 1100px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
"""
INTRO = """
# 🌀 Omni-Diffusion
Unified multimodal **understanding and generation** with a **masked discrete diffusion**
language model ([paper](https://huggingface.co/papers/2603.06577) ·
[model](https://huggingface.co/lijiang/Omni-Diffusion) ·
[code](https://github.com/VITA-MLLM/Omni-Diffusion)).
One model jointly models discrete tokens of text and images. This demo exposes its
**text-to-image**, **visual question answering**, and **text** capabilities.
"""
with gr.Blocks(theme=gr.themes.Citrus(), css=CSS) as demo:
with gr.Column(elem_id="col-container"):
gr.Markdown(INTRO)
with gr.Tabs():
# ---- Text to Image ----
with gr.Tab("Text → Image"):
with gr.Row():
t2i_prompt = gr.Textbox(
show_label=False,
placeholder="A landscape background with double exposure glasses of wine…",
container=False,
scale=4,
)
t2i_btn = gr.Button("Generate", variant="primary", scale=1)
t2i_out = gr.Image(label="Generated image", type="pil", height=384)
with gr.Accordion("Advanced settings", open=False):
t2i_steps = gr.Slider(32, 512, value=260, step=1, label="Diffusion steps")
t2i_cfg = gr.Slider(0.0, 5.0, value=0.0, step=0.5, label="CFG scale")
t2i_seed = gr.Number(value=42, precision=0, label="Seed")
gr.Examples(
examples=[
["The image shows a landscape background with double exposure glasses of wine, displaying a hyperealistic and detailed view of the subject."],
["A group of 1920s girls at college immersed in their studies at a dark academia university."],
["A super realistic and hyper-detailed 8k image of a fantasy night scene with a beach under the full moon."],
],
inputs=[t2i_prompt],
outputs=t2i_out,
fn=text_to_image,
cache_examples=False,
run_on_click=True,
)
t2i_btn.click(
text_to_image,
inputs=[t2i_prompt, t2i_steps, t2i_cfg, t2i_seed],
outputs=t2i_out,
api_name="text_to_image",
)
# ---- Visual QA ----
with gr.Tab("Image → Text (VQA)"):
with gr.Row():
vqa_image = gr.Image(label="Input image", type="pil", height=320)
with gr.Column():
vqa_question = gr.Textbox(
label="Question",
placeholder="Is the glass of orange juice half empty or half full?",
)
vqa_btn = gr.Button("Ask", variant="primary")
vqa_out = gr.Textbox(label="Answer", lines=4)
with gr.Accordion("Advanced settings", open=False):
vqa_steps = gr.Slider(16, 256, value=64, step=1, label="Diffusion steps")
vqa_max = gr.Slider(16, 512, value=128, step=1, label="Max new tokens")
gr.Examples(
examples=[
["examples/vqa_0.png", "Is the glass of orange juice half empty or half full?"],
["examples/svqa_0.jpg", "What is happening in this image?"],
],
inputs=[vqa_image, vqa_question],
outputs=vqa_out,
fn=visual_qa,
cache_examples=False,
run_on_click=True,
)
vqa_btn.click(
visual_qa,
inputs=[vqa_image, vqa_question, vqa_steps, vqa_max],
outputs=vqa_out,
api_name="visual_qa",
)
# ---- Text ----
with gr.Tab("Text → Text"):
with gr.Row():
chat_in = gr.Textbox(
show_label=False,
placeholder="Ask anything…",
container=False,
scale=4,
)
chat_btn = gr.Button("Send", variant="primary", scale=1)
chat_out = gr.Textbox(label="Response", lines=6)
with gr.Accordion("Advanced settings", open=False):
chat_steps = gr.Slider(16, 256, value=64, step=1, label="Diffusion steps")
chat_max = gr.Slider(16, 512, value=256, step=1, label="Max new tokens")
gr.Examples(
examples=[
["What is the capital of France?"],
["Write a short poem about the ocean."],
],
inputs=[chat_in],
outputs=chat_out,
fn=text_chat,
cache_examples=False,
run_on_click=True,
)
chat_btn.click(
text_chat,
inputs=[chat_in, chat_steps, chat_max],
outputs=chat_out,
api_name="text_chat",
)
if __name__ == "__main__":
demo.queue().launch(mcp_server=True)