akhaliq's picture
akhaliq HF Staff
Media example chips use public URLs from vl-test-suite
7f3ba91
Raw
History Blame Contribute Delete
2.98 kB
import os
from threading import Thread
import spaces
import torch
from transformers import AutoProcessor, AutoModelForMultimodalLM, TextIteratorStreamer
from gradio import Server
from gradio.data_classes import FileData
from fastapi.responses import HTMLResponse
MODEL_ID = "meta-models/Muse-Glimmer-30B"
processor = AutoProcessor.from_pretrained(MODEL_ID)
model = AutoModelForMultimodalLM.from_pretrained(
MODEL_ID,
dtype=torch.bfloat16,
).to("cuda")
app = Server()
@app.api(name="chat")
@spaces.GPU(size="xlarge", duration=240)
def chat(message: str,
media: FileData | None,
media_url: str,
history: list,
reasoning: str) -> str:
"""Streaming multimodal chat with Muse-Glimmer-30B.
history: list of {"role": "user"|"assistant", "content": str}
media: optional image or video FileData uploaded via the Gradio client
media_url: optional public URL of an image or video
"""
messages = []
for m in history:
messages.append({"role": m["role"], "content": m["content"]})
content = []
if media:
path = media["path"]
if path.lower().rsplit(".", 1)[-1] in ("mp4", "mov", "webm", "mkv", "avi"):
content.append({"type": "video", "path": path})
else:
content.append({"type": "image", "path": path})
elif media_url:
if media_url.lower().rsplit(".", 1)[-1].split("?")[0] in ("mp4", "mov", "webm", "mkv", "avi"):
content.append({"type": "video", "url": media_url})
else:
content.append({"type": "image", "url": media_url})
if message:
content.append({"type": "text", "text": message})
messages.append({"role": "user", "content": content})
template_kwargs = dict(
add_generation_prompt=True,
tokenize=True,
return_dict=True,
return_tensors="pt",
reasoning_strength=reasoning,
)
if any(c["type"] == "video" for c in content):
template_kwargs["processor_kwargs"] = {"num_frames": 32}
inputs = processor.apply_chat_template(messages, **template_kwargs).to(model.device)
streamer = TextIteratorStreamer(processor, skip_prompt=True, skip_special_tokens=True)
# Use the model's generation_config defaults (recommended sampling
# params ship with the model); only cap length for ZeroGPU runtime.
thread = Thread(
target=model.generate,
kwargs=dict(
**inputs,
streamer=streamer,
max_new_tokens=1024,
),
)
thread.start()
reply = ""
for token in streamer:
reply += token
yield reply # raw stream; frontend splits reasoning from final reply
thread.join()
@app.get("/", response_class=HTMLResponse)
async def homepage():
html_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "index.html")
with open(html_path, "r", encoding="utf-8") as f:
return f.read()
app.launch(show_error=True)