Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 34 additions & 0 deletions eomt/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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
``<stem>_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)``)."""
Expand Down
3 changes: 2 additions & 1 deletion eomt/engine/__init__.py
Original file line number Diff line number Diff line change
@@ -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"]
Loading
Loading