diff --git a/eomt/api.py b/eomt/api.py index 061f031..565a14a 100644 --- a/eomt/api.py +++ b/eomt/api.py @@ -25,6 +25,7 @@ from .engine import evaluate as _evaluate from .engine import evaluate_detection as _evaluate_detection from .engine import predict as _predict +from .engine import track as _track from .engine import train as _train from .model import build_model, load_dinov2_backbone from .serialization import ( @@ -385,6 +386,39 @@ def predict(self, source: str | Path, *, plot: bool = False, save: str | None = imgsz=_resolve_imgsz(imgsz, self.model), **kw, ) + # ------------------------------------------------------------------ track + def track(self, source: str | Path, *, plot: bool = True, save: str | None = "runs/track", + conf_thres: float = 0.3, max_det: int = 100, mask_thresh: float = 0.5, + imgsz: int | None = None, min_hits: int = 3, with_attr: bool = False, + details: bool = False, trace: bool = False, hud: bool = True, + tracker_kwargs: dict | None = None, **kw) -> list[dict] | dict: + """Track instances across a video, returning one result dict per frame. + + Each dict carries the usual ``predict`` keys (``boxes`` / ``scores`` / + ``classes`` / optional ``masks``) plus a persistent ``track_ids`` tensor and a + ``frame`` index. For models with secondary heads, association runs on the + main class + box only, and ``aux_track`` holds the temporally-smoothed + attribute per track (``aux`` keeps the raw per-frame values). With ``plot`` an + annotated ``.mp4`` colored by track id is written under ``save`` (``hud`` adds a + ``Frame i/N | Tracks: n`` banner, ``trace=True`` draws motion trails), plus a + ``_tracks.json`` per-track summary. ``min_hits`` requires a track to + persist that many consecutive frames before it is emitted (speck suppression); + ``tracker_kwargs`` are forwarded to ``ByteTrack`` (e.g. + ``{"lost_track_buffer": 90}``). ``imgsz`` overrides the inference size. + + ``with_attr`` treats each attribute as part of the class — the tracked identity + becomes the compound class ``cls-attr1-attr2…`` (one ByteTrack per compound + class). ``details`` records the per-frame attribute evolution per track and + changes the return to ``{"frames": [...], "tracks": [...]}``. + """ + return _track( + self.model, str(source), plot=plot, save=save, + conf_thres=conf_thres, max_det=max_det, mask_thresh=mask_thresh, + imgsz=_resolve_imgsz(imgsz, self.model), min_hits=min_hits, + with_attr=with_attr, details=details, trace=trace, + hud=hud, tracker_kwargs=tracker_kwargs, **kw, + ) + # ------------------------------------------------------------------- save def save(self, path: str | Path) -> None: """Write a self-describing checkpoint (reloadable with ``EoMT(path)``).""" diff --git a/eomt/engine/__init__.py b/eomt/engine/__init__.py index a97ca54..ff031e8 100644 --- a/eomt/engine/__init__.py +++ b/eomt/engine/__init__.py @@ -1,7 +1,8 @@ """Training, validation and inference engines for EoMT instance segmentation.""" from .predict import predict +from .track import track from .train import train from .validate import evaluate, evaluate_detection, sweep -__all__ = ["train", "evaluate", "evaluate_detection", "predict", "sweep"] +__all__ = ["train", "evaluate", "evaluate_detection", "predict", "track", "sweep"] diff --git a/eomt/engine/track.py b/eomt/engine/track.py new file mode 100644 index 0000000..dc71b24 --- /dev/null +++ b/eomt/engine/track.py @@ -0,0 +1,495 @@ +"""Video object tracking for EoMT (per-frame inference + ByteTrack association). + +Runs the still-image inference (:func:`~eomt.engine.predict.predict_image`) on every +frame of a video, associates detections across frames into persistent ``track_id``s +with ByteTrack, and (optionally) writes an annotated video. For models with auxiliary +attribute heads (e.g. ``laterality`` / ``typology``) the tracker associates on the +**main class + box only** — attributes flicker frame-to-frame and must not drive +identity — while each track keeps a running mean of the attribute probabilities and +reports a temporally-smoothed label (``aux_track``). + +``cv2`` and ``supervision`` are imported lazily so that importing ``eomt`` or running +``predict`` never requires them; they are only needed when :func:`track` is called. +""" + +from __future__ import annotations + +import json +import time +import warnings +from pathlib import Path + +import numpy as np +import torch +from PIL import Image + +from ..serialization import load_model +from .predict import predict_image + +_AUX_DATA_PREFIX = "aux::" # namespaces aux probs stashed in sv.Detections.data + + +def track( + model, + source: str, + *, + plot: bool = True, + save: str | None = "runs/track", + conf_thres: float = 0.3, + max_det: int = 100, + mask_thresh: float = 0.5, + device: str = "auto", + imgsz: int | None = None, + tracker_kwargs: dict | None = None, + min_hits: int = 3, + with_attr: bool = False, + details: bool = False, + trace: bool = False, + hud: bool = True, + alpha: float = 0.35, + show_scores: bool = True, +) -> list[dict] | dict: + """Track instances across a video, returning one result dict per frame. + + ``model`` may be a loaded :class:`~eomt.model.EoMTModel` or a checkpoint path / run + folder (loaded with :func:`~eomt.serialization.load_model`). Each returned dict + carries the usual :func:`~eomt.engine.predict.predict_image` keys + (``boxes`` / ``scores`` / ``classes`` / optional ``masks`` / optional ``aux``) plus + a persistent ``track_ids`` tensor aligned to the detections, a ``frame`` index, and + — for models with secondary heads — ``aux_track`` with the temporally-smoothed + attribute ids/probs per track. When ``plot`` is set, an annotated ``.mp4`` (colored + by track id, optional ``Frame i/N | Tracks: n`` ``hud`` banner) is written under + ``save`` and, alongside it, a ``_tracks.json`` per-track summary. + + Args: + tracker_kwargs: forwarded to ``supervision.ByteTrack`` (e.g. + ``{"lost_track_buffer": 90}`` to keep an id longer when a part leaves view). + min_hits: a track must be seen this many *consecutive* frames before it is + emitted or drawn — suppresses one-frame false-positive specks. Once a track + clears the bar it stays confirmed for the rest of the video. Set ``1`` to + emit every track immediately. + with_attr: treat the attribute(s) as part of the class — the tracked identity + becomes the compound class ``cls-attr1-attr2…``. Since ByteTrack is + class-agnostic, this runs one ByteTrack **per compound class** (per-class + association), so a ``front_door-left`` detection can never match a + ``front_door-right`` track. Labels/summary report the compound name. A part + whose attribute flickers briefly spawns a short competing track (suppressed + by ``min_hits``). Default ``False`` keeps class + box association with the + attribute smoothed by its running mean. + details: return and record the full per-frame **evolution** of each attribute + per track (raw id/name/confidence + running-smoothed id per frame). When + set, the function returns ``{"frames": [...per-frame...], "tracks": + [...summary incl. evolution...]}`` instead of the per-frame list, and the + ``_tracks.json`` gains an ``evolution`` array per track. + + Returns: + By default a ``list[dict]`` (one per frame). When ``details`` is set, a dict + ``{"frames": list[dict], "tracks": list[dict]}``. + """ + import cv2 # lazy: only needed for tracking + import supervision as sv + + if isinstance(model, (str, Path)): + model = load_model(model, device=device) + + dev = next(model.parameters()).device + imgsz = int(imgsz if imgsz is not None else model.image_size) + letterbox = bool(getattr(model, "preprocess_letterbox", False)) + names = getattr(model, "names", None) + aux_specs = list(getattr(model, "aux_specs", [])) + aux_head_names = [s.name for s in aux_specs] + aux_label_names = {s.name: s.names for s in aux_specs} + aux_ns = {s.name: int(s.num_classes) for s in aux_specs} + family = getattr(model, "family", "instance") + + cap = cv2.VideoCapture(str(source)) + if not cap.isOpened(): + raise FileNotFoundError(f"could not open video: {source}") + fps = cap.get(cv2.CAP_PROP_FPS) or 30.0 + W = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) + H = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) + n_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT)) or 0 + + print( + f"[track] eomt-{getattr(model, 'size', '?')} ({family}, " + f"nc={getattr(model, 'nc', '?')}, imgsz={imgsz}) on {dev} | " + f"{source}: {W}x{H} @ {fps:.1f} fps, {n_frames or '?'} frames" + ) + + # ByteTrack is deprecated in supervision>=0.28 (removed in 0.30); we pin <0.30 and + # keep using it per the chosen design. Class-agnostic by default; one tracker per + # compound class when with_attr (see _Associator). FutureWarning silenced inside. + assoc = _Associator(sv, tracker_kwargs, with_attr=with_attr) + + text_scale = max(0.4, W / 1600.0) + thickness = max(1, round(W / 640.0)) + lookup = sv.ColorLookup.TRACK + mask_ann = sv.MaskAnnotator(color_lookup=lookup, opacity=alpha) if family != "detect" else None + box_ann = sv.BoxAnnotator(color_lookup=lookup, thickness=thickness) + label_ann = sv.LabelAnnotator(color_lookup=lookup, text_scale=text_scale, smart_position=True) + trace_ann = sv.TraceAnnotator(color_lookup=lookup, thickness=thickness) if trace else None + + save_dir = Path(save) if save else None + if save_dir is not None: + save_dir.mkdir(parents=True, exist_ok=True) + writer = None + video_path = None + if plot and save_dir is not None: + video_path = str(save_dir / f"{Path(source).stem}_track.mp4") + writer = cv2.VideoWriter(video_path, cv2.VideoWriter_fourcc(*"mp4v"), fps, (W, H)) + + # Running per-track aux history: tid -> head -> [sum_probs (ns,), count]. + hist: dict[int, dict[str, list]] = {} + # min_hits confirmation state and per-track summary accumulators. + streak: dict[int, int] = {} + last_seen: dict[int, int] = {} + confirmed: set[int] = set() + summary: dict[int, dict] = {} + # per-track attribute time-series (details=True only): tid -> {head|"class": [...]}. + evolution: dict[int, dict[str, list]] = {} + + results: list[dict] = [] + idx = 0 + t0 = time.perf_counter() + try: + while True: + ok, frame = cap.read() # BGR uint8 + if not ok: + break + image = Image.fromarray(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)) + + raw = predict_image( + model, image, device=dev, imgsz=imgsz, + conf_thres=conf_thres, max_det=max_det, + mask_thresh=mask_thresh, letterbox=letterbox, + ) + + # ---- to sv.Detections (aux probs ride along in .data so they survive) ---- + n = int(raw["num_detections"]) + data = {} + for name in aux_head_names: + probs = raw.get("aux", {}).get(name, {}).get("probs") + data[_AUX_DATA_PREFIX + name] = ( + probs.cpu().numpy() if probs is not None else np.zeros((n, 0), np.float32) + ) + base_cls = raw["classes"].cpu().numpy().astype(int) + det = sv.Detections( + xyxy=raw["boxes"].cpu().numpy().reshape(-1, 4).astype(np.float32), + confidence=raw["scores"].cpu().numpy().astype(np.float32), + class_id=base_cls, + mask=(raw["masks"].cpu().numpy().astype(bool) if "masks" in raw else None), + data=data, + ) + + # compound key per detection (base cls + per-head raw argmax); only used + # to route detections to per-class trackers when with_attr. + keys = None + if with_attr: + if aux_head_names: + attr_ids = [raw["aux"][h]["ids"].cpu().numpy().astype(int) for h in aux_head_names] + keys = [(int(base_cls[i]), *(int(a[i]) for a in attr_ids)) for i in range(n)] + else: + keys = [(int(base_cls[i]),) for i in range(n)] + + tracked, all_tids, det_keys = assoc.update(det, keys) + + # ---- min_hits: keep only tracks seen >= min_hits consecutive frames ---- + keep = np.zeros(len(tracked), bool) + for k, tid in enumerate(all_tids): + tid = int(tid) + streak[tid] = (streak.get(tid, 0) + 1) if last_seen.get(tid) == idx - 1 else 1 + last_seen[tid] = idx + if streak[tid] >= min_hits: + confirmed.add(tid) + keep[k] = tid in confirmed + tracked = tracked[keep] + m = len(tracked) + tids = all_tids[keep] + frame_keys = [det_keys[k] for k in range(len(keep)) if keep[k]] if det_keys is not None else None + + # ---- rebuild an aligned per-frame result from the tracked detections ---- + res: dict = { + "frame": idx, + "num_detections": m, + "boxes": torch.from_numpy(tracked.xyxy).float(), + "scores": torch.from_numpy(np.asarray(tracked.confidence, np.float32)), + "classes": torch.from_numpy(np.asarray(tracked.class_id, int)).long(), + "track_ids": torch.from_numpy(tids).long(), + } + if tracked.mask is not None: + res["masks"] = torch.from_numpy(tracked.mask).bool() + + # ---- aux: raw per-frame + temporally-smoothed per-track ---- + if aux_head_names: + res["aux"], res["aux_track"] = {}, {} + for name in aux_head_names: + # sv drops .data on empty frames; fall back to a correctly-shaped + # empty so downstream indexing stays aligned. + probs = np.asarray( + tracked.data.get(_AUX_DATA_PREFIX + name, np.zeros((m, aux_ns[name]), np.float32)), + np.float32, + ).reshape(m, aux_ns[name]) + res["aux"][name] = { + "ids": torch.from_numpy(probs.argmax(1)).long() if m else torch.zeros(0, dtype=torch.long), + "probs": torch.from_numpy(probs), + } + smoothed = np.zeros_like(probs) + for j in range(m): + acc = hist.setdefault(int(tids[j]), {}).setdefault( + name, [np.zeros(probs.shape[1], np.float32), 0] + ) + acc[0] += probs[j] + acc[1] += 1 + smoothed[j] = acc[0] / acc[1] + res["aux_track"][name] = { + "ids": torch.from_numpy(smoothed.argmax(1)).long() if m else torch.zeros(0, dtype=torch.long), + "probs": torch.from_numpy(smoothed), + } + + # ---- compound class name per detection (with_attr) ---- + cls_list = res["classes"].tolist() + compound_names = None + if frame_keys is not None: + compound_names = [ + _compound_name(k, names, aux_head_names, aux_label_names) for k in frame_keys + ] + res["compound"] = compound_names + + # ---- attribute evolution over time (details only) ---- + if details: + for j in range(m): + ev = evolution.setdefault(int(tids[j]), {}) + for name in aux_head_names: + rid = int(res["aux"][name]["ids"][j]) + rp = float(res["aux"][name]["probs"][j][rid]) + sid = int(res["aux_track"][name]["ids"][j]) + ev.setdefault(name, []).append({ + "frame": idx, "id": rid, + "name": aux_label_names.get(name, {}).get(rid, str(rid)), + "conf": round(rp, 4), "smoothed_id": sid, + "smoothed_name": aux_label_names.get(name, {}).get(sid, str(sid)), + }) + cname = (compound_names[j] if compound_names is not None + else (names.get(int(cls_list[j]), str(int(cls_list[j]))) if names else str(int(cls_list[j])))) + ev.setdefault("class", []).append({"frame": idx, "name": cname}) + + # ---- per-track summary accumulators (confirmed detections only) ---- + for j in range(m): + tid = int(tids[j]) + s = summary.setdefault(tid, {"cls": {}, "compound": {}, "first": idx, "last": idx, "n": 0}) + s["cls"][int(cls_list[j])] = s["cls"].get(int(cls_list[j]), 0) + 1 + if frame_keys is not None: + s["compound"][frame_keys[j]] = s["compound"].get(frame_keys[j], 0) + 1 + s["last"] = idx + s["n"] += 1 + + # ---- render ---- + if writer is not None: + scene = frame + if m: + if mask_ann is not None and tracked.mask is not None: + scene = mask_ann.annotate(scene, tracked) + if trace_ann is not None: + scene = trace_ann.annotate(scene, tracked) + scene = box_ann.annotate(scene, tracked) + scene = label_ann.annotate(scene, tracked, labels=_labels( + m, tids, res, names, aux_head_names, aux_label_names, show_scores, + compound_names=compound_names, + )) + if hud: + _draw_hud(cv2, scene, idx + 1, n_frames, m, text_scale) + writer.write(scene) + + if video_path is not None: + res["video_path"] = video_path + results.append(res) + idx += 1 + finally: + cap.release() + if writer is not None: + writer.release() + + total = time.perf_counter() - t0 + fps_out = idx / total if total > 0 else float("inf") + dest = f" -> {video_path}" if video_path else "" + print(f"[track] done: {idx} frame(s) in {total:.2f} s ({fps_out:.1f} FPS){dest}") + + # ---- per-track summary: id -> class, smoothed attrs, frame span, frames seen ---- + tracks = _build_summary( + summary, hist, names, aux_head_names, aux_label_names, + with_attr=with_attr, evolution=(evolution if details else None), + ) + if save_dir is not None: + out = save_dir / f"{Path(source).stem}_tracks.json" + out.write_text(json.dumps({"source": str(source), "num_frames": idx, "tracks": tracks}, indent=2) + "\n") + print(f"[track] {len(tracks)} track(s) -> {out}") + _print_summary(tracks) + if details: + return {"frames": results, "tracks": tracks} + return results + + +def _build_summary(summary, hist, names, aux_head_names, aux_label_names, + *, with_attr=False, evolution=None) -> list[dict]: + """One record per track: majority class, smoothed attributes, span, frames seen. + + With ``with_attr`` the reported ``class``/``class_name`` is the majority compound + class (``base_class`` keeps the plain class). With ``evolution`` each record gains + a per-frame ``evolution`` of the attribute predictions. + """ + tracks = [] + for tid, s in sorted(summary.items()): + base_cls = max(s["cls"], key=s["cls"].get) + base_name = names.get(base_cls, str(base_cls)) if names else str(base_cls) + attrs = {} + for name in aux_head_names: + acc = hist.get(tid, {}).get(name) + if acc and acc[1]: + aid = int((acc[0] / acc[1]).argmax()) + attrs[name] = aux_label_names.get(name, {}).get(aid, str(aid)) + rec = { + "track_id": tid, + "class": base_cls, + "class_name": base_name, + "attributes": attrs, + "first_frame": s["first"], + "last_frame": s["last"], + "frames_seen": s["n"], + } + if with_attr and s.get("compound"): + key = max(s["compound"], key=s["compound"].get) + rec["base_class"], rec["base_class_name"] = base_cls, base_name + rec["class"] = "-".join(str(x) for x in key) + rec["class_name"] = _compound_name(key, names, aux_head_names, aux_label_names) + if evolution is not None: + rec["evolution"] = evolution.get(tid, {}) + tracks.append(rec) + return tracks + + +def _print_summary(tracks: list[dict], top: int = 15) -> None: + """Print a compact per-track table (most-seen first).""" + if not tracks: + print("[track] no confirmed tracks") + return + ranked = sorted(tracks, key=lambda t: t["frames_seen"], reverse=True) + print(f"[track] top tracks by frames seen (of {len(tracks)}):") + for t in ranked[:top]: + attrs = " ".join(t["attributes"].values()) + span = f"{t['first_frame']}-{t['last_frame']}" + print(f" #{t['track_id']:>3} {t['class_name']:<24} {attrs:<14} frames {span} ({t['frames_seen']})") + + +def _draw_hud(cv2, scene, i: int, n: int, ntracks: int, scale: float) -> None: + """Draw a top-left ``Frame i/N | Tracks: n`` banner (n = tracks in this frame).""" + text = f"Frame {i}/{n or '?'} | Tracks: {ntracks}" + fs = max(0.6, scale * 1.4) + th = max(1, round(fs * 2)) + (tw, hh), base = cv2.getTextSize(text, cv2.FONT_HERSHEY_SIMPLEX, fs, th) + pad = int(hh * 0.5) + cv2.rectangle(scene, (0, 0), (tw + 2 * pad, hh + base + 2 * pad), (0, 0, 0), -1) + cv2.putText(scene, text, (pad, hh + pad), cv2.FONT_HERSHEY_SIMPLEX, fs, + (255, 255, 255), th, cv2.LINE_AA) + + +def _labels(m, tids, res, names, aux_head_names, aux_label_names, show_scores, + *, compound_names=None) -> list[str]: + """Build per-detection label strings: ``#id class [score] [smoothed attrs]``. + + When ``compound_names`` is given (with_attr), the compound class name is shown + (e.g. ``front_door-right``) instead of the base class plus trailing attributes. + """ + classes = res["classes"].tolist() + scores = res["scores"].tolist() + out: list[str] = [] + for j in range(m): + if compound_names is not None: + parts = [f"#{int(tids[j])} {compound_names[j]}"] + if show_scores: + parts.append(f"{float(scores[j]):.2f}") + else: + cls = int(classes[j]) + cname = names.get(cls, str(cls)) if names else str(cls) + parts = [f"#{int(tids[j])} {cname}"] + if show_scores: + parts.append(f"{float(scores[j]):.2f}") + for name in aux_head_names: + aid = int(res["aux_track"][name]["ids"][j]) + parts.append(aux_label_names.get(name, {}).get(aid, str(aid))) + out.append(" ".join(parts)) + return out + + +def _compound_name(key, names, aux_head_names, aux_label_names) -> str: + """``(base_cls, attr1_id, …)`` -> ``base_name-attr1_name-…`` (e.g. front_door-right).""" + base = int(key[0]) + parts = [names.get(base, str(base)) if names else str(base)] + for name, aid in zip(aux_head_names, key[1:]): + parts.append(aux_label_names.get(name, {}).get(int(aid), str(int(aid)))) + return "-".join(parts) + + +class _Associator: + """ByteTrack wrapper: class-agnostic by default, or one tracker per compound class. + + ``update(det, keys)`` returns ``(tracked_detections, global_track_ids, det_keys)``. + In the default mode the single ByteTrack already yields unique ids and ``det_keys`` + is ``None``. With ``with_attr`` detections are routed to a per-compound-class + ByteTrack (so attributes drive identity); each tracker's local ids are remapped to a + single global id space, and every live tracker is stepped each frame (empty subset + when its class is absent) so ``lost_track_buffer`` ages in real frames. + """ + + def __init__(self, sv, tracker_kwargs, *, with_attr: bool): + self._sv = sv + self._kw = dict(tracker_kwargs or {}) + self.with_attr = with_attr + self._trackers: dict = {} + self._gid: dict = {} + self._counter = 0 + if not with_attr: + self._trackers[None] = self._new() + + def _new(self): + with warnings.catch_warnings(): + warnings.simplefilter("ignore", FutureWarning) + return self._sv.ByteTrack(**self._kw) + + def _global(self, key, local: int) -> int: + kk = (key, local) + g = self._gid.get(kk) + if g is None: + self._counter += 1 + g = self._gid[kk] = self._counter + return g + + def update(self, det, keys): + if not self.with_attr: + tracked = self._trackers[None].update_with_detections(det) + tids = (tracked.tracker_id.astype(int) if tracked.tracker_id is not None + else np.zeros((len(tracked),), int)) + return tracked, tids, None + + groups: dict = {} + for i, k in enumerate(keys or []): + groups.setdefault(k, []).append(i) + out_dets, out_keys = [], [] + for k in set(self._trackers) | set(groups): + if k is None: + continue + idxs = np.asarray(groups.get(k, []), dtype=int) + sub = det[idxs] # empty subset ages the tracker without adding detections + tr = self._trackers.get(k) or self._trackers.setdefault(k, self._new()) + tracked = tr.update_with_detections(sub) + if len(tracked) == 0: + continue + local = tracked.tracker_id.astype(int) + tracked.tracker_id = np.array([self._global(k, int(x)) for x in local], int) + out_dets.append(tracked) + out_keys.extend([k] * len(tracked)) + if not out_dets: + empty = self._sv.Detections.empty() + empty.tracker_id = np.array([], int) + return empty, np.array([], int), [] + merged = self._sv.Detections.merge(out_dets) + return merged, merged.tracker_id.astype(int), out_keys diff --git a/pyproject.toml b/pyproject.toml index 8dd967f..f69d115 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -31,6 +31,8 @@ dependencies = [ "requests>=2.25.0", "huggingface_hub>=0.23.0", "torchao>=0.7.0", + "opencv-python>=4.8.0", + "supervision>=0.22.0,<0.30", ] [project.optional-dependencies] diff --git a/scripts/track.py b/scripts/track.py new file mode 100644 index 0000000..aafe10b --- /dev/null +++ b/scripts/track.py @@ -0,0 +1,50 @@ +#!/usr/bin/env python +"""Track segmentation instances (+ smoothed attributes) across a video. + + python scripts/track.py runs/train/eomt-l video.mp4 --conf 0.3 + python scripts/track.py PATH/TO/WEIGHTS video.mp4 --trace # draw motion trails + python scripts/track.py PATH/TO/WEIGHTS video.mp4 --with-attr # compound cls-attr tracks + python scripts/track.py PATH/TO/WEIGHTS video.mp4 --details # per-frame attribute evolution + +Writes an annotated .mp4 to runs/track/, colored by persistent track id. Each track is +labelled with its id, class and — for models with auxiliary heads — the temporally +smoothed attribute (steadier than per-frame predict output). +""" + +from __future__ import annotations + +import argparse + +from eomt import EoMT + + +def main() -> None: + p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + p.add_argument("weights", help="Checkpoint .pt or a run/weights folder.") + p.add_argument("source", help="Video file (mp4/avi/...).") + p.add_argument("--out", default="runs/track", help="Output directory for the annotated video.") + p.add_argument("--conf", type=float, default=0.3, help="Confidence threshold.") + p.add_argument("--imgsz", type=int, default=None, help="Inference image size (defaults to trained size).") + p.add_argument("--min-hits", type=int, default=3, help="Frames a track must persist before it is drawn.") + p.add_argument("--track-buffer", type=int, default=None, help="ByteTrack lost_track_buffer (frames to keep an id after it leaves view).") + p.add_argument("--with-attr", action="store_true", help="Treat attributes as part of the class (compound cls-attr tracking).") + p.add_argument("--details", action="store_true", help="Record per-frame attribute evolution per track (returns {frames, tracks}).") + p.add_argument("--trace", action="store_true", help="Draw per-track motion trails.") + p.add_argument("--no-hud", action="store_true", help="Hide the Frame/Tracks banner.") + p.add_argument("--device", default="auto") + args = p.parse_args() + + tracker_kwargs = {"lost_track_buffer": args.track_buffer} if args.track_buffer else None + model = EoMT(args.weights, device=args.device) + out = model.track( + args.source, plot=True, save=args.out, + conf_thres=args.conf, imgsz=args.imgsz, min_hits=args.min_hits, + with_attr=args.with_attr, details=args.details, + trace=args.trace, hud=not args.no_hud, tracker_kwargs=tracker_kwargs, + ) + n_frames = len(out["frames"]) if args.details else len(out) + print(f"[done] tracked {n_frames} frame(s); annotated video under {args.out}") + + +if __name__ == "__main__": + main() diff --git a/uv.lock b/uv.lock index e55f5d4..6a4e865 100644 --- a/uv.lock +++ b/uv.lock @@ -50,7 +50,7 @@ wheels = [ [[package]] name = "attr-eomt" -version = "0.1.3" +version = "1.0.0" source = { editable = "." } dependencies = [ { name = "huggingface-hub" }, @@ -59,6 +59,7 @@ dependencies = [ { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, { name = "numpy", version = "2.4.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" }, { name = "numpy", version = "2.5.0", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, + { name = "opencv-python" }, { name = "pillow" }, { name = "pycocotools" }, { name = "pyyaml" }, @@ -66,6 +67,7 @@ dependencies = [ { name = "scipy", version = "1.15.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, { name = "scipy", version = "1.17.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" }, { name = "scipy", version = "1.18.0", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, + { name = "supervision" }, { name = "torch" }, { name = "torchao" }, { name = "torchvision" }, @@ -90,12 +92,14 @@ requires-dist = [ { name = "huggingface-hub", specifier = ">=0.23.0" }, { name = "matplotlib", specifier = ">=3.5.0" }, { name = "numpy", specifier = ">=1.19.0" }, + { name = "opencv-python", specifier = ">=4.8.0" }, { name = "pillow", specifier = ">=9.1.0" }, { name = "pycocotools", specifier = ">=2.0.0" }, { name = "pytest", marker = "extra == 'dev'", specifier = ">=7.0" }, { name = "pyyaml", specifier = ">=6.0" }, { name = "requests", specifier = ">=2.25.0" }, { name = "scipy", specifier = ">=1.7.0" }, + { name = "supervision", specifier = ">=0.22.0,<0.30" }, { name = "tensorboard", marker = "extra == 'logging'", specifier = ">=2.10" }, { name = "torch", specifier = ">=2.4.0" }, { name = "torchao", specifier = ">=0.7.0" }, @@ -605,6 +609,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/e7/05/c19819d5e3d95294a6f5947fb9b9629efb316b96de511b418c53d245aae6/cycler-0.12.1-py3-none-any.whl", hash = "sha256:85cef7cff222d8644161529808465972e51340599459b8ac3ccbac5a854e0d30", size = 8321, upload-time = "2023-10-07T05:32:16.783Z" }, ] +[[package]] +name = "defusedxml" +version = "0.7.1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/0f/d5/c66da9b79e5bdb124974bfe172b4daf3c984ebd9c2a06e2b8a4dc7331c72/defusedxml-0.7.1.tar.gz", hash = "sha256:1bb3032db185915b62d7c6209c5a8792be6a32ab2fedacc84e01b52c51aa3e69", size = 75520, upload-time = "2021-03-08T10:59:26.269Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/07/6c/aa3f2f849e01cb6a001cd8554a88d4c77c5c1a31c95bdf1cf9301e6d9ef4/defusedxml-0.7.1-py2.py3-none-any.whl", hash = "sha256:a352e7e428770286cc899e2542b6cdaedb2b4953ff269a210103ec58f6198a61", size = 25604, upload-time = "2021-03-08T10:59:24.45Z" }, +] + [[package]] name = "docutils" version = "0.23" @@ -1807,6 +1820,27 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/a8/64/3708a90d1ebe202ffdeb7185f878a3c84d15c2b2c31858da2ce0583e2def/nvidia_nvtx-13.0.85-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:cb7780edb6b14107373c835bf8b72e7a178bac7367e23da7acb108f973f157a6", size = 148878, upload-time = "2025-09-04T08:28:53.627Z" }, ] +[[package]] +name = "opencv-python" +version = "5.0.0.93" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, + { name = "numpy", version = "2.4.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" }, + { name = "numpy", version = "2.5.0", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/79/4c/a438d23e09ce2033c09f7b784ad2fbdb0adf529e434101ed28f142226f98/opencv_python-5.0.0.93.tar.gz", hash = "sha256:66aac3e5b5faa48d4025816592f3af19e4bfc2c68dec067bae2dbb4ca10aa9e2", size = 81802749, upload-time = "2026-07-02T06:59:53.815Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/9c/75/76f6ade78f6102c61034f828e2a22616708df2c9504bc8d6af9dd8f73dc5/opencv_python-5.0.0.93-cp37-abi3-macosx_13_0_arm64.whl", hash = "sha256:198a75138241810206a17c829dbcc40a7cb1841cda538ca86cbbfc6c7d95f898", size = 48322443, upload-time = "2026-07-02T05:50:25.466Z" }, + { url = "https://files.pythonhosted.org/packages/15/8c/bc1bda6aae69a32e9d84fc34153ba104cd25226861eb4aea33b2cea4860d/opencv_python-5.0.0.93-cp37-abi3-macosx_14_0_x86_64.whl", hash = "sha256:6bbc32f59e1b1a7db7b39c81f63d00625f041d333037fd8702f6da52cc39108b", size = 34782755, upload-time = "2026-07-02T05:51:30.556Z" }, + { url = "https://files.pythonhosted.org/packages/f4/8a/b04776ec45d2dea08a1b176f1829201db3515d4ed16c35f8fcc9fa7beb16/opencv_python-5.0.0.93-cp37-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:e2b4272e736836f66c2d176e43ab8101f3a00d45654916399f52e150c58981ac", size = 50614064, upload-time = "2026-07-02T06:53:22.604Z" }, + { url = "https://files.pythonhosted.org/packages/95/54/eb47866b94f2b5b42dde17644b78055ef1ee05aae59962c7290e55270803/opencv_python-5.0.0.93-cp37-abi3-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:f8b6d0a212253dd26ad338c812f1f23ca118fdf05a9c8c6b9444f161aa8c5881", size = 71064711, upload-time = "2026-07-02T06:54:13.148Z" }, + { url = "https://files.pythonhosted.org/packages/93/da/962579f1e703cbf8c5422fd1f576467dcb3b5b0b0b81c1471c979764353a/opencv_python-5.0.0.93-cp37-abi3-manylinux_2_28_aarch64.whl", hash = "sha256:08d5d91d967b58d6db86073b2ad3eaef88ca4ebdfd45c9059bf59f5ded0c7ad2", size = 49798576, upload-time = "2026-07-02T06:54:33.781Z" }, + { url = "https://files.pythonhosted.org/packages/cf/4c/c73f828fdbcd37eaf21d08fa852544a3ca7c2dbb3ea76873d64f2ea413d1/opencv_python-5.0.0.93-cp37-abi3-manylinux_2_28_x86_64.whl", hash = "sha256:c8de2dec111122a02e8beb28e16c31904992dfd6186560b142a92c71403c1039", size = 73783032, upload-time = "2026-07-02T06:55:03.415Z" }, + { url = "https://files.pythonhosted.org/packages/e2/4b/edaf83b996ca5a1a3d8ccad485706b9c6d4742b13b9c4586bf1c1e7d9423/opencv_python-5.0.0.93-cp37-abi3-win32.whl", hash = "sha256:4b4b1a34c79bf8d3738e3cfe9a9e67b51a79663f6b692cbdad8c31f570da4157", size = 35564734, upload-time = "2026-07-02T05:49:57.704Z" }, + { url = "https://files.pythonhosted.org/packages/21/f0/9fa6e85cb10c8eb36a0222d27e50fe381b86ce49a55446bf39f491727564/opencv_python-5.0.0.93-cp37-abi3-win_amd64.whl", hash = "sha256:f90ba04b8f73bc5c3814037699739f0156f597338a98f05956c684e7c3ca10d2", size = 44000345, upload-time = "2026-07-02T05:49:54.971Z" }, +] + [[package]] name = "packaging" version = "26.2" @@ -2135,6 +2169,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/4b/2d/69abac8f838090bbecd5df894befb2c2619e7996a98ddb949db9f3b93225/pydantic_core-2.46.4-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:d51026d73fcfd93610abc7b27789c26b313920fcfb20e27462d74a7f8b06e983", size = 2193071, upload-time = "2026-05-06T13:38:08.682Z" }, ] +[[package]] +name = "pydeprecate" +version = "0.9.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/d4/1e/b069ea704e42fefeb523835ba3d712347ff672743e5fa9428f4201525328/pydeprecate-0.9.0.tar.gz", hash = "sha256:ed49a36506f63a8d4ae84ffedff1e2879733d536c7a6aaafdf8bfcd39520059b", size = 151667, upload-time = "2026-06-05T19:51:17.43Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/47/87/a155f16c1c17adccaa2ed2256484fe1c6ad37b925e8e39695ed8860bc0b5/pydeprecate-0.9.0-py3-none-any.whl", hash = "sha256:c9501dae1a2567a3c7f7b04221ad5a97d1ad972a02ddb729a04299b435fedd7c", size = 102725, upload-time = "2026-06-05T19:51:16.235Z" }, +] + [[package]] name = "pygments" version = "2.20.0" @@ -2722,6 +2765,32 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/c1/d4/59e74daffcb57a07668852eeeb6035af9f32cbfd7a1d2511f17d2fe6a738/smmap-5.0.3-py3-none-any.whl", hash = "sha256:c106e05d5a61449cf6ba9a1e650227ecfb141590d2a98412103ff35d89fc7b2f", size = 24390, upload-time = "2026-03-09T03:43:24.361Z" }, ] +[[package]] +name = "supervision" +version = "0.29.1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "defusedxml" }, + { name = "matplotlib", version = "3.10.9", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, + { name = "matplotlib", version = "3.11.0", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.11'" }, + { name = "numpy", version = "2.2.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, + { name = "numpy", version = "2.4.6", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" }, + { name = "numpy", version = "2.5.0", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, + { name = "opencv-python" }, + { name = "pillow" }, + { name = "pydeprecate" }, + { name = "pyyaml" }, + { name = "requests" }, + { name = "scipy", version = "1.15.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.11'" }, + { name = "scipy", version = "1.17.1", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version == '3.11.*'" }, + { name = "scipy", version = "1.18.0", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, + { name = "tqdm" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/dc/ec/293e5693ef5f32b54533b680b909a96976925adbb5c3bbfbf53b199182e0/supervision-0.29.1.tar.gz", hash = "sha256:2db9f968a3253e6778573c39bbb8a8e3fa586ece6c98fbdc41d8fd8508eda5af", size = 246950, upload-time = "2026-06-23T19:55:37.717Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/1e/7c/ed95842ac2e323f6a0fcaf5a5bb56b62c0596a4b177872da9f7605ca1ceb/supervision-0.29.1-py3-none-any.whl", hash = "sha256:8518f6bad0d19ff7a959d1b9eb1b54517aded814aa7f237598fed0e3a772c281", size = 280218, upload-time = "2026-06-23T19:55:36.114Z" }, +] + [[package]] name = "sympy" version = "1.14.0"