refcaptioner / app.py
multimodalart's picture
multimodalart HF Staff
Fix tag highlighting: chip the escaped tag form after html.escape
2b2ef11 verified
Raw History Blame Contribute Delete
15.2 kB
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 (&lt;Image_N&gt;) so the rest
# of the caption stays safely escaped while tags render as styled chips.
return re.sub(r"&lt;(Image_\d+)&gt;", 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)