akhaliq's picture
akhaliq HF Staff
Media example chips use public URLs from vl-test-suite
7f3ba91
Raw
History Blame
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)