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 (&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)