Spaces:
Sleeping
Sleeping
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)
|