File size: 17,458 Bytes
f40d97f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b4ac5f7
 
f40d97f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b4ac5f7
 
 
 
 
 
 
 
 
 
 
f40d97f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6c5be39
 
 
 
f40d97f
 
 
 
 
 
 
 
 
6c5be39
 
 
 
 
 
f40d97f
6c5be39
 
 
 
 
 
b4ac5f7
6c5be39
 
f40d97f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
61801ee
 
 
 
 
 
 
 
 
 
 
 
 
 
f40d97f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
61801ee
f40d97f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4f55c10
f40d97f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4f55c10
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
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
"""Molmo2Fish β€” interactive fish tracking in ARIS sonar video with natural-language guidance.

Paper: "Teach a Molmo2Fish: Towards interactive fish tracking with natural language
guidance" (arXiv 2608.18602). Model: tidalove/Molmo2Fish.

The demo mirrors the paper's two-stage correction loop:
  1. an initial pass ("track all fish") produces `<tracks coords="...">fish</tracks>`
  2. the user types a plain-English critique and the model re-emits corrected tracks,
     conditioned on the video, its own previous answer, and the critique.
"""

import os

os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")

import spaces  # noqa: E402  β€” must precede torch / CUDA-touching imports

import re  # noqa: E402
import tempfile  # noqa: E402
import time  # noqa: E402
from collections import defaultdict  # noqa: E402

import cv2  # noqa: E402
import gradio as gr  # noqa: E402
import imageio.v2 as imageio  # noqa: E402
import numpy as np  # noqa: E402
import torch  # noqa: E402
from transformers import AutoModelForImageTextToText, AutoProcessor  # noqa: E402

MODEL_ID = "tidalove/Molmo2Fish"

# Matches the released video_preprocessor_config.json of tidalove/Molmo2Fish.
NUM_FRAMES = 128
FRAME_SAMPLE_MODE = "uniform_last_frame"
MAX_FPS = 2
SAMPLING_FPS = 2

TRACK_STYLE = "video_point_track_per_frame"
DEFAULT_PROMPT = "track all fish"
IM_END_TOKEN_ID = 151937     # <im_end>    (config.image_end_token_id)
FRAME_END_TOKEN_ID = 151944  # <frame_end> (config.frame_end_token_id)

# html-v2 pointing format, exactly as in olmo/preprocessing/point_formatter.py
COORD_RE = re.compile(r"<(?:points|tracks).*? coords=\"([0-9\t:;, .]+)\"/?>")
FRAME_RE = re.compile(r"(?:^|\t|:|,|;)([0-9\.]+) ([0-9\. ]+)")
POINTS_RE = re.compile(r"([0-9]+) ([0-9]{3,4}) ([0-9]{3,4})")

PALETTE = [
    (240, 82, 156),   # the authors' pink (scripts/unified_demo.py)
    (66, 214, 255),
    (124, 252, 118),
    (255, 196, 61),
    (186, 132, 255),
    (255, 122, 92),
    (0, 255, 214),
    (255, 255, 120),
]

print(f"Loading {MODEL_ID} …", flush=True)
processor = AutoProcessor.from_pretrained(
    MODEL_ID, trust_remote_code=True, padding_side="left"
)
model = AutoModelForImageTextToText.from_pretrained(
    MODEL_ID, trust_remote_code=True, dtype=torch.bfloat16
).to("cuda")

# The released checkpoint ships a mismatch: processor_config.json has
# use_frame_special_tokens=true (so the prompt gets <frame_start>/<frame_end>
# around every frame, matching training β€” see olmo/models/molmo2/molmo2.py,
# which asserts it), but config.json has it false, so the model counts
# <im_end> instead and asserts "Expected 0 videos, but got 1". Align them.
if processor.use_frame_special_tokens and not model.config.use_frame_special_tokens:
    print("[molmo2fish] aligning config.use_frame_special_tokens -> True", flush=True)
    model.config.use_frame_special_tokens = True
    model.model.config.use_frame_special_tokens = True

model.eval()
print("Model ready.", flush=True)


# --------------------------------------------------------------------------- #
# Track parsing / rendering
# --------------------------------------------------------------------------- #
def parse_tracks(text: str, width: int, height: int) -> dict:
    """Parse `<tracks coords="t id x y …">fish</tracks>` into {time: {id: (x, y)}}.

    Coordinates in the model output are normalised to 0-1000; they are scaled
    back to pixels here.
    """
    out: dict = {}
    for coord in COORD_RE.finditer(text):
        for frame in FRAME_RE.finditer(coord.group(1)):
            t = float(frame.group(1))
            per_frame = out.setdefault(t, {})
            for pt in POINTS_RE.finditer(frame.group(2)):
                idx, xs, ys = pt.group(1), pt.group(2), pt.group(3)
                x = float(xs) / 1000.0 * width
                y = float(ys) / 1000.0 * height
                if 0 <= x <= width and 0 <= y <= height:
                    per_frame.setdefault(idx, (x, y))
    return out


def render_overlay(video_path: str, tracks: dict, out_path: str) -> None:
    """Draw the parsed tracks (points + fading trails + ids) onto the source video."""
    cap = cv2.VideoCapture(video_path)
    fps = cap.get(cv2.CAP_PROP_FPS) or 6.0
    width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
    height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))

    times = sorted(tracks)
    times_arr = np.asarray(times) if times else None

    # id -> ordered list of (time, x, y), used to draw the trail behind each fish
    history = defaultdict(list)
    for t in times:
        for idx, (x, y) in tracks[t].items():
            history[idx].append((t, x, y))

    ids = sorted(history, key=lambda s: (len(s), s))
    color_of = {idx: PALETTE[i % len(PALETTE)] for i, idx in enumerate(ids)}

    radius = max(4, int(max(width, height) * 0.008))
    thickness = max(2, radius // 2)
    font_scale = max(0.5, max(width, height) / 1400.0)

    writer = imageio.get_writer(
        out_path, fps=fps, codec="libx264", quality=7,
        macro_block_size=1, pixelformat="yuv420p", ffmpeg_log_level="error",
    )
    try:
        frame_ix = 0
        while True:
            ok, frame = cap.read()
            if not ok:
                break
            rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
            if times_arr is not None:
                t_now = frame_ix / fps
                k = int(np.argmin(np.abs(times_arr - t_now)))
                t_key = times[k]
                for idx, pts in history.items():
                    trail = [(x, y) for (t, x, y) in pts if t <= t_key]
                    if len(trail) > 1:
                        poly = np.asarray(trail[-24:], dtype=np.int32).reshape(-1, 1, 2)
                        cv2.polylines(rgb, [poly], False, color_of[idx],
                                      max(1, thickness - 1), cv2.LINE_AA)
                for idx, (x, y) in tracks[t_key].items():
                    c = color_of[idx]
                    cv2.circle(rgb, (int(x), int(y)), radius, c, thickness, cv2.LINE_AA)
                    cv2.putText(rgb, str(idx), (int(x) + radius + 3, int(y) - radius - 3),
                                cv2.FONT_HERSHEY_SIMPLEX, font_scale, c,
                                max(1, thickness - 1), cv2.LINE_AA)
            writer.append_data(rgb)
            frame_ix += 1
    finally:
        writer.close()
        cap.release()


def summarise(tracks: dict) -> str:
    if not tracks:
        return "No fish tracks were returned for this clip."
    ids = {i for frame in tracks.values() for i in frame}
    return (f"**{len(ids)} track(s)** across **{len(tracks)}** sampled timesteps "
            f"(2 FPS). Track ids: {', '.join(sorted(ids, key=int))}.")


# --------------------------------------------------------------------------- #
# Model plumbing
# --------------------------------------------------------------------------- #
def build_messages(video_path: str, turns: list) -> list:
    """Chat list for Molmo2Fish. The video is attached to the *first* user turn only.

    `turns` is a list of (user_text, assistant_text_or_None), matching
    olmo/eval/vllm_runner.py::build_multi_turn_chat.
    """
    messages = []
    for i, (user_text, assistant_text) in enumerate(turns):
        content = [dict(type="text", text=user_text, style=TRACK_STYLE)]
        if i == 0:
            # Frame sampling (num_frames=128, uniform_last_frame, max_fps/sampling_fps=2)
            # comes from the model's own video_preprocessor_config.json, so the path is
            # all the processor needs β€” same as olmo/hf_model/test_molmo2.py.
            content.append(dict(type="video", video=video_path))
        messages.append({"role": "user", "content": content})
        if assistant_text is not None:
            messages.append({"role": "assistant",
                             "content": [dict(type="text", text=assistant_text)]})
    return messages


def run_model(video_path: str, turns: list, max_new_tokens: int) -> str:
    messages = build_messages(video_path, turns)
    # Reference path from the repo's own olmo/hf_model/test_molmo2.py: let the
    # Molmo2Processor decode + sample the video and expand <|video|> itself.
    inputs = processor.apply_chat_template(
        messages,
        tokenize=True,
        add_generation_prompt=True,
        return_tensors="pt",
        return_dict=True,
    )
    ids = inputs["input_ids"]
    print(
        f"[molmo2fish] input_ids={tuple(ids.shape)} "
        f"im_end={int((ids == IM_END_TOKEN_ID).sum())} "
        f"frame_end={int((ids == FRAME_END_TOKEN_ID).sum())} "
        f"keys={sorted(inputs.keys())}",
        flush=True,
    )
    inputs = {k: (v.to(model.device) if hasattr(v, "to") else v)
              for k, v in inputs.items()}
    with torch.inference_mode():
        with torch.autocast("cuda", enabled=True, dtype=torch.bfloat16):
            generated = model.generate(
                **inputs, max_new_tokens=max_new_tokens, do_sample=False
            )
    prompt_len = inputs["input_ids"].size(1)
    return processor.post_process_image_text_to_text(
        generated[:, prompt_len:],
        skip_special_tokens=True,
        clean_up_tokenization_spaces=False,
    )[0].strip()


def _video_size(video_path: str):
    cap = cv2.VideoCapture(video_path)
    w = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
    h = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
    cap.release()
    return w, h


def _infer(video_path: str, turns: list, max_new_tokens: int):
    t0 = time.perf_counter()
    raw = run_model(video_path, turns, max_new_tokens)
    elapsed = time.perf_counter() - t0
    width, height = _video_size(video_path)
    tracks = parse_tracks(raw, width, height)
    if not tracks:
        return video_path, raw, f"{summarise(tracks)}  \n_Inference: {elapsed:.1f}s_"
    out_path = tempfile.NamedTemporaryFile(suffix=".mp4", delete=False).name
    render_overlay(video_path, tracks, out_path)
    return out_path, raw, f"{summarise(tracks)}  \n_Inference: {elapsed:.1f}s_"


# --------------------------------------------------------------------------- #
# Gradio handlers
# --------------------------------------------------------------------------- #
# Measured on ZeroGPU: ~29 s for ~800 generated tokens, ~48 s for ~1600, plus
# ~10 s to render the overlay. Runtime is dominated by decoding, so scale the
# GPU reservation with the token budget instead of over-booking a flat number.
def _track_duration(video_path=None, correction_hint="", max_new_tokens=1600,
                    progress=None) -> int:
    return int(25 + 0.028 * int(max_new_tokens or 1600))


def _refine_duration(video_path=None, previous_tracks="", correction="",
                     max_new_tokens=1600, progress=None) -> int:
    return int(25 + 0.028 * int(max_new_tokens or 1600))


@spaces.GPU(duration=_track_duration)
def track_fish(
    video_path: str,
    correction_hint: str = "",
    max_new_tokens: int = 1600,
    progress=gr.Progress(track_tqdm=True),
):
    """Run the first tracking pass over a sonar clip ("track all fish").

    Args:
        video_path: path to an ARIS sonar clip (mp4).
        correction_hint: ignored here β€” it only exists so an example row can
            pre-fill the correction box alongside the video.
        max_new_tokens: generation budget for the `<tracks …>` string.

    Returns:
        (overlay video, raw model output, markdown summary)
    """
    if not video_path:
        raise gr.Error("Please provide a sonar video first.")
    return _infer(video_path, [(DEFAULT_PROMPT, None)], int(max_new_tokens))


@spaces.GPU(duration=_refine_duration)
def refine_tracks(
    video_path: str,
    previous_tracks: str,
    correction: str,
    max_new_tokens: int = 1600,
    progress=gr.Progress(track_tqdm=True),
):
    """Correct the current tracks using a natural-language instruction.

    The model sees the video, its own previous `<tracks …>` answer, and the
    critique, then re-emits a corrected track set.

    Args:
        video_path: the same sonar clip used for the first pass.
        previous_tracks: the model's previous `<tracks …>` output.
        correction: plain-English critique, e.g. "Track 1 is shifted downward".
        max_new_tokens: generation budget for the corrected `<tracks …>` string.

    Returns:
        (overlay video, raw model output, markdown summary)
    """
    if not video_path:
        raise gr.Error("Please provide a sonar video first.")
    if not previous_tracks or not previous_tracks.strip():
        raise gr.Error("Run 'Track all fish' first β€” there is nothing to correct yet.")
    if not correction or not correction.strip():
        raise gr.Error("Type a correction instruction, e.g. 'Track 1 is shifted downward'.")
    turns = [(DEFAULT_PROMPT, previous_tracks.strip()), (correction.strip(), None)]
    return _infer(video_path, turns, int(max_new_tokens))


# --------------------------------------------------------------------------- #
# UI
# --------------------------------------------------------------------------- #
CSS = """
#col-container { max-width: 1200px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
"""

EXAMPLES = [
    [
        "examples/elwha_2018-07-29.mp4",
        "Track 1 looks good overall, just slightly shifted downward from the "
        "actual fish position throughout.",
    ],
    [
        "examples/kenai_leftfar_2018-06-03.mp4",
        "Track 1 doesn't correspond to any real fish β€” you've got a false "
        "detection moving left that should be removed. The actual fish starts "
        "in the lower left around 8s and swims upward until the end of the clip, "
        "and you missed it entirely.",
    ],
    [
        "examples/nushagak_rb_f15-52.mp4",
        "You missed a fish near the top of the frame β€” please add it.",
    ],
]

with gr.Blocks(title="Molmo2Fish tracking") as demo:
    with gr.Column(elem_id="col-container"):
        gr.Markdown(
            "# 🐟 Molmo2Fish β€” interactive fish tracking\n"
            "Track salmon in ARIS **sonar** video, then fix the mistakes by *talking to the model*.\n\n"
            "Step 1 runs the model's `track all fish` pass. Step 2 feeds your plain-English "
            "critique back in β€” the model re-emits a corrected track set instead of you "
            "editing keypoints by hand.\n\n"
            "[Paper](https://huggingface.co/papers/2608.18602) Β· "
            "[Model](https://huggingface.co/tidalove/Molmo2Fish) Β· "
            "[Code](https://github.com/tidalove/molmo2fish) Β· "
            "[Data](https://huggingface.co/datasets/tidalove/cfc-track-instruction)"
        )

        with gr.Row():
            with gr.Column():
                video_in = gr.Video(label="Sonar clip", height=420)
                track_btn = gr.Button("β‘  Track all fish", variant="primary")
                correction = gr.Textbox(
                    label="β‘‘ Correction instruction",
                    placeholder="Track 2 drifts off the fish after about 6s β€” it should keep "
                                "following the fish swimming up the right side.",
                    lines=3,
                )
                refine_btn = gr.Button("β‘‘ Apply correction", variant="secondary")
            with gr.Column():
                video_out = gr.Video(label="Tracks", height=420, autoplay=True)
                summary = gr.Markdown()
                tracks_box = gr.Textbox(
                    label="Model output (html-v2 tracks) β€” edited in place by step β‘‘",
                    lines=6,
                    max_lines=12,
                )

        with gr.Accordion("Advanced", open=False):
            max_new_tokens = gr.Slider(
                256, 3072, value=1600, step=64,
                label="Max new tokens",
                info="Long clips with many fish need a bigger budget; an unclosed "
                     "</tracks> means you hit the cap.",
            )

        gr.Markdown(
            "### Examples\n"
            "Clicking a row loads the clip **and** pre-fills a real correction from the "
            "paper's CFC validation split, and runs step β‘  for you."
        )
        gr.Examples(
            examples=EXAMPLES,
            inputs=[video_in, correction],
            outputs=[video_out, tracks_box, summary],
            fn=track_fish,
            cache_examples=True,
            cache_mode="lazy",
            label="Sonar clips (CFC26, CC-BY-4.0)",
        )

        gr.Markdown(
            "Sonar clips are re-encoded from the "
            "[perona-lab/cfc26](https://huggingface.co/datasets/perona-lab/cfc26) "
            "Caltech Fish Counting release (CC-BY-4.0); correction prompts come from "
            "[tidalove/cfc-track-instruction](https://huggingface.co/datasets/tidalove/cfc-track-instruction). "
            "Tracks are predicted at 2 FPS and interpolated onto the 6 FPS source for display."
        )

    track_btn.click(
        track_fish,
        inputs=[video_in, correction, max_new_tokens],
        outputs=[video_out, tracks_box, summary],
        api_name="track_fish",
    )
    refine_btn.click(
        refine_tracks,
        inputs=[video_in, tracks_box, correction, max_new_tokens],
        outputs=[video_out, tracks_box, summary],
        api_name="refine_tracks",
    )

if __name__ == "__main__":
    demo.launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)