Spaces:
Running on Zero
Running on Zero
File size: 15,154 Bytes
f3fc14a db0cf7d f3fc14a db0cf7d f3fc14a 2b2ef11 f3fc14a 2b2ef11 f3fc14a 1dc2ce2 f3fc14a db0cf7d f3fc14a 186564b f3fc14a 1dc2ce2 f3fc14a db0cf7d f3fc14a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 | import os
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
# The runtime's torchvision (>=0.28) removed `torchvision.io.read_video`, which
# qwen-vl-utils 0.0.14 needs; pin the dispatch and swap in a PyAV decoder below.
os.environ.setdefault("FORCE_QWENVL_VIDEO_READER", "torchvision")
import spaces # noqa: E402 (must precede any torch import)
import re # noqa: E402
import time # noqa: E402
import html # noqa: E402
import torch # noqa: E402
import gradio as gr # noqa: E402
from transformers import AutoProcessor, Qwen3VLForConditionalGeneration # noqa: E402
from qwen_vl_utils import process_vision_info # noqa: E402
import qwen_vl_utils.vision_process as _vp # noqa: E402
import video_av_reader # noqa: E402
# PyAV bundles its own FFmpeg (no system libs needed) and mirrors the official
# torchvision backend 1:1 (smart_nframes + linspace indices + same metadata).
_vp.VIDEO_READER_BACKENDS["torchvision"] = video_av_reader.read_video_pyav
MODEL_ID = "TengfeiLiuCoder/RefCaptioner"
MAX_REFS = 6
# ---- Inference configuration: the authors' released recipe (Prompt_1.0) ----
CFG = {
"max_length": 18000,
"max_new_tokens": 512,
"video_fps": 2.0,
"video_min_frames": 4,
"video_max_frames": 20,
"image_max_pixels": 602112,
"video_max_pixels": 602112,
}
INSTRUCTION = """You are a multi-reference video captioning model.
Input:
- Several reference images, each labeled as <Image_1>, <Image_2>, etc.
- Some reference images may be distractors that are not visible in the video.
- One reference video.
Task:
Write one fluent English video caption that describes the visible video content and locally binds each usable reference image tag to the visual phrase it grounds.
Output rules:
- Return only the caption: no markdown, bullets, JSON, explanations, or title.
- Write one natural English paragraph, usually 4 to 7 complete sentences and about 120 to 250 words.
- Cover the main video style or format, main subject, setting, referenced appearances or objects, action progression, and useful camera, lighting, color, or mood details.
- Keep the caption grounded in visible video evidence. Do not invent unseen names, relationships, causes, dialogue, or story details.
- Do not mention audio, music, speech, dialogue, transcript, voiceover, or sound. Mention visible subtitles, logos, or text only when visually important.
Reference binding rules:
- Use only the provided tags. Do not invent tags.
- Use a tag only when its reference image can be grounded to visible video content.
- Do not force every provided tag. Omit tags whose image is not visible in the video or cannot be confidently grounded.
- Place each used tag immediately after a concrete grounded phrase, such as "the woman <Image_1>" or "the red dress <Image_2>".
- Multiple tags may be stacked as one contiguous tag group only when they refer to the same concrete visual unit, such as the same person, animal, character, vehicle, room, landscape, outfit item, prop, action pose, lighting, mood, or visual style.
- Do not stack tags merely because the images are related. A person and their clothing, accessory, carried object, background, or action are usually different visual units and should be tagged on separate phrases.
- If several tags belong to the same visual unit, keep them together as a complete tag group whenever that unit is explicitly mentioned.
- If two tags need different phrases, do not stack them.
- Attach tags to explicit noun phrases, not pronouns such as "he", "she", "it", "they", or "them".
- Do not put all tags at the end of the caption.
- Do not write phrases like "from <Image_N>", "shown in <Image_N>", "as shown in <Image_N>", "same as <Image_N>", or "similar to <Image_N>"."""
FINAL_REQUEST = (
"Now write the final caption. Use only visibly grounded provided tags exactly "
"as tag tokens, attach them to concrete visual phrases, omit ungrounded "
"distractor tags, and stack tags only when they refer to the same visual unit."
)
# ---------------------------------------------------------------------------
# Model load (module scope, eager .to("cuda") per ZeroGPU rules)
# ---------------------------------------------------------------------------
print(f"Loading {MODEL_ID} ...", flush=True)
processor = AutoProcessor.from_pretrained(MODEL_ID)
model = (
Qwen3VLForConditionalGeneration.from_pretrained(
MODEL_ID,
dtype=torch.bfloat16,
attn_implementation="sdpa",
)
.eval()
.to("cuda")
)
print("Model loaded.", flush=True)
# ---------------------------------------------------------------------------
# Prompt construction — mirrors the official inference.py message layout 1:1
# ---------------------------------------------------------------------------
def build_messages(video_path: str, image_paths: list[str]) -> list[dict]:
tags = [f"<Image_{i}>" for i in range(1, len(image_paths) + 1)]
content: list[dict] = [
{"type": "text", "text": INSTRUCTION},
{
"type": "text",
"text": (
"\n\nCurrent sample starts here.\nReference image tags, in order: "
+ ", ".join(tags)
+ "\nEach following image is the visual reference for the tag immediately before it.\n\n"
),
},
]
for tag, image in zip(tags, image_paths):
content.extend(
[
{"type": "text", "text": f"{tag} reference image:\n"},
{"type": "image", "image": image, "max_pixels": CFG["image_max_pixels"]},
{"type": "text", "text": "\n\n"},
]
)
content.extend(
[
{"type": "text", "text": "Reference video:\n"},
{
"type": "video",
"video": video_path,
"fps": CFG["video_fps"],
"min_frames": CFG["video_min_frames"],
"max_frames": CFG["video_max_frames"],
"max_pixels": CFG["video_max_pixels"],
},
{"type": "text", "text": "\n\n" + FINAL_REQUEST},
]
)
return [{"role": "user", "content": content}]
def highlight_tags(caption: str) -> str:
"""Render <Image_n> tags as highlighted chips, preserving everything else."""
def chip(m: re.Match) -> str:
return (
f'<span style="background:rgba(99,102,241,.18);border:1px solid rgba(99,102,241,.55);'
f'border-radius:6px;padding:0 6px;margin:0 2px;font-family:monospace;'
f'font-size:.92em;color:var(--body-text-color);white-space:nowrap;">{m.group(0)}</span>'
)
# Escape first, then chip the escaped tag form (<Image_N>) so the rest
# of the caption stays safely escaped while tags render as styled chips.
return re.sub(r"<(Image_\d+)>", chip, html.escape(caption))
# ---------------------------------------------------------------------------
# Inference
# ---------------------------------------------------------------------------
@spaces.GPU(duration=30)
def generate_caption(
video,
ref_1,
ref_2=None,
ref_3=None,
ref_4=None,
ref_5=None,
ref_6=None,
progress=gr.Progress(track_tqdm=True),
):
"""Generate a reference-grounded caption for a video.
Args:
video: input video (file path).
ref_1..ref_6: up to 6 reference images, in tag order (<Image_1> .. <Image_6>).
progress: Gradio progress bar (injected).
Returns:
HTML caption with highlighted <Image_n> tags, plus plain text and timing info.
"""
if not video:
raise gr.Error("Please upload a video first.")
refs = [p for p in [ref_1, ref_2, ref_3, ref_4, ref_5, ref_6] if p]
if not refs:
raise gr.Error("Please upload at least one reference image.")
started = time.perf_counter()
messages = build_messages(video, refs)
text = processor.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True, enable_thinking=False
)
image_inputs, video_inputs, video_kwargs = process_vision_info(
messages, return_video_kwargs=True
)
if isinstance(video_kwargs, dict):
video_kwargs = {
k: v
for k, v in video_kwargs.items()
if not (isinstance(v, list) and not v)
}
if isinstance(video_kwargs.get("fps"), list) and video_kwargs["fps"]:
video_kwargs["fps"] = video_kwargs["fps"][0]
inputs = processor(
text=[text],
images=image_inputs,
videos=video_inputs,
padding=True,
return_tensors="pt",
**video_kwargs,
)
n_input_tokens = int(inputs["input_ids"].shape[-1])
if n_input_tokens > CFG["max_length"]:
raise gr.Error(
f"Input too long ({n_input_tokens} tokens > {CFG['max_length']}). "
"Try a shorter video or fewer reference images."
)
inputs = inputs.to(model.device)
with torch.inference_mode():
generated = model.generate(
**inputs,
max_new_tokens=CFG["max_new_tokens"],
do_sample=False,
)
trimmed = [out[len(src):] for src, out in zip(inputs["input_ids"], generated)]
caption = processor.batch_decode(
trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False
)[0].strip()
elapsed = time.perf_counter() - started
used = sorted(set(re.findall(r"<Image_\d+>", caption)))
omitted = [f"<Image_{i}>" for i in range(1, len(refs) + 1) if f"<Image_{i}>" not in caption]
stats = (
f"⏱ {elapsed:.1f} s on GPU · 🧮 {n_input_tokens} input tokens · "
f"✅ grounded: {' '.join(used) if used else 'none'} · "
f"🚫 omitted: {' '.join(omitted) if omitted else 'none'}"
)
return highlight_tags(caption), caption, stats
# ---------------------------------------------------------------------------
# UI
# ---------------------------------------------------------------------------
CSS = """
#col-container { max-width: 1120px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
.footer a { color: var(--body-text-color); }
"""
with gr.Blocks(theme=gr.themes.Citrus(), css=CSS) as demo:
with gr.Column(elem_id="col-container"):
gr.Markdown(
"""
# 🎬 RefCaptioner: Multi-Reference Image-Grounded Video Captioning
Upload a **video** and **reference images** — RefCaptioner writes a detailed English caption and
binds each reference to the visual phrase it grounds with `<Image_n>` tags. Distractor references
(not visible in the video) are omitted. Model: [`TengfeiLiuCoder/RefCaptioner`](https://huggingface.co/TengfeiLiuCoder/RefCaptioner) ·
[paper](https://arxiv.org/abs/2607.28509) · [code](https://github.com/pkucs-Ltf/RefCaptioner)
"""
)
with gr.Row():
with gr.Column(scale=1):
video_in = gr.Video(label="🎥 Video", sources=["upload"])
with gr.Accordion("🖼️ Reference images (tag order = upload order)", open=True):
ref_components = []
with gr.Row():
r1 = gr.Image(label="<Image_1>", type="filepath", sources=["upload"])
r2 = gr.Image(label="<Image_2>", type="filepath", sources=["upload"])
ref_components += [r1, r2]
with gr.Row():
r3 = gr.Image(label="<Image_3>", type="filepath", sources=["upload"])
r4 = gr.Image(label="<Image_4>", type="filepath", sources=["upload"])
ref_components += [r3, r4]
with gr.Row():
r5 = gr.Image(label="<Image_5>", type="filepath", sources=["upload"])
r6 = gr.Image(label="<Image_6>", type="filepath", sources=["upload"])
ref_components += [r5, r6]
run_btn = gr.Button("Generate grounded caption", variant="primary")
with gr.Column(scale=1):
caption_html = gr.HTML(
label="Grounded caption (tags highlighted)",
value="<div style='opacity:.55;padding:8px 4px;'>Your caption will appear here…</div>",
)
caption_text = gr.Textbox(label="Plain-text caption", lines=10)
stats_md = gr.Markdown("")
with gr.Accordion("ℹ️ How it works", open=False):
gr.Markdown(
"""
- **Task** — describe the video *and* make phrase↔reference correspondences explicit:
each usable reference tag is placed immediately after the concrete visual phrase it grounds.
- **Distractors** — references whose content is absent from the video are deliberately
left unused, demonstrating reference selection.
- **Recipe** — this demo reproduces the authors' released inference configuration
(`Prompt_1.0`: greedy decoding, 2 FPS video sampling, ≤20 frames, 602112 max pixels
per image/frame, thinking disabled).
- **Base model** — fine-tuned from [Qwen3-VL-8B-Instruct](https://huggingface.co/Q/Qwen3-VL-8B-Instruct);
benchmarked on [MRVBench](https://huggingface.co/datasets/TengfeiLiuCoder/MRVBench).
"""
)
gr.Examples(
examples=[
# Border collie running through a field. <Image_1> is a frame of the
# same dog (grounded); <Image_2> is a husky (distractor, omitted).
["examples/border_collie_running.mp4", "examples/collie_frame.jpg", "examples/husky_dog.jpg"],
# Man jumping rope on a sidewalk. <Image_1> is a frame of the same
# man (grounded); <Image_2> is a night skyline (distractor, omitted).
["examples/man_jumping_rope.mp4", "examples/jump_rope_frame.jpg", "examples/city_skyline_night.jpg"],
],
# Examples fill video + <Image_1> + <Image_2>; ref_3..ref_6 stay at their defaults.
inputs=[video_in, r1, r2],
outputs=[caption_html, caption_text, stats_md],
fn=generate_caption,
cache_examples=True,
cache_mode="lazy",
)
gr.Markdown(
"<div class='footer' style='opacity:.7;font-size:.85em;'>Example media: "
"<a href='https://huggingface.co/datasets/linoyts/repo-to-space-example-videos'>linoyts/repo-to-space-example-videos</a> "
"and <a href='https://huggingface.co/datasets/linoyts/repo-to-space-example-inputs'>linoyts/repo-to-space-example-inputs</a> "
"(free-to-use stock media).</div>"
)
run_btn.click(
fn=generate_caption,
inputs=[video_in, *ref_components],
outputs=[caption_html, caption_text, stats_md],
api_name="generate_caption",
concurrency_limit=1,
)
if __name__ == "__main__":
demo.launch(mcp_server=True) |