From f2ee9501e728d1e00251afc644c85203dc94b05d Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Sun, 4 Oct 2026 09:28:24 -0700 Subject: [PATCH 01/28] feat(tracking): link detections into chains, store occurrence statistics, split at new sessions Add the data layer that automated tracking builds on. Each detection can now point to the detection that follows it in the same insect's path (Detection.next_detection), and each occurrence stores four statistics derived from its detections: how far the insect moved, how much its box changed size, how many distinct taxa it was labelled with, and how consistently the determination agrees with those labels. Storing them lets occurrences be sorted by them without recomputing per row; a backfill command fills in rows that existed before the fields. Regrouping captures into sessions can draw a new boundary through an occurrence. The regroup now splits such an occurrence into one per session, copying its identifications to the later pieces so no person's work is lost, and refreshes the statistics of every piece. The link across the boundary is kept because tracking never walks across sessions. Sessions can be row-locked (lock_sessions) so a regroup and a tracking run cannot write to the same session at once, and a helper refreshes the cached counts of sessions and stations after occurrences are created or split. Co-Authored-By: Claude Opus 5.5 --- .../commands/backfill_track_stats.py | 61 +++++ .../0096_detection_next_detection.py | 25 ++ .../migrations/0097_occurrence_track_stats.py | 48 ++++ ami/main/models.py | 98 +++++++- ami/main/models_future/track_stats.py | 208 ++++++++++++++++ ami/main/models_future/tracks.py | 116 +++++++++ ami/main/tasks.py | 9 + ami/main/tests.py | 231 ++++++++++++++++++ 8 files changed, 795 insertions(+), 1 deletion(-) create mode 100644 ami/main/management/commands/backfill_track_stats.py create mode 100644 ami/main/migrations/0096_detection_next_detection.py create mode 100644 ami/main/migrations/0097_occurrence_track_stats.py create mode 100644 ami/main/models_future/track_stats.py create mode 100644 ami/main/models_future/tracks.py diff --git a/ami/main/management/commands/backfill_track_stats.py b/ami/main/management/commands/backfill_track_stats.py new file mode 100644 index 000000000..4d4a8b889 --- /dev/null +++ b/ami/main/management/commands/backfill_track_stats.py @@ -0,0 +1,61 @@ +""" +Store track statistics on the occurrences of one project. + +Tracking and regrouping keep ``Occurrence.track_*`` current from the moment the fields +exist; this fills in the rows from before that, so occurrences can be sorted by motion, +size change and identification agreement. Occurrences with several detections are the +ones worth sorting by, so single-detection ones are skipped unless asked for. + +Safe to re-run: every pass recomputes from the current detections. ``updated_at`` is +left alone, so the default list order does not change. +""" + +from django.core.management.base import BaseCommand, CommandError +from django.db.models import Count + +from ami.main.models import Occurrence, Project +from ami.main.models_future.track_stats import REFRESH_BATCH_SIZE, refresh_track_stats_for_ids + + +class Command(BaseCommand): + help = "Store track statistics (motion, size ratio, distinct taxa, id agreement) on a project's occurrences." + + def add_arguments(self, parser): + parser.add_argument("--project", type=int, required=True, help="Project ID to backfill.") + parser.add_argument( + "--only-multi-detection", + action="store_true", + default=True, + help="Skip occurrences with a single detection (default).", + ) + parser.add_argument( + "--all-occurrences", + action="store_false", + dest="only_multi_detection", + help="Include single-detection occurrences as well.", + ) + + def handle(self, *args, **options): + project_id: int = options["project"] + only_multi_detection: bool = options["only_multi_detection"] + + try: + project = Project.objects.get(pk=project_id) + except Project.DoesNotExist as err: + raise CommandError(f"Project {project_id} does not exist") from err + + occurrences = Occurrence.objects.filter(project=project) + if only_multi_detection: + occurrences = occurrences.annotate(detection_count=Count("detections")).filter(detection_count__gt=1) + + # Materialize the ids first: the refresh writes to the same rows the filter reads. + ids = list(occurrences.order_by("pk").values_list("pk", flat=True)) + scope = "multi-detection occurrences" if only_multi_detection else "occurrences" + self.stdout.write(f"Project #{project.pk} ({project.name}): {len(ids)} {scope} to refresh.") + + refreshed = 0 + for start in range(0, len(ids), REFRESH_BATCH_SIZE): + refreshed += refresh_track_stats_for_ids(ids[start : start + REFRESH_BATCH_SIZE]) + self.stdout.write(f" {refreshed}/{len(ids)}") + + self.stdout.write(self.style.SUCCESS(f"Stored track statistics for {refreshed} {scope}.")) diff --git a/ami/main/migrations/0096_detection_next_detection.py b/ami/main/migrations/0096_detection_next_detection.py new file mode 100644 index 000000000..f20634eab --- /dev/null +++ b/ami/main/migrations/0096_detection_next_detection.py @@ -0,0 +1,25 @@ +# Generated by Django 4.2.10 on 2026-10-04 12:24 + +import django.db.models.deletion +from django.db import migrations, models + + +class Migration(migrations.Migration): + dependencies = [ + ("main", "0095_grant_sync_deployment_to_mldatamanager"), + ] + + operations = [ + migrations.AddField( + model_name="detection", + name="next_detection", + field=models.OneToOneField( + blank=True, + help_text="The detection that follows this one in the tracking sequence.", + null=True, + on_delete=django.db.models.deletion.SET_NULL, + related_name="previous_detection", + to="main.detection", + ), + ), + ] diff --git a/ami/main/migrations/0097_occurrence_track_stats.py b/ami/main/migrations/0097_occurrence_track_stats.py new file mode 100644 index 000000000..b773a310f --- /dev/null +++ b/ami/main/migrations/0097_occurrence_track_stats.py @@ -0,0 +1,48 @@ +# Generated by Django 4.2.10 on 2026-10-04 12:24 + +from django.db import migrations, models + + +class Migration(migrations.Migration): + dependencies = [ + ("main", "0096_detection_next_detection"), + ] + + operations = [ + migrations.AddField( + model_name="occurrence", + name="track_distinct_taxa", + field=models.IntegerField( + blank=True, + help_text="Distinct taxa among terminal classifications. See models_future/track_stats.py.", + null=True, + ), + ), + migrations.AddField( + model_name="occurrence", + name="track_id_agreement", + field=models.FloatField( + blank=True, + help_text="Share of terminal classifications naming the determination. See models_future/track_stats.py.", + null=True, + ), + ), + migrations.AddField( + model_name="occurrence", + name="track_motion", + field=models.FloatField( + blank=True, + help_text="Path length between captures as a fraction of the image diagonal. See track_stats.py.", + null=True, + ), + ), + migrations.AddField( + model_name="occurrence", + name="track_size_ratio", + field=models.FloatField( + blank=True, + help_text="Largest box area over the smallest. See models_future/track_stats.py.", + null=True, + ), + ), + ] diff --git a/ami/main/models.py b/ami/main/models.py index a1330e88a..391896000 100644 --- a/ami/main/models.py +++ b/ami/main/models.py @@ -1441,6 +1441,28 @@ def update_calculated_fields_for_events( return to_update +def update_calculated_fields_for_sessions_and_stations( + event_ids: typing.Iterable[int | None], stations_async: bool = True +) -> None: + """Refresh the cached counts of these sessions and of the stations they belong to. + + Call once after occurrences are created, merged or split, which neither the + occurrence nor the detection saves do. The station refresh scans the whole station, so by + default it runs in a background task after the transaction commits. + """ + from ami.main.tasks import refresh_deployment_cached_counts + + pks = sorted({pk for pk in event_ids if pk is not None}) + if not pks: + return + update_calculated_fields_for_events(pks=pks) + deployment_ids = list(Deployment.objects.filter(events__pk__in=pks).values_list("pk", flat=True).distinct()) + if stations_async: + transaction.on_commit(lambda: refresh_deployment_cached_counts.delay(deployment_ids)) + else: + refresh_deployment_cached_counts(deployment_ids) + + def audit_event_lengths(deployment: Deployment): logger.info("Checking for unusual event durations") @@ -1648,6 +1670,8 @@ def _group_images_into_events_locked( f"Done grouping {len(image_timestamps)} captures into {len(events)} events " f"for deployment {deployment}" ) + occurrences_split_count = _split_occurrences_at_session_boundaries(deployment, job) + # Realign Occurrence.event_id with each occurrence's detections' current # source_image.event_id. Occurrences are bound to an event once at creation # time (Detection.associate_new_occurrence and Pipeline.save_results both @@ -1728,6 +1752,7 @@ def _group_images_into_events_locked( "Events created": events_created_count, "Events touched": len(touched_event_pks), "Empty events deleted": events_deleted_empty, + "Occurrences split at a session boundary": occurrences_split_count, "Duplicate timestamps": duplicate_timestamp_count, "Ungrouped captures": ungrouped_captures_count, "Captures missing timestamp": no_timestamp_captures_count, @@ -1740,6 +1765,35 @@ def _group_images_into_events_locked( return events +def _split_occurrences_at_session_boundaries(deployment: Deployment, job: "Job | None") -> int: + """Split every occurrence in the deployment whose detections now span several sessions. + + An occurrence is expected to belong to one session, so a regroup that draws a session + boundary through it leaves one piece per session. Returns how many occurrences were split. + """ + from ami.main.models_future.tracks import split_at_session_boundaries + + spanning_ids = list( + Detection.objects.valid() + .filter(occurrence__deployment=deployment) + .values("occurrence_id") + .annotate(sessions=models.Count("source_image__event", distinct=True)) + .filter(sessions__gt=1) + .values_list("occurrence_id", flat=True) + ) + split_count = 0 + for occurrence in Occurrence.objects.filter(pk__in=spanning_ids).order_by("pk"): + pieces = split_at_session_boundaries(occurrence) + if not pieces: + continue + split_count += 1 + (job.logger if job else logger).info( + f"Split occurrence {occurrence.pk} at a session boundary; " + f"new occurrence(s) {[piece.pk for piece in pieces]} hold the later sessions." + ) + return split_count + + def deployment_events_need_update(deployment: Deployment) -> bool: """ Returns True if there are any SourceImages in the deployment @@ -3239,6 +3293,15 @@ class Detection(BaseModel): similarity_vector = models.JSONField(null=True, blank=True) + next_detection = models.OneToOneField( + "self", + on_delete=models.SET_NULL, + null=True, + blank=True, + related_name="previous_detection", + help_text="The detection that follows this one in the tracking sequence.", + ) + # For type hints classifications: models.QuerySet["Classification"] source_image_id: int @@ -3402,8 +3465,17 @@ def valid(self): - Occurrences with determination__isnull=True (no taxonomic identification, same field bug shape) """ + return self.with_real_detections().exclude(determination__isnull=True) + + def with_real_detections(self): + """ + Occurrences backed by at least one real bounding box, determined or not. + + Null-marker sentinels stay excluded exactly as in valid(), since they carry no box. + Used where undetermined occurrences must be reachable, such as the tracks export. + """ has_valid_detection = Exists(Detection.objects.valid().filter(occurrence_id=OuterRef("pk"))) - return self.filter(has_valid_detection).exclude(determination__isnull=True) + return self.filter(has_valid_detection) def with_detections_count(self): return self.annotate(detections_count=models.Count("detections", distinct=True)) @@ -3688,6 +3760,30 @@ class Occurrence(BaseModel): deployment = models.ForeignKey(Deployment, on_delete=models.SET_NULL, null=True, related_name="occurrences") project = models.ForeignKey("Project", on_delete=models.SET_NULL, null=True, related_name="occurrences") + # Statistics stored so occurrences can be sorted by them. Written by + # ``models_future.track_stats.refresh_track_stats`` whenever tracking or a regroup changes + # which detections an occurrence holds; null until then, or when it has none. + track_motion = models.FloatField( + null=True, + blank=True, + help_text="Path length between captures as a fraction of the image diagonal. See track_stats.py.", + ) + track_size_ratio = models.FloatField( + null=True, + blank=True, + help_text="Largest box area over the smallest. See models_future/track_stats.py.", + ) + track_distinct_taxa = models.IntegerField( + null=True, + blank=True, + help_text="Distinct taxa among terminal classifications. See models_future/track_stats.py.", + ) + track_id_agreement = models.FloatField( + null=True, + blank=True, + help_text="Share of terminal classifications naming the determination. See models_future/track_stats.py.", + ) + detections: models.QuerySet[Detection] identifications: models.QuerySet[Identification] diff --git a/ami/main/models_future/track_stats.py b/ami/main/models_future/track_stats.py new file mode 100644 index 000000000..492ff5b98 --- /dev/null +++ b/ami/main/models_future/track_stats.py @@ -0,0 +1,208 @@ +"""Per-occurrence statistics of an insect's path: how far it moved, how much its box changed, +and how consistently it was labelled. + +Four of the numbers are stored on the occurrence row (Occurrence.track_*) so occurrences can +be sorted by them. refresh_track_stats recomputes them in SQL whenever tracking or a regroup +changes which detections an occurrence holds, and backfill_track_stats fills in older rows. + +Definitions: +- frames: detections in the occurrence (not stored). +- motion: path length between consecutive detection centres, ordered by + (timestamp, id), divided by the image diagonal. A stationary insect scores ~0. +- size_ratio: largest bbox area over smallest, areas floored at 1.0. +- distinct_taxa: distinct taxa among terminal classifications, across every + algorithm, so two classifiers that disagree count as two taxa. +- id_agreement: share of terminal classifications naming the determination; + null when there are none. + +The image diagonal is taken from the largest capture width and height seen; when the captures +carry no dimensions it falls back to the farthest bbox corner, and to 1.0 when that is also +unknown. Single-detection occurrences score motion 0.0 and size_ratio 1.0. +""" + +from __future__ import annotations + +import math +from collections.abc import Iterable + +from django.db import connection + +from ami.main.models import Occurrence + +# The stored subset, in the order the Occurrence fields are declared. +TRACK_STAT_FIELDS = ("track_motion", "track_size_ratio", "track_distinct_taxa", "track_id_agreement") + +# Ids per statement when refreshing many rows; keeps the ANY(%s) arrays and the +# CASE expression bulk_update builds to a size Postgres plans quickly. +REFRESH_BATCH_SIZE = 500 + +_MIN_AREA = 1.0 +_ROUND_TO = 4 + + +def frame_diagonal( + max_width: float | None, max_height: float | None, max_x2: float | None, max_y2: float | None +) -> float: + """Length motion is normalised by, in the same pixel units as the bboxes.""" + if max_width and max_height: + return math.hypot(max_width, max_height) + if max_x2 or max_y2: + return math.hypot(max_x2 or 0.0, max_y2 or 0.0) or 1.0 + return 1.0 + + +def size_ratio(min_area: float | None, max_area: float | None) -> float: + if not min_area or not max_area: + return 1.0 + return round(max_area / min_area, _ROUND_TO) + + +def id_agreement(agreeing: int, terminal_count: int) -> float | None: + if terminal_count == 0: + return None + return round(agreeing / terminal_count, _ROUND_TO) + + +# One row per occurrence: detection count, summed centre-to-centre distance in capture order, +# and the extremes the diagonal and size ratio are derived from. The CASE keeps a +# detection without a bbox out of the area extremes instead of flooring it to 1.0. +_GEOMETRY_SQL = """ +WITH frames AS ( + SELECT + d.occurrence_id, + d.id, + d.timestamp, + (d.bbox->>0)::float AS x1, + (d.bbox->>1)::float AS y1, + (d.bbox->>2)::float AS x2, + (d.bbox->>3)::float AS y2, + si.width, + si.height + FROM {detection} d + LEFT JOIN {source_image} si ON si.id = d.source_image_id + WHERE d.occurrence_id = ANY(%s) +), +steps AS ( + SELECT + occurrence_id, + (x1 + x2) / 2.0 AS cx, + (y1 + y2) / 2.0 AS cy, + LAG((x1 + x2) / 2.0) OVER track AS prev_cx, + LAG((y1 + y2) / 2.0) OVER track AS prev_cy, + CASE WHEN x1 IS NULL THEN NULL ELSE GREATEST(ABS((x2 - x1) * (y2 - y1)), {min_area}) END AS area, + x2, + y2, + width, + height + FROM frames + WINDOW track AS (PARTITION BY occurrence_id ORDER BY timestamp, id) +) +SELECT + occurrence_id, + COUNT(*) AS frames, + COALESCE(SUM(SQRT(POWER(cx - prev_cx, 2) + POWER(cy - prev_cy, 2))), 0.0) AS path_length, + MIN(area) AS min_area, + MAX(area) AS max_area, + MAX(width) AS max_width, + MAX(height) AS max_height, + MAX(x2) AS max_x2, + MAX(y2) AS max_y2 +FROM steps +GROUP BY occurrence_id +""" + +# One row per occurrence that has terminal classifications. +_CLASSIFICATION_SQL = """ +SELECT + d.occurrence_id, + COUNT(DISTINCT c.taxon_id) AS distinct_taxa, + COUNT(*) AS terminal_count, + COUNT(*) FILTER (WHERE c.taxon_id = o.determination_id) AS agreeing +FROM {classification} c +JOIN {detection} d ON d.id = c.detection_id +JOIN {occurrence} o ON o.id = d.occurrence_id +WHERE d.occurrence_id = ANY(%s) AND c.terminal +GROUP BY d.occurrence_id +""" + + +def track_stats_for_occurrences(occurrence_ids: list[int]) -> dict[int, dict]: + """Stats keyed by occurrence id, in two statements scoped to the ids given. + + Occurrences with no detections are absent from the result. The cost is proportional + to the detections behind the ids, not to the project, so keep a call to a page or a + refresh batch of ids. + """ + from ami.main.models import Classification, Detection, SourceImage + + if not occurrence_ids: + return {} + + geometry_sql = _GEOMETRY_SQL.format( + detection=Detection._meta.db_table, + source_image=SourceImage._meta.db_table, + min_area=_MIN_AREA, + ) + classification_sql = _CLASSIFICATION_SQL.format( + classification=Classification._meta.db_table, + detection=Detection._meta.db_table, + occurrence=Occurrence._meta.db_table, + ) + + stats: dict[int, dict] = {} + with connection.cursor() as cursor: + cursor.execute(geometry_sql, [list(occurrence_ids)]) + for pk, frames, path_length, min_area, max_area, max_width, max_height, max_x2, max_y2 in cursor.fetchall(): + diagonal = frame_diagonal(max_width, max_height, max_x2, max_y2) + stats[pk] = { + "frames": frames, + "motion": round(float(path_length) / diagonal, _ROUND_TO), + "size_ratio": size_ratio(min_area, max_area), + "distinct_taxa": 0, + "id_agreement": None, + } + + cursor.execute(classification_sql, [list(occurrence_ids)]) + for pk, distinct, terminal_count, agreeing in cursor.fetchall(): + if pk in stats: + stats[pk]["distinct_taxa"] = distinct + stats[pk]["id_agreement"] = id_agreement(agreeing, terminal_count) + + return stats + + +def _apply_stats(occurrence: Occurrence, stats: dict | None) -> None: + occurrence.track_motion = stats["motion"] if stats else None + occurrence.track_size_ratio = stats["size_ratio"] if stats else None + occurrence.track_distinct_taxa = stats["distinct_taxa"] if stats else None + occurrence.track_id_agreement = stats["id_agreement"] if stats else None + + +def refresh_track_stats(*occurrences: Occurrence) -> None: + """Recompute the stored stats for these occurrences from their current detections. + + Three queries however many occurrences are given: the two statements of + ``track_stats_for_occurrences`` and one ``bulk_update``. The instances are updated in + place as well as the rows. An occurrence with no detections is set back to null. + Call after the determination is settled, since ``id_agreement`` is measured against + it, and without going through ``Occurrence.save()``, which would recompute it. + """ + targets = [occurrence for occurrence in occurrences if occurrence.pk is not None] + if not targets: + return + stats = track_stats_for_occurrences([occurrence.pk for occurrence in targets]) + for occurrence in targets: + _apply_stats(occurrence, stats.get(occurrence.pk)) + Occurrence.objects.bulk_update(targets, TRACK_STAT_FIELDS) + + +def refresh_track_stats_for_ids(occurrence_ids: Iterable[int]) -> int: + """Refresh stored stats by id, in batches of ``REFRESH_BATCH_SIZE``; returns the count. + + Writes through ``bulk_update`` so ``updated_at`` is left alone: the list sorts by it + by default, and a backfill must not reorder every occurrence in a project. + """ + ids = list(occurrence_ids) + for start in range(0, len(ids), REFRESH_BATCH_SIZE): + refresh_track_stats(*(Occurrence(pk=pk) for pk in ids[start : start + REFRESH_BATCH_SIZE])) + return len(ids) diff --git a/ami/main/models_future/tracks.py b/ami/main/models_future/tracks.py new file mode 100644 index 000000000..7f3dda946 --- /dev/null +++ b/ami/main/models_future/tracks.py @@ -0,0 +1,116 @@ +"""Operations that keep tracked occurrences consistent with session boundaries. + +Tracking links the detections of one insect through ``Detection.next_detection`` and attaches +the chain to a single occurrence. Chains never cross a session boundary, but regrouping +captures into sessions can draw a new boundary through an existing occurrence. The functions +here lock sessions against concurrent writers and split such an occurrence into one per session. +""" + +from __future__ import annotations + +from collections.abc import Iterable + +from django.db import transaction +from django.db.models import F + +from ami.main.models import Detection, Event, Identification, Occurrence, SourceImage, update_occurrence_determination +from ami.main.models_future.track_stats import refresh_track_stats + +# Order of detections within an occurrence: capture time, then capture, then detection. +# The split and the tracks export both use it, so they agree on what "next" is. +CAPTURE_ORDER = (F("source_image__timestamp").asc(nulls_last=True), "source_image_id", "pk") + + +def lock_sessions(event_ids: Iterable[int | None]) -> None: + """Hold a row lock on each session until the surrounding transaction ends. + + Tracking runs and regroup splits both call this before reading what they change, so one + waits for the other. Locking in id order keeps two writers from deadlocking. + """ + pks = sorted({pk for pk in event_ids if pk is not None}) + if not pks: + return + list(Event.objects.select_for_update().filter(pk__in=pks).order_by("pk").values_list("pk", flat=True)) + + +def _move_to_new_occurrence(occurrence: Occurrence, detections: list[Detection]) -> Occurrence: + """Attach ``detections`` to a new occurrence beside ``occurrence``. + + The new occurrence takes the session of its first detection's capture, which after a + regroup need not be the session of ``occurrence``. + """ + new_occurrence = Occurrence.objects.create( + event_id=detections[0].source_image.event_id, + deployment=occurrence.deployment, + project=occurrence.project, + ) + Detection.objects.filter(pk__in=[d.pk for d in detections]).update(occurrence=new_occurrence) + return new_occurrence + + +@transaction.atomic +def split_at_session_boundaries(occurrence: Occurrence) -> list[Occurrence]: + """Split an occurrence whose detections fall in several sessions into one per session. + + The piece in the earliest session keeps this occurrence and its identifications; each + later piece is a new occurrence holding copies of them. The link between the last + detection of one piece and the first of the next is kept, since tracking stops at session + boundaries and so never walks across it. Returns the new occurrences in time order, or an + empty list when nothing was split. + """ + sessions = SourceImage.objects.filter(detections__occurrence=occurrence).values_list("event_id", flat=True) + lock_sessions([occurrence.event_id, *sessions]) + try: + occurrence.refresh_from_db() + except Occurrence.DoesNotExist: + return [] + detections = occurrence.detections.select_related("source_image").order_by(*CAPTURE_ORDER) + by_session: dict[int, list[Detection]] = {} + for detection in detections: + # A capture with no session stays with the earliest piece. + if detection.source_image.event_id is not None: + by_session.setdefault(detection.source_image.event_id, []).append(detection) + if len(by_session) < 2: + return [] + + earliest_event_id, *later_event_ids = by_session + pieces = [_move_to_new_occurrence(occurrence, by_session[event_id]) for event_id in later_event_ids] + + if occurrence.event_id != earliest_event_id: + occurrence.event_id = earliest_event_id + Occurrence.objects.filter(pk=occurrence.pk).update(event_id=earliest_event_id) + + _copy_identifications(occurrence, pieces) + for piece in [occurrence, *pieces]: + update_occurrence_determination(piece, save=True) + refresh_track_stats(occurrence, *pieces) + return pieces + + +def _copy_identifications(source: Occurrence, targets: list[Occurrence]) -> None: + """Give each target a copy of every identification on ``source``, dated as the original. + + Written with ``bulk_create`` to skip ``Identification.save()``, which would withdraw + the user's other identifications on the target. The caller recomputes determinations. + """ + originals = list(source.identifications.all()) + if not originals or not targets: + return + note = f"Copied from occurrence {source.pk} when regrouping split it at a session boundary." + pairs = [(original, target) for target in targets for original in originals] + copies = Identification.objects.bulk_create( + [ + Identification( + occurrence=target, + user_id=original.user_id, + taxon_id=original.taxon_id, + withdrawn=original.withdrawn, + comment=f"{original.comment}\n{note}" if original.comment else note, + ) + for original, target in pairs + ] + ) + # created_at is auto_now_add, so the original date can only be written after the insert. + for copy, (original, _) in zip(copies, pairs): + copy.created_at = original.created_at + Identification.objects.bulk_update(copies, ["created_at"]) diff --git a/ami/main/tasks.py b/ami/main/tasks.py index 16f927a3f..19ea81321 100644 --- a/ami/main/tasks.py +++ b/ami/main/tasks.py @@ -23,3 +23,12 @@ def refresh_project_cached_counts(project_id: int) -> None: logger.info(f"Refreshing cached counts for project {project.pk} ({project.name})") project.update_related_calculated_fields() + + +@celery_app.task(ignore_result=True) +def refresh_deployment_cached_counts(deployment_ids: list[int]) -> None: + """Refresh the cached counts of these stations after occurrences were created, merged or split.""" + from ami.main.models import Deployment + + for deployment in Deployment.objects.filter(pk__in=deployment_ids): + deployment.update_calculated_fields(save=True) diff --git a/ami/main/tests.py b/ami/main/tests.py index ef58ce8d9..070cd70e4 100644 --- a/ami/main/tests.py +++ b/ami/main/tests.py @@ -1,5 +1,6 @@ import copy import datetime +import io import logging import typing from io import BytesIO @@ -8538,3 +8539,233 @@ def test_identifications_browsable_page(self): html = self._get_html("/api/v2/identifications/") self._assert_number_input(html, "occurrence") self._assert_number_input(html, "taxon") + + +class TestRegroupSplitsOccurrences(TestCase): + """Regrouping never leaves one occurrence spanning two sessions. + + The captures are two bursts three hours apart: one session under a 6-hour gap, two + under a 2-hour gap. An occurrence built across all of them is what an earlier grouping + leaves behind when a later regroup draws a boundary through it. + """ + + def setUp(self) -> None: + self.project, self.deployment = setup_test_project(reuse=False) + create_taxa(project=self.project) + self.taxon, self.other_taxon = list(Taxon.objects.filter(projects=self.project).order_by("pk")[:2]) + self.user = User.objects.create_user(email="regroup-identifier@insectai.org") # type: ignore[attr-defined] + start = datetime.datetime(2024, 6, 1, 22, 0) + self.captures = [ + SourceImage.objects.create( + deployment=self.deployment, + timestamp=start + datetime.timedelta(minutes=minutes), + path=f"test/regroup-split-{i}.jpg", + width=640, + height=480, + ) + for i, minutes in enumerate([0, 1, 2, 180, 181, 182]) + ] + + def _group(self, gap_hours: int) -> list[Event]: + group_images_into_events(self.deployment, max_time_gap=datetime.timedelta(hours=gap_hours)) + for capture in self.captures: + capture.refresh_from_db() + return list(Event.objects.filter(deployment=self.deployment).order_by("start")) + + def _make_occurrence(self, captures: list[SourceImage]) -> tuple[Occurrence, list[Detection]]: + occurrence = Occurrence.objects.create( + event=captures[0].event, deployment=self.deployment, project=self.project + ) + detections = [] + for capture in captures: + detection = Detection.objects.create( + source_image=capture, timestamp=capture.timestamp, bbox=[10, 10, 40, 40], occurrence=occurrence + ) + detection.classifications.create(taxon=self.taxon, score=0.9, timestamp=capture.timestamp) + detections.append(detection) + for earlier, later in zip(detections, detections[1:]): + earlier.next_detection = later + earlier.save(update_fields=["next_detection"]) + occurrence.save() + return occurrence, detections + + def _split_one_occurrence(self) -> tuple[Occurrence, Occurrence, list[Detection], list[Event]]: + self._group(gap_hours=6) + occurrence, detections = self._make_occurrence(self.captures) + events = self._group(gap_hours=2) + self.assertEqual(len(events), 2) + piece = Occurrence.objects.exclude(pk=occurrence.pk).get(deployment=self.deployment) + occurrence.refresh_from_db() + return occurrence, piece, detections, events + + def _detection_ids(self, occurrence: Occurrence) -> list[int]: + return list(occurrence.detections.order_by("timestamp").values_list("pk", flat=True)) + + def test_an_occurrence_across_a_new_boundary_is_split_into_one_per_session(self): + occurrence, piece, detections, (first, second) = self._split_one_occurrence() + + self.assertEqual(self._detection_ids(occurrence), [d.pk for d in detections[:3]]) + self.assertEqual(self._detection_ids(piece), [d.pk for d in detections[3:]]) + self.assertEqual(occurrence.event_id, first.pk) + self.assertEqual(piece.event_id, second.pk) + self.assertEqual(piece.determination_id, self.taxon.pk) + self.assertEqual(first.occurrences_count, 1) + self.assertEqual(second.occurrences_count, 1) + self.assertEqual(piece.track_motion, 0.0) + + def test_the_link_between_the_pieces_is_kept(self): + _, _, detections, _ = self._split_one_occurrence() + + boundary = Detection.objects.get(pk=detections[2].pk) + self.assertEqual(boundary.next_detection_id, detections[3].pk) + + def test_merging_sessions_leaves_occurrences_untouched(self): + self._group(gap_hours=2) + early, early_detections = self._make_occurrence(self.captures[:3]) + late, late_detections = self._make_occurrence(self.captures[3:]) + + (merged,) = self._group(gap_hours=6) + + self.assertEqual(Occurrence.objects.filter(deployment=self.deployment).count(), 2) + self.assertEqual(self._detection_ids(early), [d.pk for d in early_detections]) + self.assertEqual(self._detection_ids(late), [d.pk for d in late_detections]) + self.assertEqual( + set(Occurrence.objects.filter(deployment=self.deployment).values_list("event_id", flat=True)), + {merged.pk}, + ) + + def test_identifications_are_copied_to_every_piece(self): + self._group(gap_hours=6) + occurrence, _ = self._make_occurrence(self.captures) + superseded = Identification.objects.create(occurrence=occurrence, user=self.user, taxon=self.taxon) + current = Identification.objects.create( + occurrence=occurrence, user=self.user, taxon=self.other_taxon, comment="Wing pattern checked." + ) + + self._group(gap_hours=2) + + occurrence.refresh_from_db() + piece = Occurrence.objects.exclude(pk=occurrence.pk).get(deployment=self.deployment) + self.assertEqual(set(occurrence.identifications.values_list("pk", flat=True)), {superseded.pk, current.pk}) + note = f"Copied from occurrence {occurrence.pk} when regrouping split it at a session boundary." + copies = {(i.taxon_id, i.user_id, i.withdrawn, i.created_at, i.comment) for i in piece.identifications.all()} + superseded.refresh_from_db() + self.assertEqual( + copies, + { + (self.taxon.pk, self.user.pk, True, superseded.created_at, note), + (self.other_taxon.pk, self.user.pk, False, current.created_at, f"Wing pattern checked.\n{note}"), + }, + ) + self.assertFalse( + piece.identifications.filter( + models.Q(agreed_with_identification__isnull=False) | models.Q(agreed_with_prediction__isnull=False) + ).exists() + ) + self.assertEqual(occurrence.determination_id, self.other_taxon.pk) + self.assertEqual(piece.determination_id, self.other_taxon.pk) + + +class TrackStatsTestCase(TestCase): + """The stored occurrence statistics follow their documented definitions and refresh in bulk.""" + + FRAME_WIDTH = 300 + FRAME_HEIGHT = 400 # a 300x400 image has a diagonal of exactly 500 + + def setUp(self) -> None: + self.project, self.deployment = setup_test_project(reuse=False) + create_taxa(project=self.project) + create_captures(deployment=self.deployment, num_nights=1, images_per_night=3, interval_minutes=1) + SourceImage.objects.filter(deployment=self.deployment).update(width=self.FRAME_WIDTH, height=self.FRAME_HEIGHT) + self.captures = list(SourceImage.objects.filter(deployment=self.deployment).order_by("timestamp")) + self.event = self.captures[0].event + taxa = list(Taxon.objects.filter(projects=self.project).order_by("pk")[:3]) + assert len(taxa) == 3, "Fixture must provide three taxa to disagree with" + self.taxon_a, self.taxon_b, self.taxon_c = taxa + + # Centres (5,5) -> (35,45) -> (65,85): two steps of 50 px over a 500 px diagonal. + # Areas 100, 100, 400. Terminal labels A, A, B plus a non-terminal C to ignore. + self.multi = self._make_occurrence( + [ + ([0, 0, 10, 10], [(self.taxon_a, 0.9, True)]), + ([30, 40, 40, 50], [(self.taxon_a, 0.85, True)]), + ([55, 75, 75, 95], [(self.taxon_b, 0.8, True), (self.taxon_c, 0.5, False)]), + ] + ) + self.single = self._make_occurrence([([10, 10, 20, 20], [(self.taxon_a, 0.7, True)])]) + + def _make_occurrence(self, frames: list[tuple[list[int], list[tuple[Taxon, float, bool]]]]) -> Occurrence: + occurrence = Occurrence.objects.create(event=self.event, deployment=self.deployment, project=self.project) + for capture, (bbox, labels) in zip(self.captures, frames): + detection = Detection.objects.create( + source_image=capture, timestamp=capture.timestamp, bbox=bbox, occurrence=occurrence + ) + for taxon, score, terminal in labels: + detection.classifications.create( + taxon=taxon, score=score, timestamp=capture.timestamp, terminal=terminal + ) + # Pin the determination the agreement is measured against, independent of how save() picks one. + Occurrence.objects.filter(pk=occurrence.pk).update(determination=self.taxon_a, determination_score=0.9) + return occurrence + + def _stored(self, occurrence: Occurrence) -> dict: + from ami.main.models_future.track_stats import TRACK_STAT_FIELDS + + return Occurrence.objects.filter(pk=occurrence.pk).values(*TRACK_STAT_FIELDS).get() + + def test_stats_follow_the_documented_definitions(self): + from ami.main.models_future.track_stats import refresh_track_stats + + refresh_track_stats(self.multi, self.single) + + stored = self._stored(self.multi) + self.assertAlmostEqual(stored["track_motion"], 100 / 500, places=4) + self.assertAlmostEqual(stored["track_size_ratio"], 4.0, places=4) + self.assertEqual(stored["track_distinct_taxa"], 2, "A non-terminal classification must not count as a taxon") + self.assertAlmostEqual(stored["track_id_agreement"], 2 / 3, places=4) + self.assertEqual( + self._stored(self.single), + {"track_motion": 0.0, "track_size_ratio": 1.0, "track_distinct_taxa": 1, "track_id_agreement": 1.0}, + ) + + def test_stats_are_null_until_stored_and_the_refresh_updates_the_instance(self): + from ami.main.models_future.track_stats import refresh_track_stats + + self.assertEqual(set(self._stored(self.multi).values()), {None}) + refresh_track_stats(self.multi) + self.assertEqual(self.multi.track_motion, 0.2) + + def test_refresh_is_three_queries_however_many_occurrences(self): + from cachalot.api import cachalot_disabled + + from ami.main.models_future.track_stats import refresh_track_stats + + with cachalot_disabled(), CaptureQueriesContext(connection) as ctx: + refresh_track_stats(self.multi, self.single) + self.assertEqual(len(ctx.captured_queries), 3) + + def test_an_occurrence_with_no_detections_is_set_back_to_null(self): + from ami.main.models_future.track_stats import refresh_track_stats + + refresh_track_stats(self.multi) + self.multi.detections.update(occurrence=None) + refresh_track_stats(self.multi) + self.assertEqual(set(self._stored(self.multi).values()), {None}) + + def test_the_backfill_command_stores_stats_for_multi_detection_occurrences(self): + from django.core.management import call_command + + out = io.StringIO() + call_command("backfill_track_stats", project=self.project.pk, stdout=out) + self.assertIn("1 multi-detection occurrences", out.getvalue()) + self.assertEqual(self._stored(self.multi)["track_motion"], 0.2) + self.assertIsNone(self._stored(self.single)["track_motion"], "Single detections are skipped by default") + + call_command("backfill_track_stats", project=self.project.pk, only_multi_detection=False, stdout=out) + self.assertEqual(self._stored(self.single)["track_motion"], 0.0) + + def test_the_backfill_command_refuses_an_unknown_project(self): + from django.core.management import CommandError, call_command + + with self.assertRaises(CommandError): + call_command("backfill_track_stats", project=0) From 48e020564cf21df13a8d6fa7ed979612fe87f892 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Sun, 4 Oct 2026 09:28:28 -0700 Subject: [PATCH 02/28] feat(exports): add a CSV export with one row per detection of each occurrence Add the tracks_csv export format and an export_tracks management command. Each row is one detection with its occurrence, session, capture time, position in the occurrence, bounding box, image size, best label, and the id of the next detection in the chain. The file is meant for inspecting how tracking grouped detections and for building a benchmark, and the API export and the command share one column definition so their files compare directly. Occurrences are read in chunks so the query count grows with the number of chunks and not with the number of rows. Co-Authored-By: Claude Opus 5.5 --- ami/exports/format_types.py | 21 +++ .../migrations/0002_add_tracks_csv_format.py | 24 +++ ami/exports/registry.py | 1 + ami/exports/tests.py | 150 ++++++++++++++++ ami/exports/tracks.py | 161 ++++++++++++++++++ ami/main/management/commands/export_tracks.py | 46 +++++ 6 files changed, 403 insertions(+) create mode 100644 ami/exports/migrations/0002_add_tracks_csv_format.py create mode 100644 ami/exports/tracks.py create mode 100644 ami/main/management/commands/export_tracks.py diff --git a/ami/exports/format_types.py b/ami/exports/format_types.py index a3f4c82d0..87bab8f91 100644 --- a/ami/exports/format_types.py +++ b/ami/exports/format_types.py @@ -252,3 +252,24 @@ def export(self): self.update_job_progress(records_exported) self.update_export_stats(file_temp_path=temp_file.name) return temp_file.name # Return the file path + + +class TracksCSVExporter(BaseExporter): + """One row per detection of every occurrence in scope; see ami/exports/tracks.py for the columns.""" + + file_format = "csv" + + def get_queryset(self): + return Occurrence.objects.with_real_detections().filter(project=self.project) # type: ignore[union-attr] + + def export(self): + from ami.exports.tracks import write_tracks_csv + + temp_file = tempfile.NamedTemporaryFile(delete=False, suffix=".csv", mode="w", newline="", encoding="utf-8") + with open(temp_file.name, "w", newline="", encoding="utf-8") as csvfile: + rows = write_tracks_csv(self.queryset, csvfile, on_chunk=self.update_job_progress) + self.update_export_stats(file_temp_path=temp_file.name) + # A tracks file holds one record per detection, not per occurrence. + self.data_export.record_count = rows + self.data_export.save(update_fields=["record_count"]) + return temp_file.name diff --git a/ami/exports/migrations/0002_add_tracks_csv_format.py b/ami/exports/migrations/0002_add_tracks_csv_format.py new file mode 100644 index 000000000..0a48301da --- /dev/null +++ b/ami/exports/migrations/0002_add_tracks_csv_format.py @@ -0,0 +1,24 @@ +# Generated by Django 4.2.10 on 2026-10-04 12:24 + +from django.db import migrations, models + + +class Migration(migrations.Migration): + dependencies = [ + ("exports", "0001_initial"), + ] + + operations = [ + migrations.AlterField( + model_name="dataexport", + name="format", + field=models.CharField( + choices=[ + ("occurrences_api_json", "occurrences_api_json"), + ("occurrences_simple_csv", "occurrences_simple_csv"), + ("tracks_csv", "tracks_csv"), + ], + max_length=255, + ), + ), + ] diff --git a/ami/exports/registry.py b/ami/exports/registry.py index 29a4cc0e7..8574c77c6 100644 --- a/ami/exports/registry.py +++ b/ami/exports/registry.py @@ -27,3 +27,4 @@ def get_supported_formats(cls): ExportRegistry.register("occurrences_api_json")(format_types.JSONExporter) ExportRegistry.register("occurrences_simple_csv")(format_types.CSVExporter) +ExportRegistry.register("tracks_csv")(format_types.TracksCSVExporter) diff --git a/ami/exports/tests.py b/ami/exports/tests.py index 866b1af61..48eaf75b1 100644 --- a/ami/exports/tests.py +++ b/ami/exports/tests.py @@ -1,4 +1,5 @@ import csv +import io import json import logging @@ -551,3 +552,152 @@ def test_csv_has_all_new_fields(self): ] for field in expected_fields: self.assertIn(field, headers, f"Missing CSV field: {field}") + + +class TracksExportTest(TestCase): + """The tracks CSV: one row per detection, a fixed column contract, and bounded queries.""" + + EXPECTED_HEADER = ( + "occurrence_id,detection_id,event_id,deployment_id,source_image_id,timestamp,detection_index," + "detection_count,bbox_x1,bbox_y1,bbox_x2,bbox_y2,image_width,image_height,detection_label," + "detection_score,occurrence_determination,occurrence_determination_score,next_detection_id" + ) + + def setUp(self): + self.project, self.deployment = setup_test_project(reuse=False) + self.user = self.project.owner + create_captures(deployment=self.deployment, num_nights=1, images_per_night=3, interval_minutes=1) + group_images_into_events(self.deployment) + create_taxa(self.project) + self.taxon = Taxon.objects.filter(projects=self.project).first() + self.algorithm, _ = Algorithm.objects.get_or_create( + name="test-classifier", defaults={"key": "test-classifier"} + ) + self.captures = list(self.project.captures.order_by("timestamp")) + self.captures[0].width, self.captures[0].height = 4096, 2160 + self.captures[0].save() + # Three occurrences of three detections each, created latest capture first so that + # the export cannot get capture order from detection pks. + self.occurrences = [self._make_occurrence(offset=i * 100) for i in range(3)] + + def _make_occurrence(self, offset: int) -> Occurrence: + occurrence = Occurrence.objects.create( + project=self.project, + deployment=self.deployment, + event=self.captures[0].event, + determination=self.taxon, + determination_score=0.9, + ) + for capture in reversed(self.captures): + detection = Detection.objects.create( + source_image=capture, + timestamp=capture.timestamp, + bbox=[offset, offset, offset + 10, offset + 20], + occurrence=occurrence, + ) + detection.classifications.create( + taxon=self.taxon, score=0.8, timestamp=capture.timestamp, algorithm=self.algorithm, terminal=True + ) + return occurrence + + def _rows(self, occurrences=None, **kwargs) -> list[dict[str, str]]: + from ami.exports.tracks import iter_track_rows + + return list(iter_track_rows(occurrences or Occurrence.objects.filter(project=self.project), **kwargs)) + + def test_format_export_header_is_the_contract(self): + data_export = DataExport.objects.create(user=self.user, project=self.project, format="tracks_csv") + file_path = data_export.run_export().replace("/media/", "") + with default_storage.open(file_path, "r") as f: + lines = f.read().splitlines() + default_storage.delete(file_path) + data_export.refresh_from_db() + + self.assertEqual(lines[0], self.EXPECTED_HEADER) + self.assertEqual(len(lines) - 1, 9) + self.assertEqual(data_export.record_count, 9, "A tracks export counts detection rows") + + def test_detection_index_follows_capture_time(self): + rows = [row for row in self._rows() if row["occurrence_id"] == str(self.occurrences[0].pk)] + by_capture = {int(row["source_image_id"]): row for row in rows} + for index, capture in enumerate(self.captures): + row = by_capture[capture.pk] + self.assertEqual(row["detection_index"], str(index)) + self.assertEqual(row["detection_count"], "3") + self.assertEqual(row["timestamp"], capture.timestamp.isoformat()) + first = by_capture[self.captures[0].pk] + self.assertEqual((first["image_width"], first["image_height"]), ("4096", "2160")) + self.assertEqual((first["bbox_x1"], first["bbox_y2"]), ("0", "20")) + self.assertEqual((first["detection_label"], first["detection_score"]), (self.taxon.name, "0.8")) + self.assertEqual(by_capture[self.captures[1].pk]["image_width"], "") + + def test_next_detection_id_carries_the_chain(self): + detections = list(self.occurrences[0].detections.order_by("source_image__timestamp")) + detections[0].next_detection = detections[1] + detections[0].save(update_fields=["next_detection"]) + + rows = {int(row["detection_id"]): row for row in self._rows()} + self.assertEqual(rows[detections[0].pk]["next_detection_id"], str(detections[1].pk)) + self.assertEqual(rows[detections[1].pk]["next_detection_id"], "") + + def test_query_count_is_one_pair_per_chunk(self): + from cachalot.api import cachalot_disabled + from django.db import connection + from django.test.utils import CaptureQueriesContext + + # Three occurrences in chunks of two: occurrences, detections, occurrences, detections. + with cachalot_disabled(), CaptureQueriesContext(connection) as ctx: + rows = self._rows(chunk_size=2) + self.assertEqual(len(rows), 9) + self.assertEqual(len(ctx.captured_queries), 4) + + # Doubling the detections per occurrence adds no queries. + for occurrence in self.occurrences: + for detection in list(occurrence.detections.all()): + detection.pk = None + detection.next_detection = None + detection.save() + with cachalot_disabled(), CaptureQueriesContext(connection) as ctx: + rows = self._rows(chunk_size=2) + self.assertEqual(len(rows), 18) + self.assertEqual(len(ctx.captured_queries), 4) + + def _run_command(self, **options) -> tuple[list[dict[str, str]], str]: + from django.core.management import call_command + + stdout, stderr = io.StringIO(), io.StringIO() + call_command("export_tracks", project=self.project.pk, stdout=stdout, stderr=stderr, **options) + lines = stdout.getvalue().splitlines() + self.assertEqual(lines[0], self.EXPECTED_HEADER) + return list(csv.DictReader(lines)), stderr.getvalue() + + def test_management_command_writes_the_same_csv(self): + rows, summary = self._run_command() + self.assertEqual(len(rows), 9) + self.assertIn("Wrote 9 detection rows", summary) + + def test_management_command_event_filter(self): + import datetime + + from ami.main.models import Event + + first_event = self.captures[0].event + other_event = Event.objects.create( + project=self.project, + deployment=self.deployment, + group_by="2030-01-01", + start=datetime.datetime(2030, 1, 1, 22, 0), + ) + moved = self.occurrences[2] + moved.event = other_event + moved.save(update_fields=["event"]) + + rows, _ = self._run_command(events=[first_event.pk]) + self.assertEqual( + {row["occurrence_id"] for row in rows}, {str(self.occurrences[0].pk), str(self.occurrences[1].pk)} + ) + self.assertEqual(len(rows), 6) + + rows, _ = self._run_command(events=[first_event.pk, other_event.pk]) + self.assertEqual({row["occurrence_id"] for row in rows}, {str(o.pk) for o in self.occurrences}) + self.assertEqual(len(rows), 9) diff --git a/ami/exports/tracks.py b/ami/exports/tracks.py new file mode 100644 index 000000000..aba6f330a --- /dev/null +++ b/ami/exports/tracks.py @@ -0,0 +1,161 @@ +""" +One row per detection of every occurrence in scope, for inspecting and benchmarking tracking. + +The export format and the ``export_tracks`` management command both write through +``iter_track_rows()``, so ``TRACKS_CSV_COLUMNS`` is the single definition of the columns. +""" + +import csv +import datetime +import typing +from collections.abc import Callable, Iterator + +from django.db import models +from django.db.models import OuterRef, Subquery + +from ami.main.models import BEST_MACHINE_PREDICTION_ORDER, Classification, Detection, Occurrence +from ami.main.models_future.tracks import CAPTURE_ORDER + +TRACKS_CSV_COLUMNS: typing.Final = ( + "occurrence_id", + "detection_id", + "event_id", + "deployment_id", + "source_image_id", + "timestamp", + "detection_index", + "detection_count", + "bbox_x1", + "bbox_y1", + "bbox_x2", + "bbox_y2", + "image_width", + "image_height", + "detection_label", + "detection_score", + "occurrence_determination", + "occurrence_determination_score", + "next_detection_id", +) + +DEFAULT_CHUNK_SIZE: typing.Final = 500 + + +def _cell(value) -> str: + """Render one value the way every tracks CSV writes it: blank for None, lowercase booleans.""" + if value is None: + return "" + if isinstance(value, bool): + return "true" if value else "false" + if isinstance(value, datetime.datetime): + return value.isoformat() + return str(value) + + +def _detections_for(occurrence_ids: list[int]) -> models.QuerySet: + """Real detections of the given occurrences, in capture order, with their own best label.""" + best_classification = Classification.objects.filter(detection=OuterRef("pk")).order_by( + *BEST_MACHINE_PREDICTION_ORDER + ) + return ( + Detection.objects.valid() # type: ignore[attr-defined] Custom queryset method + .filter(occurrence_id__in=occurrence_ids) + .annotate( + label=Subquery(best_classification.values("taxon__name")[:1]), + label_score=Subquery(best_classification.values("score")[:1]), + ) + .order_by("occurrence_id", *CAPTURE_ORDER) + .values( + "pk", + "occurrence_id", + "bbox", + "next_detection_id", + "source_image_id", + "source_image__timestamp", + "source_image__width", + "source_image__height", + "source_image__event_id", + "source_image__deployment_id", + "label", + "label_score", + ) + ) + + +def iter_track_rows( + occurrences: models.QuerySet[Occurrence], + chunk_size: int = DEFAULT_CHUNK_SIZE, + on_chunk: Callable[[int], None] | None = None, +) -> Iterator[dict[str, str]]: + """ + Yield one row per detection of each occurrence in ``occurrences``, keyed by TRACKS_CSV_COLUMNS. + + Occurrences are read in pk order, ``chunk_size`` at a time, with one detection query per + chunk, so the query count grows with the number of chunks and never with the rows. + ``on_chunk`` receives the running count of occurrences read, for progress reporting. + """ + scope = ( + occurrences.order_by("pk") + .values("pk", "determination__name", "determination_score") + .distinct() # A capture-set filter joins through detections and repeats occurrences. + ) + occurrences_read = 0 + last_pk = 0 + while True: + chunk = list(scope.filter(pk__gt=last_pk)[:chunk_size]) + if not chunk: + break + last_pk = chunk[-1]["pk"] + occurrences_read += len(chunk) + + detections_by_occurrence: dict[int, list[dict]] = {} + for detection in _detections_for([row["pk"] for row in chunk]): + detections_by_occurrence.setdefault(detection["occurrence_id"], []).append(detection) + + for occurrence in chunk: + detections = detections_by_occurrence.get(occurrence["pk"], []) + for detection_index, detection in enumerate(detections): + bbox = detection["bbox"] + x1, y1, x2, y2 = bbox if isinstance(bbox, list) and len(bbox) == 4 else (None,) * 4 + yield { + "occurrence_id": _cell(occurrence["pk"]), + "detection_id": _cell(detection["pk"]), + "event_id": _cell(detection["source_image__event_id"]), + "deployment_id": _cell(detection["source_image__deployment_id"]), + "source_image_id": _cell(detection["source_image_id"]), + "timestamp": _cell(detection["source_image__timestamp"]), + "detection_index": _cell(detection_index), + "detection_count": _cell(len(detections)), + "bbox_x1": _cell(x1), + "bbox_y1": _cell(y1), + "bbox_x2": _cell(x2), + "bbox_y2": _cell(y2), + "image_width": _cell(detection["source_image__width"]), + "image_height": _cell(detection["source_image__height"]), + "detection_label": _cell(detection["label"]), + "detection_score": _cell(detection["label_score"]), + "occurrence_determination": _cell(occurrence["determination__name"]), + "occurrence_determination_score": _cell(occurrence["determination_score"]), + "next_detection_id": _cell(detection["next_detection_id"]), + } + + if on_chunk: + on_chunk(occurrences_read) + if len(chunk) < chunk_size: + break + + +def write_tracks_csv( + occurrences: models.QuerySet[Occurrence], + stream: typing.TextIO, + chunk_size: int = DEFAULT_CHUNK_SIZE, + on_chunk: Callable[[int], None] | None = None, +) -> int: + """Write the header and every detection row to ``stream``; return the number of detection rows.""" + writer = csv.DictWriter(stream, fieldnames=TRACKS_CSV_COLUMNS) + writer.writeheader() + rows = 0 + for row in iter_track_rows(occurrences, chunk_size=chunk_size, on_chunk=on_chunk): + writer.writerow(row) + rows += 1 + return rows diff --git a/ami/main/management/commands/export_tracks.py b/ami/main/management/commands/export_tracks.py new file mode 100644 index 000000000..32feb13b3 --- /dev/null +++ b/ami/main/management/commands/export_tracks.py @@ -0,0 +1,46 @@ +""" +Write a project's occurrences as CSV, one row per detection, for inspecting and benchmarking tracking. + +The columns are the ``tracks_csv`` export format's (``ami/exports/tracks.py``), so a file from +this command and a file from the export API can be compared directly. Read-only. +""" + +from django.core.management.base import BaseCommand, CommandError + +from ami.exports.tracks import write_tracks_csv +from ami.main.models import Occurrence, Project + + +class Command(BaseCommand): + help = "Write a project's occurrences (one row per detection) as CSV to a file or stdout." + + def add_arguments(self, parser): + parser.add_argument("--project", type=int, required=True, help="Project ID to export.") + parser.add_argument( + "--event", + type=int, + action="append", + dest="events", + default=[], + help="Only occurrences of this session (event) ID. Repeat to export several.", + ) + parser.add_argument("--output", "-o", default="-", help="File path to write, or - for stdout (default).") + + def handle(self, *args, **options): + project_id: int = options["project"] + if not Project.objects.filter(pk=project_id).exists(): + raise CommandError(f"Project {project_id} does not exist") + + occurrences = Occurrence.objects.with_real_detections() # type: ignore[union-attr] + occurrences = occurrences.filter(project_id=project_id) + if options["events"]: + occurrences = occurrences.filter(event_id__in=options["events"]) + + output: str = options["output"] + if output == "-": + rows = write_tracks_csv(occurrences, self.stdout) + else: + with open(output, "w", newline="", encoding="utf-8") as stream: + rows = write_tracks_csv(occurrences, stream) + # Keep stdout pure CSV when the file goes there; the summary goes to stderr. + (self.stderr if output == "-" else self.stdout).write(f"Wrote {rows} detection rows.") From b5c5add0fe92574823987192bfe02521947d65d0 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Sun, 4 Oct 2026 09:34:39 -0700 Subject: [PATCH 03/28] feat(post-processing): generate admin action forms from a task's config schema Operators set a post-processing task's options on an admin confirmation page. Until now each task hand-wrote a Django form that repeated the labels, help text, defaults and bounds already declared on its pydantic config. A task with many options would have to keep the two in step. schema_form_fields builds Django form fields from a pydantic config class: bool, int and float fields and optional versions of them, using the field title as the label, the description as help text, the default as the initial value, and ge/le as min and max. Optional fields are not required and a blank value becomes None. Strict limits (gt/lt) are not expressible as form bounds, so they stay in the schema, whose error the admin action already maps back onto the field. SchemaActionForm wraps it for a task that names its schema and the scope fields to leave out. The existing class masking and size filter forms are unchanged. Co-Authored-By: Claude Opus 5.5 --- ami/ml/post_processing/admin/forms.py | 51 +++++++++++++++++++++++++++ 1 file changed, 51 insertions(+) diff --git a/ami/ml/post_processing/admin/forms.py b/ami/ml/post_processing/admin/forms.py index c9c162808..ed05e90f0 100644 --- a/ami/ml/post_processing/admin/forms.py +++ b/ami/ml/post_processing/admin/forms.py @@ -9,6 +9,10 @@ """ from __future__ import annotations +from collections.abc import Collection +from typing import Any + +import pydantic from django import forms @@ -33,3 +37,50 @@ def __init__(self, *args, scope_queryset=None, **kwargs): def to_config(self) -> dict: """Return ``cleaned_data`` shaped for ``Job.params['config']``.""" return dict(self.cleaned_data) + + +def schema_form_fields(schema: type[pydantic.BaseModel], exclude: Collection[str] = ()) -> dict[str, forms.Field]: + """Build Django form fields from a pydantic (v1) config schema. + + The schema stays the single source of truth for defaults, titles, help text and numeric bounds. + Supports ``bool``, ``int`` and ``float`` fields and optional versions of them; an optional field is + not required and a blank value becomes ``None``. Strict limits (``gt``/``lt``) are left to the + schema, whose error the admin action shows on the field. + """ + fields: dict[str, forms.Field] = {} + for name, model_field in schema.__fields__.items(): + if name in exclude: + continue + info = model_field.field_info + kwargs: dict[str, Any] = { + "label": info.title or name.replace("_", " ").capitalize(), + "help_text": info.description or "", + "initial": model_field.default, + } + python_type = model_field.type_ + if python_type is bool: # bool is checked first because it subclasses int + fields[name] = forms.BooleanField(required=False, **kwargs) + continue + if issubclass(python_type, int): + field_class: type[forms.Field] = forms.IntegerField + elif issubclass(python_type, float): + field_class = forms.FloatField + else: + raise TypeError(f"No form field for {schema.__name__}.{name} of type {python_type!r}") + if info.ge is not None: + kwargs["min_value"] = info.ge + if info.le is not None: + kwargs["max_value"] = info.le + fields[name] = field_class(required=model_field.required is True and not model_field.allow_none, **kwargs) + return fields + + +class SchemaActionForm(BasePostProcessingActionForm): + """Action form whose fields are generated from ``schema``, minus the scope fields in ``exclude_fields``.""" + + schema: type[pydantic.BaseModel] + exclude_fields: Collection[str] = () + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.fields.update(schema_form_fields(self.schema, self.exclude_fields)) From 8711c7b6e40a6c2a9ee35bcf2a811dcdb023bfe0 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Sun, 4 Oct 2026 09:34:40 -0700 Subject: [PATCH 04/28] feat(tracking): run occurrence tracking as a post-processing task from the admin Tracking links each detection to the matching detection in the next capture and folds every chain into one occurrence per session. It is now a registered post-processing task, started from the Django admin on capture sets and on sessions (events), through the shared action factory. Every cost setting is a field on TrackingConfig and the confirmation form is generated from it, so the help text operators read is the schema's description. The defaults reproduce the plain sum of (1 - overlap) + (1 - size ratio) + distance / image diagonal with a cutoff of 1.0. Four optional limits (minimum overlap, minimum size ratio, maximum distance, maximum time between captures) rule a pair out entirely when enabled. The cost defaults are starting points pending experiments and are tuned for captures about 20 seconds apart. Matching runs over processed captures only, meaning captures with at least one detection row, so an unprocessed capture no longer separates its neighbours and leaves a sampled session with nothing to compare. Chains stop at a session boundary. Identifications move onto the surviving occurrence before merged ones are deleted. When a merge changes an occurrence's determination, a terminal classification by the tracking algorithm records the winning prediction in applied_to. Statistics are stored for every settled occurrence, session and station counts are refreshed, and each session runs under a session lock and its own transaction. The sessions changelist creates one job per project. Co-Authored-By: Claude Opus 5.5 --- ami/main/admin.py | 22 +- ami/ml/post_processing/__init__.py | 1 + .../post_processing/admin/tracking_actions.py | 65 ++ ami/ml/post_processing/admin/tracking_form.py | 15 + ami/ml/post_processing/registry.py | 2 + .../tests/test_tracking_admin.py | 146 +++++ .../tests/test_tracking_task.py | 341 ++++++++++ ami/ml/post_processing/tracking_task.py | 595 ++++++++++++++++++ ami/tests/fixtures/tracking.py | 61 ++ 9 files changed, 1247 insertions(+), 1 deletion(-) create mode 100644 ami/ml/post_processing/admin/tracking_actions.py create mode 100644 ami/ml/post_processing/admin/tracking_form.py create mode 100644 ami/ml/post_processing/tests/test_tracking_admin.py create mode 100644 ami/ml/post_processing/tests/test_tracking_task.py create mode 100644 ami/ml/post_processing/tracking_task.py create mode 100644 ami/tests/fixtures/tracking.py diff --git a/ami/main/admin.py b/ami/main/admin.py index 47b383a00..b73fd1ea5 100644 --- a/ami/main/admin.py +++ b/ami/main/admin.py @@ -16,8 +16,11 @@ from ami.ml.post_processing.admin.actions import make_post_processing_action from ami.ml.post_processing.admin.class_masking_form import ClassMaskingActionForm from ami.ml.post_processing.admin.small_size_filter_form import SmallSizeFilterActionForm +from ami.ml.post_processing.admin.tracking_actions import build_tracking_jobs_for_events +from ami.ml.post_processing.admin.tracking_form import TrackingActionForm from ami.ml.post_processing.class_masking import ClassMaskingTask from ami.ml.post_processing.small_size_filter import SmallSizeFilterTask +from ami.ml.post_processing.tracking_task import TrackingTask from ami.ml.tasks import remove_duplicate_classifications from .models import ( @@ -314,8 +317,16 @@ def update_calculated_fields(self, request: HttpRequest, queryset: QuerySet[Even update_calculated_fields_for_events(qs=queryset) self.message_user(request, f"Updated {queryset.count()} events.") + # One Job per project, since a Job belongs to a single project and the changelist can span several. + run_tracking = make_post_processing_action( + TrackingTask, + TrackingActionForm, + build_jobs=build_tracking_jobs_for_events, + description="Run Occurrence tracking on the selected sessions (async)", + ) + list_filter = ("deployment", "project", "start") - actions = [update_calculated_fields] + actions = [update_calculated_fields, run_tracking] @admin.register(SourceImage) @@ -861,11 +872,20 @@ def populate_collection_async(self, request: HttpRequest, queryset: QuerySet[Sou f"Post-processing: {task_cls.name} on Capture Set {collection.pk}" ), ) + run_tracking = make_post_processing_action( + TrackingTask, + TrackingActionForm, + scope_resolver=lambda collection: {"source_image_collection_id": collection.pk}, + name_resolver=lambda task_cls, collection: ( + f"Post-processing: {task_cls.name} on Capture Set {collection.pk}" + ), + ) actions = [ populate_collection, populate_collection_async, run_small_size_filter, run_class_masking, + run_tracking, ] # Hide images many-to-many field from form. This would list all source images in the database. diff --git a/ami/ml/post_processing/__init__.py b/ami/ml/post_processing/__init__.py index c94be9ae9..56e3e2fec 100644 --- a/ami/ml/post_processing/__init__.py +++ b/ami/ml/post_processing/__init__.py @@ -1,2 +1,3 @@ from . import class_masking # noqa: F401 from . import small_size_filter # noqa: F401 +from . import tracking_task # noqa: F401 diff --git a/ami/ml/post_processing/admin/tracking_actions.py b/ami/ml/post_processing/admin/tracking_actions.py new file mode 100644 index 000000000..61966b17b --- /dev/null +++ b/ami/ml/post_processing/admin/tracking_actions.py @@ -0,0 +1,65 @@ +"""Job builder for the Sessions (Events) entry point to Occurrence tracking. + +Tracking runs per session, and the Events changelist lets an operator pick sessions from several +projects at once. A Job belongs to exactly one project, so the selection is split by project: +one Job per project, carrying that project's session ids. +""" +from __future__ import annotations + +import collections +from typing import Any + +import pydantic +from django.db import transaction + +from ami.jobs.models import Job +from ami.ml.post_processing.admin.actions import ConfigValidationErrors +from ami.ml.post_processing.base import BasePostProcessingTask + + +def build_tracking_jobs_for_events( + *, + config: dict[str, Any], + queryset, + task_cls: type[BasePostProcessingTask], + form_field_names: set[str], + **_unused: Any, +) -> list[int]: + """Enqueue one tracking Job per project across the selected sessions.""" + events_by_project: dict[Any, list[int]] = collections.defaultdict(list) + orphans: list[int] = [] + for event in queryset: + if event.project_id is None: + orphans.append(event.pk) + else: + events_by_project[event.project_id].append(event.pk) + + errors: list[tuple[str | None, str]] = [] + if orphans: + errors.append((None, f"Session(s) {sorted(orphans)} have no project. Set their project first.")) + + validated: list[tuple[int, pydantic.BaseModel]] = [] + for project_id, event_ids in events_by_project.items(): + try: + validated.append((project_id, task_cls.config_schema(**{**config, "event_ids": sorted(event_ids)}))) + except pydantic.ValidationError as exc: + for err in exc.errors(): + loc = err.get("loc") or () + field = str(loc[0]) if loc and str(loc[0]) in form_field_names else None + errors.append((field, err.get("msg", "Invalid value"))) + + if errors: + raise ConfigValidationErrors(errors) + + job_pks: list[int] = [] + with transaction.atomic(): + for project_id, model in validated: + job = Job.objects.create( + name=f"Post-processing: {task_cls.name} on {len(model.event_ids)} session(s)", + project_id=project_id, + job_type_key="post_processing", + params={"task": task_cls.key, "config": model.dict()}, + ) + job.enqueue() + job_pks.append(job.pk) + return job_pks diff --git a/ami/ml/post_processing/admin/tracking_form.py b/ami/ml/post_processing/admin/tracking_form.py new file mode 100644 index 000000000..27f257a5f --- /dev/null +++ b/ami/ml/post_processing/admin/tracking_form.py @@ -0,0 +1,15 @@ +from __future__ import annotations + +from ami.ml.post_processing.admin.forms import SchemaActionForm +from ami.ml.post_processing.tracking_task import TrackingConfig + + +class TrackingActionForm(SchemaActionForm): + """Knobs surfaced when an admin triggers Occurrence tracking. + + Every field is generated from ``TrackingConfig``; the scope (capture set or sessions) is supplied + by the admin entry point. + """ + + schema = TrackingConfig + exclude_fields = ("source_image_collection_id", "event_ids") diff --git a/ami/ml/post_processing/registry.py b/ami/ml/post_processing/registry.py index 308be18ae..3ad5923bc 100644 --- a/ami/ml/post_processing/registry.py +++ b/ami/ml/post_processing/registry.py @@ -1,10 +1,12 @@ # Registry of available post-processing tasks from ami.ml.post_processing.class_masking import ClassMaskingTask from ami.ml.post_processing.small_size_filter import SmallSizeFilterTask +from ami.ml.post_processing.tracking_task import TrackingTask POSTPROCESSING_TASKS = { SmallSizeFilterTask.key: SmallSizeFilterTask, ClassMaskingTask.key: ClassMaskingTask, + TrackingTask.key: TrackingTask, } diff --git a/ami/ml/post_processing/tests/test_tracking_admin.py b/ami/ml/post_processing/tests/test_tracking_admin.py new file mode 100644 index 000000000..1565bc13b --- /dev/null +++ b/ami/ml/post_processing/tests/test_tracking_admin.py @@ -0,0 +1,146 @@ +"""Admin-action tests for the Occurrence tracking trigger. + +The Sessions (events) changelist can span projects, so its selection becomes one Job per project; the +capture-set changelist uses the shared one-Job-per-row path. The form is generated from ``TrackingConfig``. +""" +from django.contrib import admin as django_admin +from django.test import Client, TestCase +from django.urls import reverse + +from ami.jobs.models import Job +from ami.main.models import Project, SourceImageCollection +from ami.ml.post_processing.tracking_task import TrackingConfig +from ami.tests.fixtures.main import create_captures, setup_test_project +from ami.users.models import User + +COST_FIELDS = ( + "cost_threshold", + "iou_weight", + "size_weight", + "distance_weight", + "min_iou", + "min_size_ratio", + "max_distance", + "max_capture_interval_seconds", +) + + +class _TrackingAdminCase(TestCase): + @classmethod + def setUpTestData(cls) -> None: + cls.superuser = User.objects.create_superuser(email=f"trackadmin+{cls.__name__}@example.com", password="x") + cls.project, cls.deployment = setup_test_project(reuse=False) + create_captures(deployment=cls.deployment, num_nights=1, images_per_night=2, interval_minutes=1) + cls.event = cls.project.events.first() + assert cls.event is not None + + def setUp(self) -> None: + self.client = Client() + self.client.force_login(self.superuser) + + @staticmethod + def _knobs(**overrides) -> dict: + """Form payload for a confirmed submit. Unchecked booleans are absent, as in HTML.""" + return { + "confirm": "yes", + "cost_threshold": "1.0", + "iou_weight": "1.0", + "size_weight": "1.0", + "distance_weight": "1.0", + "skip_if_human_identifications": "on", + "require_fresh_event": "on", + **overrides, + } + + +class TestEventAdminTrackingAction(_TrackingAdminCase): + def _post(self, data: dict, pks: list[int] | None = None): + selected = [str(pk) for pk in (pks or [self.event.pk])] + return self.client.post( + reverse("admin:main_event_changelist"), + data={"action": "run_tracking", django_admin.helpers.ACTION_CHECKBOX_NAME: selected, **data}, + ) + + def test_renders_every_tunable_with_its_help_text_and_creates_no_job(self): + response = self._post({}) + self.assertEqual(response.status_code, 200) + self.assertContains(response, "Run Occurrence tracking") + for name in COST_FIELDS: + self.assertContains(response, f'name="{name}"') + self.assertContains(response, TrackingConfig.__fields__[name].field_info.title) + self.assertContains(response, "about 20 seconds apart") + self.assertNotContains(response, 'name="event_ids"') + self.assertEqual(Job.objects.filter(job_type_key="post_processing").count(), 0) + + def test_creates_one_job_per_project_carrying_its_own_events_and_non_default_values(self): + other_project, other_deployment = setup_test_project(reuse=False) + create_captures(deployment=other_deployment, num_nights=1, images_per_night=2, interval_minutes=1) + other_event = other_project.events.first() + assert other_event is not None + + response = self._post( + self._knobs( + cost_threshold="0.35", min_iou="0.25", max_capture_interval_seconds="45", require_fresh_event="" + ), + pks=[self.event.pk, other_event.pk], + ) + self.assertEqual(response.status_code, 302) + + jobs = {job.project_id: job for job in Job.objects.filter(job_type_key="post_processing")} + self.assertEqual(set(jobs), {self.project.pk, other_project.pk}) + for job in jobs.values(): + config = job.params["config"] + self.assertEqual(job.params["task"], "tracking") + self.assertEqual(config["cost_threshold"], 0.35) + self.assertEqual(config["min_iou"], 0.25) + self.assertEqual(config["max_capture_interval_seconds"], 45) + self.assertIsNone(config["max_distance"]) + self.assertTrue(config["skip_if_human_identifications"]) + self.assertFalse(config["require_fresh_event"]) + self.assertEqual(jobs[self.project.pk].params["config"]["event_ids"], [self.event.pk]) + self.assertEqual(jobs[other_project.pk].params["config"]["event_ids"], [other_event.pk]) + + def test_an_out_of_range_value_is_shown_on_the_form_and_creates_no_job(self): + response = self._post(self._knobs(min_iou="1.5")) + self.assertEqual(response.status_code, 200) + self.assertTrue(response.context["form"].errors.get("min_iou")) + self.assertEqual(Job.objects.filter(job_type_key="post_processing").count(), 0) + + def test_a_zero_interval_is_refused_by_the_schema_and_shown_on_the_field(self): + response = self._post(self._knobs(max_capture_interval_seconds="0")) + self.assertEqual(response.status_code, 200) + self.assertTrue(response.context["form"].errors.get("max_capture_interval_seconds")) + self.assertEqual(Job.objects.filter(job_type_key="post_processing").count(), 0) + + def test_sessions_without_a_project_are_refused(self): + orphan = self.event + type(orphan).objects.filter(pk=orphan.pk).update(project=None) + response = self._post(self._knobs()) + self.assertEqual(response.status_code, 200) + self.assertContains(response, "have no project") + self.assertEqual(Job.objects.filter(job_type_key="post_processing").count(), 0) + self.assertTrue(Project.objects.filter(pk=self.project.pk).exists()) + + +class TestCollectionAdminTrackingAction(_TrackingAdminCase): + def test_creates_a_job_scoped_to_the_capture_set(self): + collection = SourceImageCollection.objects.create(project=self.project, name="Tracking admin test set") + collection.images.set(self.event.captures.all()) + + response = self.client.post( + reverse("admin:main_sourceimagecollection_changelist"), + data={ + "action": "run_tracking", + django_admin.helpers.ACTION_CHECKBOX_NAME: [str(collection.pk)], + **self._knobs(distance_weight="2.5", max_distance="0.1"), + }, + ) + self.assertEqual(response.status_code, 302) + + job = Job.objects.get(job_type_key="post_processing") + self.assertEqual(job.project_id, self.project.pk) + config = job.params["config"] + self.assertEqual(config["source_image_collection_id"], collection.pk) + self.assertEqual(config["event_ids"], []) + self.assertEqual(config["distance_weight"], 2.5) + self.assertEqual(config["max_distance"], 0.1) diff --git a/ami/ml/post_processing/tests/test_tracking_task.py b/ami/ml/post_processing/tests/test_tracking_task.py new file mode 100644 index 000000000..5e59ce40f --- /dev/null +++ b/ami/ml/post_processing/tests/test_tracking_task.py @@ -0,0 +1,341 @@ +import datetime +import logging +import typing + +import pydantic +from django.test import SimpleTestCase, TestCase + +from ami.jobs.models import Job +from ami.main.models import Classification, Detection, Event, Identification, Occurrence, Taxon +from ami.ml.models import Algorithm +from ami.ml.post_processing.registry import get_postprocessing_task +from ami.ml.post_processing.tracking_task import ( + TrackingConfig, + TrackingTask, + assign_occurrences_from_detection_chains, + captures_too_far_apart, + pair_cost, + select_links, +) +from ami.tests.fixtures.main import create_taxa, setup_test_project +from ami.tests.fixtures.tracking import add_detection, create_session +from ami.users.tests.factories import UserFactory + +logger = logging.getLogger(__name__) + +BOX = [100, 100, 200, 200] +DIAG = 1000 * 2**0.5 + + +def _config(**kwargs) -> TrackingConfig: + return TrackingConfig(event_ids=[1], **kwargs) + + +class TestTrackingConfig(SimpleTestCase): + def test_defaults_are_the_plain_sum_baseline_with_every_limit_off(self): + config = _config() + self.assertEqual( + (config.cost_threshold, config.iou_weight, config.size_weight, config.distance_weight), + (1.0, 1.0, 1.0, 1.0), + ) + for name in ("min_iou", "min_size_ratio", "max_distance", "max_capture_interval_seconds"): + self.assertIsNone(getattr(config, name), name) + self.assertTrue(config.skip_if_human_identifications) + self.assertTrue(config.require_fresh_event) + + def test_exactly_one_scope(self): + with self.assertRaises(pydantic.ValidationError): + TrackingConfig() + with self.assertRaises(pydantic.ValidationError): + TrackingConfig(source_image_collection_id=1, event_ids=[1]) + TrackingConfig(source_image_collection_id=1) + + def test_values_outside_their_range_are_rejected(self): + for bad in ( + {"min_iou": 1.5}, + {"min_size_ratio": -0.1}, + {"max_distance": -1}, + {"max_capture_interval_seconds": 0}, + {"iou_weight": -1}, + {"cost_threshold": -1}, + {"unknown_option": 1}, + ): + with self.subTest(bad), self.assertRaises(pydantic.ValidationError): + _config(**bad) + + def test_every_tunable_has_a_title_and_help_text(self): + for name, field in TrackingConfig.__fields__.items(): + if name in ("source_image_collection_id", "event_ids"): + continue + self.assertTrue(field.field_info.title, name) + self.assertTrue(field.field_info.description, name) + + def test_task_is_registered(self): + self.assertIs(get_postprocessing_task("tracking"), TrackingTask) + + +class TestPairCost(SimpleTestCase): + def test_default_cost_is_the_plain_sum_of_the_three_terms(self): + shifted = [150, 100, 250, 200] + # IoU with the +1 pixel convention: overlap 51x101, union 2*101*101 - 51*101. + expected = (1 - 51 * 101 / (2 * 101 * 101 - 51 * 101)) + 0.0 + 50 / DIAG + self.assertAlmostEqual(pair_cost(BOX, shifted, DIAG, _config()), expected) + + def test_identical_boxes_cost_nothing(self): + self.assertAlmostEqual(pair_cost(BOX, BOX, DIAG, _config()), 0.0) + + def test_weights_scale_their_term(self): + shifted = [150, 100, 250, 200] + base = pair_cost(BOX, shifted, DIAG, _config()) + no_overlap_term = pair_cost(BOX, shifted, DIAG, _config(iou_weight=0)) + self.assertAlmostEqual(base - no_overlap_term, 1 - 51 * 101 / (2 * 101 * 101 - 51 * 101)) + self.assertAlmostEqual(pair_cost(BOX, shifted, DIAG, _config(distance_weight=2)), base + 50 / DIAG) + smaller = [100, 100, 150, 150] + self.assertAlmostEqual( + pair_cost(BOX, smaller, DIAG, _config(size_weight=0)), + pair_cost(BOX, smaller, DIAG, _config(size_weight=3)) - 3 * (1 - 51 * 51 / (101 * 101)), + ) + + def test_each_enabled_limit_rejects_a_pair_that_fails_it(self): + shifted = [150, 100, 250, 200] # IoU about 0.33, same size, centres 50 px apart + smaller = [100, 100, 150, 150] # smaller box: size ratio about 0.25 + self.assertIsNotNone(pair_cost(BOX, shifted, DIAG, _config(min_iou=0.3))) + self.assertIsNone(pair_cost(BOX, shifted, DIAG, _config(min_iou=0.5))) + self.assertIsNotNone(pair_cost(BOX, smaller, DIAG, _config(min_size_ratio=0.2))) + self.assertIsNone(pair_cost(BOX, smaller, DIAG, _config(min_size_ratio=0.5))) + self.assertIsNotNone(pair_cost(BOX, shifted, DIAG, _config(max_distance=0.05))) + self.assertIsNone(pair_cost(BOX, shifted, DIAG, _config(max_distance=0.03))) + + def test_a_pair_failing_a_limit_is_not_a_candidate_even_at_a_huge_cutoff(self): + class Det: + def __init__(self, pk, bbox): + self.pk, self.bbox = pk, bbox + + links = select_links( + [Det(1, BOX)], [Det(2, [150, 100, 250, 200])], DIAG, _config(min_iou=0.9, cost_threshold=99) + ) + self.assertEqual(links, []) + + +class TestIntervalLimit(SimpleTestCase): + def test_interval_limit(self): + class Capture: + def __init__(self, timestamp): + self.timestamp = timestamp + + t0 = datetime.datetime(2026, 7, 1, 22, 0, 0) + near, far = Capture(t0 + datetime.timedelta(seconds=20)), Capture(t0 + datetime.timedelta(seconds=45)) + first = Capture(t0) + self.assertFalse(captures_too_far_apart(first, far, _config())) + config = _config(max_capture_interval_seconds=30) + self.assertFalse(captures_too_far_apart(first, near, config)) + self.assertTrue(captures_too_far_apart(first, far, config)) + self.assertTrue(captures_too_far_apart(first, Capture(None), config)) + + +class _TrackingCase(TestCase): + def setUp(self) -> None: + self.project, self.deployment = setup_test_project(reuse=False) + create_taxa(self.project) + self.taxa = list(Taxon.objects.filter(projects=self.project, rank="SPECIES").order_by("name")) + + def run_task(self, event: Event, job: Job | None = None, **config) -> TrackingTask: + task = TrackingTask(job=job, logger=logger, event_ids=[event.pk], **config) + task.run() + return task + + def occurrence_sizes(self, event: Event) -> list[int]: + return sorted(o.detections.count() for o in Occurrence.objects.filter(event=event)) + + +class TestTrackingRun(_TrackingCase): + def test_a_still_insect_is_folded_into_one_occurrence_and_stats_are_stored(self): + captures = create_session(self.deployment, [[BOX], [BOX], [BOX]], self.taxa[0]) + event = captures[0].event + self.assertEqual(Occurrence.objects.filter(event=event).count(), 3) + + self.run_task(event) + + self.assertEqual(self.occurrence_sizes(event), [3]) + occurrence = Occurrence.objects.get(event=event) + self.assertIsNotNone(occurrence.track_motion) + self.assertIsNotNone(occurrence.track_size_ratio) + self.assertEqual(Detection.objects.filter(source_image__event=event, next_detection__isnull=False).count(), 2) + event.refresh_from_db() + self.assertEqual(event.occurrences_count, 1) + + def test_two_insects_are_matched_one_to_one_by_lowest_cost(self): + left, right = [100, 100, 200, 200], [700, 700, 800, 800] + captures = create_session( + self.deployment, [[left, right], [[110, 100, 210, 200], [705, 700, 805, 800]]], self.taxa[0] + ) + self.run_task(captures[0].event) + self.assertEqual(self.occurrence_sizes(captures[0].event), [2, 2]) + + def test_an_unprocessed_capture_between_processed_ones_does_not_break_the_chain(self): + captures = create_session(self.deployment, [[BOX], None, [BOX]], self.taxa[0]) + event = captures[0].event + + self.run_task(event) + + self.assertEqual(self.occurrence_sizes(event), [2]) + + def test_a_processed_capture_with_no_insects_is_part_of_the_sequence(self): + """A null-bbox marker row makes a capture processed, so the empty capture separates its neighbours.""" + captures = create_session(self.deployment, [[BOX], [], [BOX]], self.taxa[0]) + event = captures[0].event + + self.run_task(event) + + self.assertEqual(self.occurrence_sizes(event), [1, 1]) + + def test_the_interval_limit_stops_links_across_a_long_gap(self): + captures = create_session(self.deployment, [[BOX], [BOX], [BOX]], self.taxa[0], interval_seconds=60) + event = captures[0].event + job = Job.objects.create(name="t", project=self.project, job_type_key="post_processing") + job.progress.add_stage("Post-processing", key="post_processing") + job.save() + + self.run_task(event, job=job, max_capture_interval_seconds=30) + + self.assertEqual(self.occurrence_sizes(event), [1, 1, 1]) + job.refresh_from_db() + params = {p.name: p.value for p in job.progress.get_stage("post_processing").params} + self.assertEqual(params["Capture pairs too far apart to compare"], 2) + + def test_a_chain_stops_at_a_session_boundary(self): + """A link stored between two sessions never merges their occurrences.""" + day = datetime.datetime(2026, 7, 1, 22, 0, 0) + first = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0], start=day) + second = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0], start=day + datetime.timedelta(hours=6)) + self.assertNotEqual(first[0].event_id, second[0].event_id) + last_of_first = first[1].detections.get() + first_of_second = second[0].detections.get() + last_of_first.next_detection = first_of_second + last_of_first.save(update_fields=["next_detection"]) + + for event in (first[0].event, second[0].event): + self.run_task(event) + + self.assertEqual(self.occurrence_sizes(first[0].event), [2]) + self.assertEqual(self.occurrence_sizes(second[0].event), [2]) + first_of_second.refresh_from_db() + self.assertNotEqual(first_of_second.occurrence_id, first[0].detections.get().occurrence_id) + + def test_a_session_with_a_single_processed_capture_is_skipped_with_a_reason(self): + captures = create_session(self.deployment, [[BOX], None], self.taxa[0]) + job = Job.objects.create(name="t", project=self.project, job_type_key="post_processing") + job.progress.add_stage("Post-processing", key="post_processing") + job.save() + + self.run_task(captures[0].event, job=job) + + job.refresh_from_db() + params = {p.name: p.value for p in job.progress.get_stage("post_processing").params} + self.assertEqual(params["Sessions skipped"], 1) + self.assertIn("fewer than two processed captures", params["Result"]) + + +class TestGuards(_TrackingCase): + def test_a_session_that_was_already_tracked_is_skipped_unless_the_guard_is_off(self): + captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) + event = captures[0].event + self.run_task(event) + add_detection(captures[1], [500, 500, 560, 560], self.taxa[0]) + self.assertEqual(self.occurrence_sizes(event), [1, 2]) + + self.run_task(event) + self.assertEqual(self.occurrence_sizes(event), [1, 2]) + self.assertEqual(Detection.objects.filter(source_image__event=event, next_detection__isnull=False).count(), 1) + + self.run_task(event, require_fresh_event=False) + self.assertEqual(self.occurrence_sizes(event), [1, 2]) + + def test_human_identifications_skip_the_session_unless_the_guard_is_off(self): + captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) + event = captures[0].event + occurrence = Occurrence.objects.filter(event=event).first() + Identification.objects.create(user=UserFactory(), taxon=self.taxa[1], occurrence=occurrence) + + self.run_task(event) + self.assertEqual(self.occurrence_sizes(event), [1, 1]) + + self.run_task(event, skip_if_human_identifications=False) + self.assertEqual(self.occurrence_sizes(event), [2]) + + +class TestMerging(_TrackingCase): + def test_identifications_move_onto_the_keeper_instead_of_being_deleted(self): + captures = create_session(self.deployment, [[BOX], [BOX], [BOX]], self.taxa[0]) + event = captures[0].event + user = UserFactory() + for occurrence in Occurrence.objects.filter(event=event): + Identification.objects.create(user=user, taxon=self.taxa[1], occurrence=occurrence) + + self.run_task(event, skip_if_human_identifications=False) + + keeper = Occurrence.objects.get(event=event) + self.assertEqual(keeper.detections.count(), 3) + self.assertEqual(Identification.objects.filter(user=user).count(), 3) + self.assertEqual(set(Identification.objects.values_list("occurrence_id", flat=True)), {keeper.pk}) + + def _two_detection_chain(self, first_score: float, second_score: float): + captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) + first, second = (c.detections.get() for c in captures) + for det, taxon, score in ((first, self.taxa[0], first_score), (second, self.taxa[1], second_score)): + Classification.objects.filter(detection=det).update(taxon=taxon, score=score) + det.occurrence.save() + first.next_detection = second + first.save(update_fields=["next_detection"]) + return captures, first, second + + def test_a_changed_determination_is_recorded_with_the_winning_prediction_as_applied_to(self): + captures, first, second = self._two_detection_chain(first_score=0.3, second_score=0.9) + tracking_algorithm = Algorithm.objects.create(name="Test tracking", key="test-tracking") + + assign_occurrences_from_detection_chains(captures, logger, record_as=tracking_algorithm) + + keeper = Occurrence.objects.get(pk=first.occurrence_id) + self.assertEqual(keeper.determination, self.taxa[1]) + record = Classification.objects.get(algorithm=tracking_algorithm) + self.assertEqual((record.detection_id, record.taxon, record.terminal), (second.pk, self.taxa[1], True)) + self.assertEqual((record.applied_to.detection_id, record.applied_to.taxon), (second.pk, self.taxa[1])) + self.assertNotEqual(record.applied_to.algorithm_id, tracking_algorithm.pk) + + assign_occurrences_from_detection_chains(captures, logger, record_as=tracking_algorithm) + self.assertEqual(Classification.objects.filter(algorithm=tracking_algorithm).count(), 1) + + def test_an_unchanged_determination_records_nothing(self): + captures, first, _ = self._two_detection_chain(first_score=0.9, second_score=0.3) + tracking_algorithm = Algorithm.objects.create(name="Test tracking", key="test-tracking") + + assign_occurrences_from_detection_chains(captures, logger, record_as=tracking_algorithm) + + self.assertEqual(Occurrence.objects.get(pk=first.occurrence_id).determination, self.taxa[0]) + self.assertFalse(Classification.objects.filter(algorithm=tracking_algorithm).exists()) + + def test_a_run_records_the_determination_under_the_task_algorithm(self): + captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) + for capture, taxon, score in zip(captures, self.taxa, (0.3, 0.9)): + Classification.objects.filter(detection__source_image=capture).update(taxon=taxon, score=score) + capture.detections.get().occurrence.save() + + task = self.run_task(captures[0].event) + + self.assertEqual(Classification.objects.filter(algorithm=task.algorithm).count(), 1) + + +class TestTrackingJobMetrics(_TrackingCase): + def test_result_line_says_what_was_tracked(self): + captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) + job = Job.objects.create(name="t", project=self.project, job_type_key="post_processing") + job.progress.add_stage("Post-processing", key="post_processing") + job.save() + + self.run_task(captures[0].event, job=job) + + job.refresh_from_db() + params: dict[str, typing.Any] = {p.name: p.value for p in job.progress.get_stage("post_processing").params} + self.assertEqual(params["Sessions tracked"], 1) + self.assertEqual(params["Detection links created"], 1) + self.assertEqual(params["Result"], "Tracked 1 session(s).") diff --git a/ami/ml/post_processing/tracking_task.py b/ami/ml/post_processing/tracking_task.py new file mode 100644 index 000000000..c321547e3 --- /dev/null +++ b/ami/ml/post_processing/tracking_task.py @@ -0,0 +1,595 @@ +import collections +import logging +import math +import typing +from collections.abc import Iterable, Iterator, Sequence + +import pydantic +from django.db import transaction +from django.db.models import Count, Exists, OuterRef +from django.utils import timezone + +from ami.main.models import ( + Classification, + Detection, + Event, + Identification, + Occurrence, + SourceImage, + SourceImageCollection, + update_calculated_fields_for_sessions_and_stations, + update_occurrence_determination, +) +from ami.main.models_future.track_stats import refresh_track_stats_for_ids +from ami.main.models_future.tracks import lock_sessions +from ami.ml.models import Algorithm +from ami.ml.post_processing.base import BasePostProcessingTask + +COST_NOTE = ( + "The default is a starting point that is still being tuned by experiment. " + "It suits captures taken about 20 seconds apart." +) + + +class TrackingConfig(pydantic.BaseModel): + """Scope and tunables for a tracking run. + + Scope: exactly one of ``source_image_collection_id`` or ``event_ids`` says + which sessions to track. A capture set is the bulk path; an explicit event + list is what the Events admin page sends. + + The matching cost between two detections in consecutive captures is + ``iou_weight * (1 - IoU) + size_weight * (1 - size ratio) + distance_weight * (distance / diagonal)``. + Two detections are linked only when the cost is below ``cost_threshold`` and + every enabled limit passes. The field titles and descriptions are the help text + shown on the admin form. + """ + + source_image_collection_id: int | None = None + event_ids: list[int] = [] + + cost_threshold: float = pydantic.Field( + 1.0, + title="Cost cutoff", + ge=0, + description=( + "Two detections in neighbouring captures are only linked when their matching cost is below this " + "value. The cost adds up how little the boxes overlap, how different their sizes are, and how far " + "apart their centres are. Lower values link fewer detections and make fewer mistakes. " + COST_NOTE + ), + ) + iou_weight: float = pydantic.Field( + 1.0, + title="Overlap weight", + ge=0, + description=("How strongly poor overlap between two boxes raises the cost. 0 ignores overlap. " + COST_NOTE), + ) + size_weight: float = pydantic.Field( + 1.0, + title="Size weight", + ge=0, + description=("How strongly a difference in box area raises the cost. 0 ignores size. " + COST_NOTE), + ) + distance_weight: float = pydantic.Field( + 1.0, + title="Distance weight", + ge=0, + description=( + "How strongly the distance between box centres, measured as a share of the image diagonal, " + "raises the cost. 0 ignores distance. " + COST_NOTE + ), + ) + + min_iou: float | None = pydantic.Field( + None, + title="Minimum overlap", + ge=0, + le=1, + description=( + "Never link two detections whose boxes overlap by less than this (0 to 1, where 1 is identical " + "boxes). Leave blank for no limit." + ), + ) + min_size_ratio: float | None = pydantic.Field( + None, + title="Minimum size ratio", + ge=0, + le=1, + description=( + "Never link two detections when the smaller box has less than this share of the area of the larger " + "one (0 to 1). Leave blank for no limit." + ), + ) + max_distance: float | None = pydantic.Field( + None, + title="Maximum distance", + ge=0, + description=( + "Never link two detections whose box centres are further apart than this share of the image " + "diagonal (for example 0.1 is ten percent). Leave blank for no limit." + ), + ) + max_capture_interval_seconds: float | None = pydantic.Field( + None, + title="Maximum time between captures", + gt=0, + description=( + "Never link detections in two neighbouring captures taken further apart than this many seconds. " + "Leave blank for no limit." + ), + ) + + skip_if_human_identifications: bool = pydantic.Field( + True, + title="Skip sessions with human identifications", + description="Leave a session alone when someone has already identified one of its occurrences.", + ) + require_fresh_event: bool = pydantic.Field( + True, + title="Only track sessions that have not been tracked", + description=( + "Skip a session when any of its occurrences already holds more than one detection. Turn this off " + "to track the session again." + ), + ) + + @pydantic.root_validator(skip_on_failure=True) + def _exactly_one_scope(cls, values: dict) -> dict: + scopes = [values.get("source_image_collection_id"), values.get("event_ids") or None] + if sum(s is not None for s in scopes) != 1: + raise ValueError("Provide exactly one of source_image_collection_id or event_ids") + return values + + class Config: + extra = "forbid" + + +def iou(bb1, bb2) -> float: + xA = max(bb1[0], bb2[0]) + yA = max(bb1[1], bb2[1]) + xB = min(bb1[2], bb2[2]) + yB = min(bb1[3], bb2[3]) + inter = max(0, xB - xA + 1) * max(0, yB - yA + 1) + area1 = (bb1[2] - bb1[0] + 1) * (bb1[3] - bb1[1] + 1) + area2 = (bb2[2] - bb2[0] + 1) * (bb2[3] - bb2[1] + 1) + union = area1 + area2 - inter + return inter / union if union > 0 else 0.0 + + +def box_ratio(bb1, bb2) -> float: + area1 = (bb1[2] - bb1[0] + 1) * (bb1[3] - bb1[1] + 1) + area2 = (bb2[2] - bb2[0] + 1) * (bb2[3] - bb2[1] + 1) + return min(area1, area2) / max(area1, area2) + + +def distance_ratio(bb1, bb2, img_diag: float) -> float: + cx1 = (bb1[0] + bb1[2]) / 2 + cy1 = (bb1[1] + bb1[3]) / 2 + cx2 = (bb2[0] + bb2[2]) / 2 + cy2 = (bb2[1] + bb2[3]) / 2 + dist = math.sqrt((cx2 - cx1) ** 2 + (cy2 - cy1) ** 2) + return dist / img_diag if img_diag > 0 else 1.0 + + +def image_diagonal(width: int, height: int) -> int: + return int(math.ceil(math.sqrt(width**2 + height**2))) + + +def pair_cost(bb1, bb2, diag: float, config: TrackingConfig) -> float | None: + """Matching cost between two detections; lower means more likely the same insect. + + Returns None when the pair fails an enabled limit, so it is never a candidate however + low its cost. With default weights the cost is the plain sum of the three terms. + """ + overlap = iou(bb1, bb2) + size_ratio = box_ratio(bb1, bb2) + distance = distance_ratio(bb1, bb2, diag) + if config.min_iou is not None and overlap < config.min_iou: + return None + if config.min_size_ratio is not None and size_ratio < config.min_size_ratio: + return None + if config.max_distance is not None and distance > config.max_distance: + return None + return ( + config.iou_weight * (1 - overlap) + config.size_weight * (1 - size_ratio) + config.distance_weight * distance + ) + + +def captures_too_far_apart(first: SourceImage, second: SourceImage, config: TrackingConfig) -> bool: + """Is the gap between two captures over the interval limit? A missing timestamp counts as too far.""" + if config.max_capture_interval_seconds is None: + return False + if first.timestamp is None or second.timestamp is None: + return True + return abs((second.timestamp - first.timestamp).total_seconds()) > config.max_capture_interval_seconds + + +def event_is_fresh(event: Event) -> tuple[bool, str]: + """Has this event's detections already been grouped into chains? + + The guard keeps tracking away from events that were already consolidated: merging + those again can delete an occurrence that carries identifications. An occurrence + spanning more than one detection is the signal for that. + + A detection with no occurrence at all is not that signal, since the chain walk creates an + occurrence for a chain that has none. Occurrences are found through their detections' + captures, not ``Occurrence.event``, so one that reaches into this session from another + session also counts. + """ + multi_detection_occurrences = ( + Occurrence.objects.filter(pk__in=Detection.objects.filter(source_image__event=event).values("occurrence_id")) + .annotate(_n=Count("detections")) + .filter(_n__gt=1) + .count() + ) + if multi_detection_occurrences: + return False, f"{multi_detection_occurrences} occurrence(s) already span >1 detection" + return True, "" + + +def processed_captures(event: Event) -> list[SourceImage]: + """Captures of a session that have at least one detection row, oldest first. + + A null-bbox marker row counts: it records that a capture was processed and found nothing. + Captures nobody processed are left out, so they cannot break the adjacency of their neighbours. + """ + return list( + SourceImage.objects.filter(event=event) + .filter(Exists(Detection.objects.filter(source_image=OuterRef("pk")))) + .order_by("timestamp", "pk") + ) + + +def record_tracking_determination(occurrence: Occurrence, algorithm: Algorithm) -> Classification | None: + """Leave a terminal classification by the tracking algorithm after a merge changed the determination. + + The row carries the winning prediction and points back at it through ``applied_to``, the way + class masking does, so the history shows what tracking decided. Nothing is written when the + winner already came from the tracking algorithm. + """ + winner = occurrence.best_prediction + if winner is None or winner.detection_id is None or winner.taxon_id is None: + return None + if winner.algorithm_id == algorithm.pk: + return None + return Classification.objects.create( + detection=winner.detection, + taxon=winner.taxon, + score=winner.score, + terminal=True, + algorithm=algorithm, + timestamp=timezone.now(), + applied_to=winner, + ) + + +def assign_occurrences_from_detection_chains( + source_images: Sequence[SourceImage], logger: logging.Logger, record_as: Algorithm | None = None +) -> dict[str, int]: + """Fold each chain of linked detections into one occurrence, keeping the first existing one. + + A chain never leaves the given captures, which belong to one session, so a link that crosses a + session boundary starts a new chain on each side. Identifications move onto the keeper before the + occurrences that held them are deleted, because deleting an occurrence deletes its identifications. + With ``record_as`` set, a merge that changes the keeper's determination leaves a classification by + that algorithm. Statistics are stored for every occurrence the chains settle on. + """ + image_ids = [image.pk for image in source_images] + detections = list( + Detection.objects.valid().filter(source_image_id__in=image_ids).select_related("source_image", "occurrence") + ) + by_id = {det.pk: det for det in detections} + # Walk each capture in time order, so a chain starts at its earliest detection. + position = {image_id: i for i, image_id in enumerate(image_ids)} + detections.sort(key=lambda d: (position[d.source_image_id], d.pk)) + has_previous = {det.next_detection_id for det in detections if det.next_detection_id in by_id} + + visited: set[int] = set() + settled: set[int] = set() + created = merged = identifications_moved = determinations_recorded = 0 + existing = Occurrence.objects.filter(detections__source_image_id__in=image_ids).distinct().count() + + for det in detections: + if det.pk in visited or det.pk in has_previous: + continue + chain: list[Detection] = [] + current: Detection | None = det + while current is not None and current.pk not in visited: + chain.append(current) + visited.add(current.pk) + current = by_id.get(current.next_detection_id) if current.next_detection_id else None + + old_occ_ids = {d.occurrence_id for d in chain if d.occurrence_id} + # Coherent chains need no change but still get their statistics stored below. + if len(old_occ_ids) == 1 and all(d.occurrence_id is not None for d in chain): + settled.update(old_occ_ids) + continue + + keeper: Occurrence | None = next((d.occurrence for d in chain if d.occurrence_id), None) + previous_determination_id = keeper.determination_id if keeper is not None else None + if keeper is None: + first_image = chain[0].source_image + keeper = Occurrence.objects.create( + event=first_image.event, deployment=first_image.deployment, project=first_image.project + ) + created += 1 + + for d in chain: + if d.occurrence_id != keeper.pk: + d.occurrence = keeper + d.save(update_fields=["occurrence"]) + + doomed = old_occ_ids - {keeper.pk} + if doomed: + identifications_moved += Identification.objects.filter(occurrence_id__in=doomed).update(occurrence=keeper) + Occurrence.objects.filter(pk__in=doomed).delete() + merged += len(doomed) + + # Only the determination is written, so a column this run did not change keeps its stored value. + if update_occurrence_determination(keeper, save=False): + keeper.save(update_determination=False, update_fields=["determination", "determination_score"]) + if record_as is not None and keeper.determination_id != previous_determination_id: + if record_tracking_determination(keeper, record_as) is not None: + determinations_recorded += 1 + settled.add(keeper.pk) + + # Stored once every determination is settled, since id_agreement is measured against it. + stats_stored = refresh_track_stats_for_ids(settled) + + new_count = Occurrence.objects.filter(detections__source_image_id__in=image_ids).distinct().count() + logger.info( + f"Created {created} occurrences and merged {merged} across {len(image_ids)} captures " + f"(occurrences before: {existing}, after: {new_count}). Moved {identifications_moved} identification(s), " + f"recorded {determinations_recorded} determination change(s), stored statistics for {stats_stored}." + ) + return { + "occurrences_before": existing, + "occurrences_after": new_count, + "occurrences_created": created, + "occurrences_merged": merged, + "identifications_moved": identifications_moved, + "determinations_recorded": determinations_recorded, + } + + +def select_links( + current_detections: Sequence[Detection], + next_detections: Sequence[Detection], + diag: float, + config: TrackingConfig, +) -> list[tuple[Detection, Detection, float]]: + """The links to make between two adjacent captures, lowest cost first; nothing is saved. + + A pair is a candidate when it passes every enabled limit and its cost is below the cutoff. + Candidates are taken lowest cost first, and each detection is linked at most once on either side. + """ + candidates: list[tuple[Detection, Detection, float]] = [] + for det in current_detections: + for nxt in next_detections: + cost = pair_cost(det.bbox, nxt.bbox, diag, config) + if cost is not None and cost < config.cost_threshold: + candidates.append((det, nxt, cost)) + + # Secondary keys keep tied costs deterministic across runs. + candidates.sort(key=lambda x: (x[2], x[0].pk, x[1].pk)) + + claimed_current: set[int] = set() + claimed_next: set[int] = set() + links: list[tuple[Detection, Detection, float]] = [] + for det, nxt, cost in candidates: + if det.pk in claimed_current or nxt.pk in claimed_next: + continue + claimed_current.add(det.pk) + claimed_next.add(nxt.pk) + links.append((det, nxt, cost)) + return links + + +def save_links(links: Iterable[tuple[Detection, Detection, float]], logger: logging.Logger) -> None: + """Store each link as ``next_detection``, first detaching any other detection that points at the target.""" + links = list(links) + if not links: + return + Detection.objects.filter(next_detection_id__in=[nxt.pk for _, nxt, _ in links]).update(next_detection=None) + for det, nxt, cost in links: + det.next_detection = nxt + logger.debug(f"Linked detection {det.pk} -> {nxt.pk} (cost {cost:.4f})") + Detection.objects.bulk_update([det for det, _, _ in links], ["next_detection"]) + + +def iter_transition_links( + source_images: Sequence[SourceImage], config: TrackingConfig, logger: logging.Logger +) -> Iterator[list[tuple[Detection, Detection, float]] | None]: + """Yield the proposed links for each pair of consecutive captures, in order, saving nothing. + + Yields None for a transition that is not compared (the earlier capture has no dimensions) and an + empty list for one over the interval limit. + """ + transitions = len(source_images) - 1 + for i in range(transitions): + cur, nxt = source_images[i], source_images[i + 1] + if captures_too_far_apart(cur, nxt, config): + yield [] + continue + if not cur.width or not cur.height: + logger.warning(f"Capture {cur.pk} has no dimensions; not comparing it with the next capture.") + yield None + continue + yield select_links( + list(cur.detections.valid()), + list(nxt.detections.valid()), + image_diagonal(cur.width, cur.height), + config, + ) + + +def assign_occurrences_by_tracking_images( + event: Event, + logger: logging.Logger, + config: TrackingConfig, + progress_cb: typing.Callable[[float], None] | None = None, + record_as: Algorithm | None = None, +) -> dict[str, int]: + """Link the detections of one session's processed captures and fold the chains into occurrences.""" + source_images = processed_captures(event) + if len(source_images) < 2: + logger.warning(f"Session {event.pk}: fewer than two processed captures ({len(source_images)}).") + return {} + + transitions = len(source_images) - 1 + links = skipped_for_dimensions = 0 + skipped_for_interval = sum( + captures_too_far_apart(source_images[i], source_images[i + 1], config) for i in range(transitions) + ) + # Per-session atomic boundary: a crash mid-session rolls back this session only. + with transaction.atomic(): + for i, proposed in enumerate(iter_transition_links(source_images, config, logger)): + if proposed is None: + skipped_for_dimensions += 1 + else: + save_links(proposed, logger) + links += len(proposed) + if progress_cb: + progress_cb((i + 1) / transitions) + counters = assign_occurrences_from_detection_chains(source_images, logger, record_as=record_as) + + counters["links_created"] = links + counters["transitions_too_far_apart"] = skipped_for_interval + counters["transitions_without_dimensions"] = skipped_for_dimensions + return counters + + +def nothing_tracked_summary(skip_reasons: collections.Counter[str]) -> str: + """The line a job shows when every session in scope was skipped, with the count per reason.""" + total = sum(skip_reasons.values()) + reasons = "; ".join(f"{count} because {reason}" for reason, count in skip_reasons.most_common()) + return f"Nothing was tracked: {total} session(s) skipped ({reasons})." + + +class TrackingTask(BasePostProcessingTask): + """Link detections across consecutive processed captures and fold each chain into one occurrence. + + Sets each detection's ``next_detection`` link from bounding-box overlap, size and distance, + then merges every chain into a single occurrence per session. + """ + + key = "tracking" + name = "Occurrence tracking" + config_schema = TrackingConfig + + config: TrackingConfig + + def _resolve_events(self) -> list[Event]: + """Return the sessions to track, from either scope in the config. + + When a job is attached, sessions outside ``job.project`` are dropped with a warning, which + guards against a trigger smuggling in sessions from a project the operator cannot see. + """ + if self.config.source_image_collection_id is not None: + collection = SourceImageCollection.objects.filter(pk=self.config.source_image_collection_id).first() + if collection is None: + raise ValueError(f"Capture set {self.config.source_image_collection_id} not found.") + qs = Event.objects.filter(captures__collections=collection).distinct() + requested: list[int] | None = None + else: + requested = list(self.config.event_ids) + qs = Event.objects.filter(pk__in=requested) + + if self.job and self.job.project_id: + cross_project = list(qs.exclude(project_id=self.job.project_id).values_list("pk", flat=True)) + if cross_project: + self.logger.warning( + f"Dropping {len(cross_project)} session(s) outside job project " + f"{self.job.project_id}: {cross_project}" + ) + qs = qs.filter(project_id=self.job.project_id) + + events = list(qs.order_by("pk").distinct()) + if requested is not None: + missing = set(requested) - {e.pk for e in events} + if missing: + self.logger.warning(f"Tracking requested {sorted(missing)} but those sessions were not found.") + return events + + def _skip_reason(self, event: Event) -> str | None: + """Why this session must not be tracked, or None. Called under the session lock.""" + if self.config.require_fresh_event: + fresh, detail = event_is_fresh(event) + if not fresh: + self.logger.info(f"Skipping session {event.pk}: already tracked or edited ({detail}).") + return "it was already tracked or edited" + if ( + self.config.skip_if_human_identifications + and Occurrence.objects.filter(event=event, identifications__isnull=False).exists() + ): + self.logger.info(f"Skipping session {event.pk}: has human identifications.") + return "it has human identifications" + return None + + def run(self) -> None: + self.logger.info(f"Tracking starting with config: {self.config.dict()}") + + events = self._resolve_events() + total = len(events) + self.logger.info(f"Tracking: {total} session(s) in scope") + + totals: collections.Counter[str] = collections.Counter() + tracked_event_ids: list[int] = [] + # Why each session was skipped, so a run that tracks nothing can say so. + skip_reasons: collections.Counter[str] = collections.Counter() + + for idx, event in enumerate(events, start=1): + self.logger.info(f"Tracking session {idx}/{total} (id={event.pk})") + # The checks and the writes share one lock on the session, so an edit made since the + # job started is seen and an edit made during the run waits. + with transaction.atomic(): + lock_sessions([event.pk]) + reason = self._skip_reason(event) + if reason is not None: + totals["skipped"] += 1 + skip_reasons[reason] += 1 + continue + + def _stage_progress(p: float, _idx=idx, _total=total) -> None: + self.update_progress(((_idx - 1) + p) / _total) + + counters = assign_occurrences_by_tracking_images( + event=event, + logger=self.logger, + config=self.config, + record_as=self.algorithm, + progress_cb=_stage_progress, + ) + if not counters: + totals["skipped"] += 1 + skip_reasons["it has fewer than two processed captures"] += 1 + continue + totals["tracked"] += 1 + tracked_event_ids.append(event.pk) + for key in ("links_created", "occurrences_merged", "transitions_too_far_apart"): + totals[key] += counters.get(key, 0) + + # Merging occurrences changes the session and station counts, which no save refreshes. + # This already runs in a background job, so the station refresh stays inline. + update_calculated_fields_for_sessions_and_stations(tracked_event_ids, stations_async=False) + + metrics: dict[str, typing.Any] = { + "Sessions tracked": totals["tracked"], + "Sessions skipped": totals["skipped"], + "Detection links created": totals["links_created"], + "Occurrences merged": totals["occurrences_merged"], + } + if self.config.max_capture_interval_seconds is not None: + metrics["Capture pairs too far apart to compare"] = totals["transitions_too_far_apart"] + # The job still succeeds when every session is skipped, so this line is written on every run: + # a retry keeps text params, and a stale line would contradict the counts. + if totals["tracked"]: + metrics["Result"] = f"Tracked {totals['tracked']} session(s)." + elif skip_reasons: + metrics["Result"] = nothing_tracked_summary(skip_reasons) + self.logger.warning(metrics["Result"]) + else: + metrics["Result"] = "Nothing was tracked: no sessions in scope." + self.report_stage_metrics(metrics) + self.update_progress(1.0) + self.logger.info(f"Tracking finished: {dict(totals)}") diff --git a/ami/tests/fixtures/tracking.py b/ami/tests/fixtures/tracking.py new file mode 100644 index 000000000..010cbc61d --- /dev/null +++ b/ami/tests/fixtures/tracking.py @@ -0,0 +1,61 @@ +"""Small synthetic capture sessions for tracking tests. + +A session is described by the bounding boxes in each capture. Every box becomes a detection with its own +occurrence and a classification, which is the state a pipeline leaves behind before tracking runs. +""" +import datetime + +from ami.main.models import Classification, Deployment, Detection, Occurrence, SourceImage, Taxon + +IMAGE_SIZE = (1000, 1000) + + +def create_session( + deployment: Deployment, + boxes_per_capture: list[list[list[int]] | None], + taxon: Taxon, + start: datetime.datetime | None = None, + interval_seconds: float = 20, + score: float = 0.9, +) -> list[SourceImage]: + """Create captures spaced ``interval_seconds`` apart and return them in time order. + + ``None`` leaves a capture unprocessed (no detection rows); an empty list gives it a null-bbox marker + row, as a processed capture with no insects has. The deployment groups the captures into sessions. + """ + start = start or datetime.datetime(2026, 7, 1, 22, 0, 0) + captures: list[SourceImage] = [] + for i, boxes in enumerate(boxes_per_capture): + timestamp = start + datetime.timedelta(seconds=i * interval_seconds) + capture = SourceImage.objects.create( + deployment=deployment, + project=deployment.project, + timestamp=timestamp, + path=f"tracking/{timestamp:%Y%m%d%H%M%S}_{i}.jpg", + width=IMAGE_SIZE[0], + height=IMAGE_SIZE[1], + ) + captures.append(capture) + deployment.save(update_calculated_fields=True, regroup_async=False) + for capture, boxes in zip(captures, boxes_per_capture): + capture.refresh_from_db() + if boxes is None: + continue + if not boxes: + Detection.objects.create(source_image=capture, bbox=None, timestamp=capture.timestamp) + for bbox in boxes: + add_detection(capture, bbox, taxon, score) + return captures + + +def add_detection(capture: SourceImage, bbox: list[int], taxon: Taxon, score: float = 0.9) -> Detection: + """Add one detection with its own occurrence and a terminal classification.""" + occurrence = Occurrence.objects.create(event=capture.event, deployment=capture.deployment, project=capture.project) + detection = Detection.objects.create( + source_image=capture, occurrence=occurrence, bbox=bbox, timestamp=capture.timestamp + ) + Classification.objects.create( + detection=detection, taxon=taxon, score=score, timestamp=capture.timestamp, terminal=True + ) + occurrence.save() + return detection From 7b389d38e90e1eef39920b5348bc423f18e0d5ed Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Sun, 4 Oct 2026 09:37:26 -0700 Subject: [PATCH 05/28] docs(tracking): say that re-tracking a session adds links but never undoes them The help text on the fresh-session guard suggested that turning it off tracks a session again from scratch. A run only adds links and merges; splitting an occurrence that an earlier run merged is occurrence editing, which comes later. Co-Authored-By: Claude Opus 5.5 --- ami/ml/post_processing/tracking_task.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/ami/ml/post_processing/tracking_task.py b/ami/ml/post_processing/tracking_task.py index c321547e3..13d8947c4 100644 --- a/ami/ml/post_processing/tracking_task.py +++ b/ami/ml/post_processing/tracking_task.py @@ -128,8 +128,8 @@ class TrackingConfig(pydantic.BaseModel): True, title="Only track sessions that have not been tracked", description=( - "Skip a session when any of its occurrences already holds more than one detection. Turn this off " - "to track the session again." + "Skip a session when any of its occurrences already holds more than one detection. Turned off, a " + "run can add links and merges to such a session, but it never undoes earlier ones." ), ) From c3c9b6a43b55657faa6007e635bd1571fb498de9 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Mon, 5 Oct 2026 16:55:12 -0700 Subject: [PATCH 06/28] refactor(exports): drop the dedicated tracks export The per-detection CSV export of occurrences duplicated what the detections export in #1395 is meant to provide, so it is removed together with its management command, its export format migration, and the queryset helper that only it used. The ami/exports app is back to its state before this branch. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01QT59KePky4u4nbCsTAggtc --- ami/exports/format_types.py | 21 --- .../migrations/0002_add_tracks_csv_format.py | 24 --- ami/exports/registry.py | 1 - ami/exports/tests.py | 150 ---------------- ami/exports/tracks.py | 161 ------------------ ami/main/management/commands/export_tracks.py | 46 ----- ami/main/models.py | 11 +- 7 files changed, 1 insertion(+), 413 deletions(-) delete mode 100644 ami/exports/migrations/0002_add_tracks_csv_format.py delete mode 100644 ami/exports/tracks.py delete mode 100644 ami/main/management/commands/export_tracks.py diff --git a/ami/exports/format_types.py b/ami/exports/format_types.py index 87bab8f91..a3f4c82d0 100644 --- a/ami/exports/format_types.py +++ b/ami/exports/format_types.py @@ -252,24 +252,3 @@ def export(self): self.update_job_progress(records_exported) self.update_export_stats(file_temp_path=temp_file.name) return temp_file.name # Return the file path - - -class TracksCSVExporter(BaseExporter): - """One row per detection of every occurrence in scope; see ami/exports/tracks.py for the columns.""" - - file_format = "csv" - - def get_queryset(self): - return Occurrence.objects.with_real_detections().filter(project=self.project) # type: ignore[union-attr] - - def export(self): - from ami.exports.tracks import write_tracks_csv - - temp_file = tempfile.NamedTemporaryFile(delete=False, suffix=".csv", mode="w", newline="", encoding="utf-8") - with open(temp_file.name, "w", newline="", encoding="utf-8") as csvfile: - rows = write_tracks_csv(self.queryset, csvfile, on_chunk=self.update_job_progress) - self.update_export_stats(file_temp_path=temp_file.name) - # A tracks file holds one record per detection, not per occurrence. - self.data_export.record_count = rows - self.data_export.save(update_fields=["record_count"]) - return temp_file.name diff --git a/ami/exports/migrations/0002_add_tracks_csv_format.py b/ami/exports/migrations/0002_add_tracks_csv_format.py deleted file mode 100644 index 0a48301da..000000000 --- a/ami/exports/migrations/0002_add_tracks_csv_format.py +++ /dev/null @@ -1,24 +0,0 @@ -# Generated by Django 4.2.10 on 2026-10-04 12:24 - -from django.db import migrations, models - - -class Migration(migrations.Migration): - dependencies = [ - ("exports", "0001_initial"), - ] - - operations = [ - migrations.AlterField( - model_name="dataexport", - name="format", - field=models.CharField( - choices=[ - ("occurrences_api_json", "occurrences_api_json"), - ("occurrences_simple_csv", "occurrences_simple_csv"), - ("tracks_csv", "tracks_csv"), - ], - max_length=255, - ), - ), - ] diff --git a/ami/exports/registry.py b/ami/exports/registry.py index 8574c77c6..29a4cc0e7 100644 --- a/ami/exports/registry.py +++ b/ami/exports/registry.py @@ -27,4 +27,3 @@ def get_supported_formats(cls): ExportRegistry.register("occurrences_api_json")(format_types.JSONExporter) ExportRegistry.register("occurrences_simple_csv")(format_types.CSVExporter) -ExportRegistry.register("tracks_csv")(format_types.TracksCSVExporter) diff --git a/ami/exports/tests.py b/ami/exports/tests.py index 48eaf75b1..866b1af61 100644 --- a/ami/exports/tests.py +++ b/ami/exports/tests.py @@ -1,5 +1,4 @@ import csv -import io import json import logging @@ -552,152 +551,3 @@ def test_csv_has_all_new_fields(self): ] for field in expected_fields: self.assertIn(field, headers, f"Missing CSV field: {field}") - - -class TracksExportTest(TestCase): - """The tracks CSV: one row per detection, a fixed column contract, and bounded queries.""" - - EXPECTED_HEADER = ( - "occurrence_id,detection_id,event_id,deployment_id,source_image_id,timestamp,detection_index," - "detection_count,bbox_x1,bbox_y1,bbox_x2,bbox_y2,image_width,image_height,detection_label," - "detection_score,occurrence_determination,occurrence_determination_score,next_detection_id" - ) - - def setUp(self): - self.project, self.deployment = setup_test_project(reuse=False) - self.user = self.project.owner - create_captures(deployment=self.deployment, num_nights=1, images_per_night=3, interval_minutes=1) - group_images_into_events(self.deployment) - create_taxa(self.project) - self.taxon = Taxon.objects.filter(projects=self.project).first() - self.algorithm, _ = Algorithm.objects.get_or_create( - name="test-classifier", defaults={"key": "test-classifier"} - ) - self.captures = list(self.project.captures.order_by("timestamp")) - self.captures[0].width, self.captures[0].height = 4096, 2160 - self.captures[0].save() - # Three occurrences of three detections each, created latest capture first so that - # the export cannot get capture order from detection pks. - self.occurrences = [self._make_occurrence(offset=i * 100) for i in range(3)] - - def _make_occurrence(self, offset: int) -> Occurrence: - occurrence = Occurrence.objects.create( - project=self.project, - deployment=self.deployment, - event=self.captures[0].event, - determination=self.taxon, - determination_score=0.9, - ) - for capture in reversed(self.captures): - detection = Detection.objects.create( - source_image=capture, - timestamp=capture.timestamp, - bbox=[offset, offset, offset + 10, offset + 20], - occurrence=occurrence, - ) - detection.classifications.create( - taxon=self.taxon, score=0.8, timestamp=capture.timestamp, algorithm=self.algorithm, terminal=True - ) - return occurrence - - def _rows(self, occurrences=None, **kwargs) -> list[dict[str, str]]: - from ami.exports.tracks import iter_track_rows - - return list(iter_track_rows(occurrences or Occurrence.objects.filter(project=self.project), **kwargs)) - - def test_format_export_header_is_the_contract(self): - data_export = DataExport.objects.create(user=self.user, project=self.project, format="tracks_csv") - file_path = data_export.run_export().replace("/media/", "") - with default_storage.open(file_path, "r") as f: - lines = f.read().splitlines() - default_storage.delete(file_path) - data_export.refresh_from_db() - - self.assertEqual(lines[0], self.EXPECTED_HEADER) - self.assertEqual(len(lines) - 1, 9) - self.assertEqual(data_export.record_count, 9, "A tracks export counts detection rows") - - def test_detection_index_follows_capture_time(self): - rows = [row for row in self._rows() if row["occurrence_id"] == str(self.occurrences[0].pk)] - by_capture = {int(row["source_image_id"]): row for row in rows} - for index, capture in enumerate(self.captures): - row = by_capture[capture.pk] - self.assertEqual(row["detection_index"], str(index)) - self.assertEqual(row["detection_count"], "3") - self.assertEqual(row["timestamp"], capture.timestamp.isoformat()) - first = by_capture[self.captures[0].pk] - self.assertEqual((first["image_width"], first["image_height"]), ("4096", "2160")) - self.assertEqual((first["bbox_x1"], first["bbox_y2"]), ("0", "20")) - self.assertEqual((first["detection_label"], first["detection_score"]), (self.taxon.name, "0.8")) - self.assertEqual(by_capture[self.captures[1].pk]["image_width"], "") - - def test_next_detection_id_carries_the_chain(self): - detections = list(self.occurrences[0].detections.order_by("source_image__timestamp")) - detections[0].next_detection = detections[1] - detections[0].save(update_fields=["next_detection"]) - - rows = {int(row["detection_id"]): row for row in self._rows()} - self.assertEqual(rows[detections[0].pk]["next_detection_id"], str(detections[1].pk)) - self.assertEqual(rows[detections[1].pk]["next_detection_id"], "") - - def test_query_count_is_one_pair_per_chunk(self): - from cachalot.api import cachalot_disabled - from django.db import connection - from django.test.utils import CaptureQueriesContext - - # Three occurrences in chunks of two: occurrences, detections, occurrences, detections. - with cachalot_disabled(), CaptureQueriesContext(connection) as ctx: - rows = self._rows(chunk_size=2) - self.assertEqual(len(rows), 9) - self.assertEqual(len(ctx.captured_queries), 4) - - # Doubling the detections per occurrence adds no queries. - for occurrence in self.occurrences: - for detection in list(occurrence.detections.all()): - detection.pk = None - detection.next_detection = None - detection.save() - with cachalot_disabled(), CaptureQueriesContext(connection) as ctx: - rows = self._rows(chunk_size=2) - self.assertEqual(len(rows), 18) - self.assertEqual(len(ctx.captured_queries), 4) - - def _run_command(self, **options) -> tuple[list[dict[str, str]], str]: - from django.core.management import call_command - - stdout, stderr = io.StringIO(), io.StringIO() - call_command("export_tracks", project=self.project.pk, stdout=stdout, stderr=stderr, **options) - lines = stdout.getvalue().splitlines() - self.assertEqual(lines[0], self.EXPECTED_HEADER) - return list(csv.DictReader(lines)), stderr.getvalue() - - def test_management_command_writes_the_same_csv(self): - rows, summary = self._run_command() - self.assertEqual(len(rows), 9) - self.assertIn("Wrote 9 detection rows", summary) - - def test_management_command_event_filter(self): - import datetime - - from ami.main.models import Event - - first_event = self.captures[0].event - other_event = Event.objects.create( - project=self.project, - deployment=self.deployment, - group_by="2030-01-01", - start=datetime.datetime(2030, 1, 1, 22, 0), - ) - moved = self.occurrences[2] - moved.event = other_event - moved.save(update_fields=["event"]) - - rows, _ = self._run_command(events=[first_event.pk]) - self.assertEqual( - {row["occurrence_id"] for row in rows}, {str(self.occurrences[0].pk), str(self.occurrences[1].pk)} - ) - self.assertEqual(len(rows), 6) - - rows, _ = self._run_command(events=[first_event.pk, other_event.pk]) - self.assertEqual({row["occurrence_id"] for row in rows}, {str(o.pk) for o in self.occurrences}) - self.assertEqual(len(rows), 9) diff --git a/ami/exports/tracks.py b/ami/exports/tracks.py deleted file mode 100644 index aba6f330a..000000000 --- a/ami/exports/tracks.py +++ /dev/null @@ -1,161 +0,0 @@ -""" -One row per detection of every occurrence in scope, for inspecting and benchmarking tracking. - -The export format and the ``export_tracks`` management command both write through -``iter_track_rows()``, so ``TRACKS_CSV_COLUMNS`` is the single definition of the columns. -""" - -import csv -import datetime -import typing -from collections.abc import Callable, Iterator - -from django.db import models -from django.db.models import OuterRef, Subquery - -from ami.main.models import BEST_MACHINE_PREDICTION_ORDER, Classification, Detection, Occurrence -from ami.main.models_future.tracks import CAPTURE_ORDER - -TRACKS_CSV_COLUMNS: typing.Final = ( - "occurrence_id", - "detection_id", - "event_id", - "deployment_id", - "source_image_id", - "timestamp", - "detection_index", - "detection_count", - "bbox_x1", - "bbox_y1", - "bbox_x2", - "bbox_y2", - "image_width", - "image_height", - "detection_label", - "detection_score", - "occurrence_determination", - "occurrence_determination_score", - "next_detection_id", -) - -DEFAULT_CHUNK_SIZE: typing.Final = 500 - - -def _cell(value) -> str: - """Render one value the way every tracks CSV writes it: blank for None, lowercase booleans.""" - if value is None: - return "" - if isinstance(value, bool): - return "true" if value else "false" - if isinstance(value, datetime.datetime): - return value.isoformat() - return str(value) - - -def _detections_for(occurrence_ids: list[int]) -> models.QuerySet: - """Real detections of the given occurrences, in capture order, with their own best label.""" - best_classification = Classification.objects.filter(detection=OuterRef("pk")).order_by( - *BEST_MACHINE_PREDICTION_ORDER - ) - return ( - Detection.objects.valid() # type: ignore[attr-defined] Custom queryset method - .filter(occurrence_id__in=occurrence_ids) - .annotate( - label=Subquery(best_classification.values("taxon__name")[:1]), - label_score=Subquery(best_classification.values("score")[:1]), - ) - .order_by("occurrence_id", *CAPTURE_ORDER) - .values( - "pk", - "occurrence_id", - "bbox", - "next_detection_id", - "source_image_id", - "source_image__timestamp", - "source_image__width", - "source_image__height", - "source_image__event_id", - "source_image__deployment_id", - "label", - "label_score", - ) - ) - - -def iter_track_rows( - occurrences: models.QuerySet[Occurrence], - chunk_size: int = DEFAULT_CHUNK_SIZE, - on_chunk: Callable[[int], None] | None = None, -) -> Iterator[dict[str, str]]: - """ - Yield one row per detection of each occurrence in ``occurrences``, keyed by TRACKS_CSV_COLUMNS. - - Occurrences are read in pk order, ``chunk_size`` at a time, with one detection query per - chunk, so the query count grows with the number of chunks and never with the rows. - ``on_chunk`` receives the running count of occurrences read, for progress reporting. - """ - scope = ( - occurrences.order_by("pk") - .values("pk", "determination__name", "determination_score") - .distinct() # A capture-set filter joins through detections and repeats occurrences. - ) - occurrences_read = 0 - last_pk = 0 - while True: - chunk = list(scope.filter(pk__gt=last_pk)[:chunk_size]) - if not chunk: - break - last_pk = chunk[-1]["pk"] - occurrences_read += len(chunk) - - detections_by_occurrence: dict[int, list[dict]] = {} - for detection in _detections_for([row["pk"] for row in chunk]): - detections_by_occurrence.setdefault(detection["occurrence_id"], []).append(detection) - - for occurrence in chunk: - detections = detections_by_occurrence.get(occurrence["pk"], []) - for detection_index, detection in enumerate(detections): - bbox = detection["bbox"] - x1, y1, x2, y2 = bbox if isinstance(bbox, list) and len(bbox) == 4 else (None,) * 4 - yield { - "occurrence_id": _cell(occurrence["pk"]), - "detection_id": _cell(detection["pk"]), - "event_id": _cell(detection["source_image__event_id"]), - "deployment_id": _cell(detection["source_image__deployment_id"]), - "source_image_id": _cell(detection["source_image_id"]), - "timestamp": _cell(detection["source_image__timestamp"]), - "detection_index": _cell(detection_index), - "detection_count": _cell(len(detections)), - "bbox_x1": _cell(x1), - "bbox_y1": _cell(y1), - "bbox_x2": _cell(x2), - "bbox_y2": _cell(y2), - "image_width": _cell(detection["source_image__width"]), - "image_height": _cell(detection["source_image__height"]), - "detection_label": _cell(detection["label"]), - "detection_score": _cell(detection["label_score"]), - "occurrence_determination": _cell(occurrence["determination__name"]), - "occurrence_determination_score": _cell(occurrence["determination_score"]), - "next_detection_id": _cell(detection["next_detection_id"]), - } - - if on_chunk: - on_chunk(occurrences_read) - if len(chunk) < chunk_size: - break - - -def write_tracks_csv( - occurrences: models.QuerySet[Occurrence], - stream: typing.TextIO, - chunk_size: int = DEFAULT_CHUNK_SIZE, - on_chunk: Callable[[int], None] | None = None, -) -> int: - """Write the header and every detection row to ``stream``; return the number of detection rows.""" - writer = csv.DictWriter(stream, fieldnames=TRACKS_CSV_COLUMNS) - writer.writeheader() - rows = 0 - for row in iter_track_rows(occurrences, chunk_size=chunk_size, on_chunk=on_chunk): - writer.writerow(row) - rows += 1 - return rows diff --git a/ami/main/management/commands/export_tracks.py b/ami/main/management/commands/export_tracks.py deleted file mode 100644 index 32feb13b3..000000000 --- a/ami/main/management/commands/export_tracks.py +++ /dev/null @@ -1,46 +0,0 @@ -""" -Write a project's occurrences as CSV, one row per detection, for inspecting and benchmarking tracking. - -The columns are the ``tracks_csv`` export format's (``ami/exports/tracks.py``), so a file from -this command and a file from the export API can be compared directly. Read-only. -""" - -from django.core.management.base import BaseCommand, CommandError - -from ami.exports.tracks import write_tracks_csv -from ami.main.models import Occurrence, Project - - -class Command(BaseCommand): - help = "Write a project's occurrences (one row per detection) as CSV to a file or stdout." - - def add_arguments(self, parser): - parser.add_argument("--project", type=int, required=True, help="Project ID to export.") - parser.add_argument( - "--event", - type=int, - action="append", - dest="events", - default=[], - help="Only occurrences of this session (event) ID. Repeat to export several.", - ) - parser.add_argument("--output", "-o", default="-", help="File path to write, or - for stdout (default).") - - def handle(self, *args, **options): - project_id: int = options["project"] - if not Project.objects.filter(pk=project_id).exists(): - raise CommandError(f"Project {project_id} does not exist") - - occurrences = Occurrence.objects.with_real_detections() # type: ignore[union-attr] - occurrences = occurrences.filter(project_id=project_id) - if options["events"]: - occurrences = occurrences.filter(event_id__in=options["events"]) - - output: str = options["output"] - if output == "-": - rows = write_tracks_csv(occurrences, self.stdout) - else: - with open(output, "w", newline="", encoding="utf-8") as stream: - rows = write_tracks_csv(occurrences, stream) - # Keep stdout pure CSV when the file goes there; the summary goes to stderr. - (self.stderr if output == "-" else self.stdout).write(f"Wrote {rows} detection rows.") diff --git a/ami/main/models.py b/ami/main/models.py index 391896000..59df1ec2a 100644 --- a/ami/main/models.py +++ b/ami/main/models.py @@ -3465,17 +3465,8 @@ def valid(self): - Occurrences with determination__isnull=True (no taxonomic identification, same field bug shape) """ - return self.with_real_detections().exclude(determination__isnull=True) - - def with_real_detections(self): - """ - Occurrences backed by at least one real bounding box, determined or not. - - Null-marker sentinels stay excluded exactly as in valid(), since they carry no box. - Used where undetermined occurrences must be reachable, such as the tracks export. - """ has_valid_detection = Exists(Detection.objects.valid().filter(occurrence_id=OuterRef("pk"))) - return self.filter(has_valid_detection) + return self.filter(has_valid_detection).exclude(determination__isnull=True) def with_detections_count(self): return self.annotate(detections_count=models.Count("detections", distinct=True)) From 334f7f4b2695f1c9831a358da3e2bd33bce0405f Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Mon, 5 Oct 2026 16:56:46 -0700 Subject: [PATCH 07/28] refactor(tracking): drop the stored occurrence statistics The four per-occurrence statistics (motion, size ratio, distinct taxa, label agreement) are removed from the occurrence table, together with the code that computed them, the backfill command and the migration that added the columns. Tracking and regrouping no longer refresh them. The numbers are meant to return as snapshots in the post-processing results of #1461, and as sortable columns only once the sorting interface exists. The only migration this branch adds is the one for the detection link. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01QT59KePky4u4nbCsTAggtc --- .../commands/backfill_track_stats.py | 61 ----- .../migrations/0097_occurrence_track_stats.py | 48 ---- ami/main/models.py | 24 -- ami/main/models_future/track_stats.py | 208 ------------------ ami/main/models_future/tracks.py | 2 - ami/main/tests.py | 107 --------- .../tests/test_tracking_task.py | 5 +- ami/ml/post_processing/tracking_task.py | 13 +- 8 files changed, 4 insertions(+), 464 deletions(-) delete mode 100644 ami/main/management/commands/backfill_track_stats.py delete mode 100644 ami/main/migrations/0097_occurrence_track_stats.py delete mode 100644 ami/main/models_future/track_stats.py diff --git a/ami/main/management/commands/backfill_track_stats.py b/ami/main/management/commands/backfill_track_stats.py deleted file mode 100644 index 4d4a8b889..000000000 --- a/ami/main/management/commands/backfill_track_stats.py +++ /dev/null @@ -1,61 +0,0 @@ -""" -Store track statistics on the occurrences of one project. - -Tracking and regrouping keep ``Occurrence.track_*`` current from the moment the fields -exist; this fills in the rows from before that, so occurrences can be sorted by motion, -size change and identification agreement. Occurrences with several detections are the -ones worth sorting by, so single-detection ones are skipped unless asked for. - -Safe to re-run: every pass recomputes from the current detections. ``updated_at`` is -left alone, so the default list order does not change. -""" - -from django.core.management.base import BaseCommand, CommandError -from django.db.models import Count - -from ami.main.models import Occurrence, Project -from ami.main.models_future.track_stats import REFRESH_BATCH_SIZE, refresh_track_stats_for_ids - - -class Command(BaseCommand): - help = "Store track statistics (motion, size ratio, distinct taxa, id agreement) on a project's occurrences." - - def add_arguments(self, parser): - parser.add_argument("--project", type=int, required=True, help="Project ID to backfill.") - parser.add_argument( - "--only-multi-detection", - action="store_true", - default=True, - help="Skip occurrences with a single detection (default).", - ) - parser.add_argument( - "--all-occurrences", - action="store_false", - dest="only_multi_detection", - help="Include single-detection occurrences as well.", - ) - - def handle(self, *args, **options): - project_id: int = options["project"] - only_multi_detection: bool = options["only_multi_detection"] - - try: - project = Project.objects.get(pk=project_id) - except Project.DoesNotExist as err: - raise CommandError(f"Project {project_id} does not exist") from err - - occurrences = Occurrence.objects.filter(project=project) - if only_multi_detection: - occurrences = occurrences.annotate(detection_count=Count("detections")).filter(detection_count__gt=1) - - # Materialize the ids first: the refresh writes to the same rows the filter reads. - ids = list(occurrences.order_by("pk").values_list("pk", flat=True)) - scope = "multi-detection occurrences" if only_multi_detection else "occurrences" - self.stdout.write(f"Project #{project.pk} ({project.name}): {len(ids)} {scope} to refresh.") - - refreshed = 0 - for start in range(0, len(ids), REFRESH_BATCH_SIZE): - refreshed += refresh_track_stats_for_ids(ids[start : start + REFRESH_BATCH_SIZE]) - self.stdout.write(f" {refreshed}/{len(ids)}") - - self.stdout.write(self.style.SUCCESS(f"Stored track statistics for {refreshed} {scope}.")) diff --git a/ami/main/migrations/0097_occurrence_track_stats.py b/ami/main/migrations/0097_occurrence_track_stats.py deleted file mode 100644 index b773a310f..000000000 --- a/ami/main/migrations/0097_occurrence_track_stats.py +++ /dev/null @@ -1,48 +0,0 @@ -# Generated by Django 4.2.10 on 2026-10-04 12:24 - -from django.db import migrations, models - - -class Migration(migrations.Migration): - dependencies = [ - ("main", "0096_detection_next_detection"), - ] - - operations = [ - migrations.AddField( - model_name="occurrence", - name="track_distinct_taxa", - field=models.IntegerField( - blank=True, - help_text="Distinct taxa among terminal classifications. See models_future/track_stats.py.", - null=True, - ), - ), - migrations.AddField( - model_name="occurrence", - name="track_id_agreement", - field=models.FloatField( - blank=True, - help_text="Share of terminal classifications naming the determination. See models_future/track_stats.py.", - null=True, - ), - ), - migrations.AddField( - model_name="occurrence", - name="track_motion", - field=models.FloatField( - blank=True, - help_text="Path length between captures as a fraction of the image diagonal. See track_stats.py.", - null=True, - ), - ), - migrations.AddField( - model_name="occurrence", - name="track_size_ratio", - field=models.FloatField( - blank=True, - help_text="Largest box area over the smallest. See models_future/track_stats.py.", - null=True, - ), - ), - ] diff --git a/ami/main/models.py b/ami/main/models.py index 59df1ec2a..eaf6ed15a 100644 --- a/ami/main/models.py +++ b/ami/main/models.py @@ -3751,30 +3751,6 @@ class Occurrence(BaseModel): deployment = models.ForeignKey(Deployment, on_delete=models.SET_NULL, null=True, related_name="occurrences") project = models.ForeignKey("Project", on_delete=models.SET_NULL, null=True, related_name="occurrences") - # Statistics stored so occurrences can be sorted by them. Written by - # ``models_future.track_stats.refresh_track_stats`` whenever tracking or a regroup changes - # which detections an occurrence holds; null until then, or when it has none. - track_motion = models.FloatField( - null=True, - blank=True, - help_text="Path length between captures as a fraction of the image diagonal. See track_stats.py.", - ) - track_size_ratio = models.FloatField( - null=True, - blank=True, - help_text="Largest box area over the smallest. See models_future/track_stats.py.", - ) - track_distinct_taxa = models.IntegerField( - null=True, - blank=True, - help_text="Distinct taxa among terminal classifications. See models_future/track_stats.py.", - ) - track_id_agreement = models.FloatField( - null=True, - blank=True, - help_text="Share of terminal classifications naming the determination. See models_future/track_stats.py.", - ) - detections: models.QuerySet[Detection] identifications: models.QuerySet[Identification] diff --git a/ami/main/models_future/track_stats.py b/ami/main/models_future/track_stats.py deleted file mode 100644 index 492ff5b98..000000000 --- a/ami/main/models_future/track_stats.py +++ /dev/null @@ -1,208 +0,0 @@ -"""Per-occurrence statistics of an insect's path: how far it moved, how much its box changed, -and how consistently it was labelled. - -Four of the numbers are stored on the occurrence row (Occurrence.track_*) so occurrences can -be sorted by them. refresh_track_stats recomputes them in SQL whenever tracking or a regroup -changes which detections an occurrence holds, and backfill_track_stats fills in older rows. - -Definitions: -- frames: detections in the occurrence (not stored). -- motion: path length between consecutive detection centres, ordered by - (timestamp, id), divided by the image diagonal. A stationary insect scores ~0. -- size_ratio: largest bbox area over smallest, areas floored at 1.0. -- distinct_taxa: distinct taxa among terminal classifications, across every - algorithm, so two classifiers that disagree count as two taxa. -- id_agreement: share of terminal classifications naming the determination; - null when there are none. - -The image diagonal is taken from the largest capture width and height seen; when the captures -carry no dimensions it falls back to the farthest bbox corner, and to 1.0 when that is also -unknown. Single-detection occurrences score motion 0.0 and size_ratio 1.0. -""" - -from __future__ import annotations - -import math -from collections.abc import Iterable - -from django.db import connection - -from ami.main.models import Occurrence - -# The stored subset, in the order the Occurrence fields are declared. -TRACK_STAT_FIELDS = ("track_motion", "track_size_ratio", "track_distinct_taxa", "track_id_agreement") - -# Ids per statement when refreshing many rows; keeps the ANY(%s) arrays and the -# CASE expression bulk_update builds to a size Postgres plans quickly. -REFRESH_BATCH_SIZE = 500 - -_MIN_AREA = 1.0 -_ROUND_TO = 4 - - -def frame_diagonal( - max_width: float | None, max_height: float | None, max_x2: float | None, max_y2: float | None -) -> float: - """Length motion is normalised by, in the same pixel units as the bboxes.""" - if max_width and max_height: - return math.hypot(max_width, max_height) - if max_x2 or max_y2: - return math.hypot(max_x2 or 0.0, max_y2 or 0.0) or 1.0 - return 1.0 - - -def size_ratio(min_area: float | None, max_area: float | None) -> float: - if not min_area or not max_area: - return 1.0 - return round(max_area / min_area, _ROUND_TO) - - -def id_agreement(agreeing: int, terminal_count: int) -> float | None: - if terminal_count == 0: - return None - return round(agreeing / terminal_count, _ROUND_TO) - - -# One row per occurrence: detection count, summed centre-to-centre distance in capture order, -# and the extremes the diagonal and size ratio are derived from. The CASE keeps a -# detection without a bbox out of the area extremes instead of flooring it to 1.0. -_GEOMETRY_SQL = """ -WITH frames AS ( - SELECT - d.occurrence_id, - d.id, - d.timestamp, - (d.bbox->>0)::float AS x1, - (d.bbox->>1)::float AS y1, - (d.bbox->>2)::float AS x2, - (d.bbox->>3)::float AS y2, - si.width, - si.height - FROM {detection} d - LEFT JOIN {source_image} si ON si.id = d.source_image_id - WHERE d.occurrence_id = ANY(%s) -), -steps AS ( - SELECT - occurrence_id, - (x1 + x2) / 2.0 AS cx, - (y1 + y2) / 2.0 AS cy, - LAG((x1 + x2) / 2.0) OVER track AS prev_cx, - LAG((y1 + y2) / 2.0) OVER track AS prev_cy, - CASE WHEN x1 IS NULL THEN NULL ELSE GREATEST(ABS((x2 - x1) * (y2 - y1)), {min_area}) END AS area, - x2, - y2, - width, - height - FROM frames - WINDOW track AS (PARTITION BY occurrence_id ORDER BY timestamp, id) -) -SELECT - occurrence_id, - COUNT(*) AS frames, - COALESCE(SUM(SQRT(POWER(cx - prev_cx, 2) + POWER(cy - prev_cy, 2))), 0.0) AS path_length, - MIN(area) AS min_area, - MAX(area) AS max_area, - MAX(width) AS max_width, - MAX(height) AS max_height, - MAX(x2) AS max_x2, - MAX(y2) AS max_y2 -FROM steps -GROUP BY occurrence_id -""" - -# One row per occurrence that has terminal classifications. -_CLASSIFICATION_SQL = """ -SELECT - d.occurrence_id, - COUNT(DISTINCT c.taxon_id) AS distinct_taxa, - COUNT(*) AS terminal_count, - COUNT(*) FILTER (WHERE c.taxon_id = o.determination_id) AS agreeing -FROM {classification} c -JOIN {detection} d ON d.id = c.detection_id -JOIN {occurrence} o ON o.id = d.occurrence_id -WHERE d.occurrence_id = ANY(%s) AND c.terminal -GROUP BY d.occurrence_id -""" - - -def track_stats_for_occurrences(occurrence_ids: list[int]) -> dict[int, dict]: - """Stats keyed by occurrence id, in two statements scoped to the ids given. - - Occurrences with no detections are absent from the result. The cost is proportional - to the detections behind the ids, not to the project, so keep a call to a page or a - refresh batch of ids. - """ - from ami.main.models import Classification, Detection, SourceImage - - if not occurrence_ids: - return {} - - geometry_sql = _GEOMETRY_SQL.format( - detection=Detection._meta.db_table, - source_image=SourceImage._meta.db_table, - min_area=_MIN_AREA, - ) - classification_sql = _CLASSIFICATION_SQL.format( - classification=Classification._meta.db_table, - detection=Detection._meta.db_table, - occurrence=Occurrence._meta.db_table, - ) - - stats: dict[int, dict] = {} - with connection.cursor() as cursor: - cursor.execute(geometry_sql, [list(occurrence_ids)]) - for pk, frames, path_length, min_area, max_area, max_width, max_height, max_x2, max_y2 in cursor.fetchall(): - diagonal = frame_diagonal(max_width, max_height, max_x2, max_y2) - stats[pk] = { - "frames": frames, - "motion": round(float(path_length) / diagonal, _ROUND_TO), - "size_ratio": size_ratio(min_area, max_area), - "distinct_taxa": 0, - "id_agreement": None, - } - - cursor.execute(classification_sql, [list(occurrence_ids)]) - for pk, distinct, terminal_count, agreeing in cursor.fetchall(): - if pk in stats: - stats[pk]["distinct_taxa"] = distinct - stats[pk]["id_agreement"] = id_agreement(agreeing, terminal_count) - - return stats - - -def _apply_stats(occurrence: Occurrence, stats: dict | None) -> None: - occurrence.track_motion = stats["motion"] if stats else None - occurrence.track_size_ratio = stats["size_ratio"] if stats else None - occurrence.track_distinct_taxa = stats["distinct_taxa"] if stats else None - occurrence.track_id_agreement = stats["id_agreement"] if stats else None - - -def refresh_track_stats(*occurrences: Occurrence) -> None: - """Recompute the stored stats for these occurrences from their current detections. - - Three queries however many occurrences are given: the two statements of - ``track_stats_for_occurrences`` and one ``bulk_update``. The instances are updated in - place as well as the rows. An occurrence with no detections is set back to null. - Call after the determination is settled, since ``id_agreement`` is measured against - it, and without going through ``Occurrence.save()``, which would recompute it. - """ - targets = [occurrence for occurrence in occurrences if occurrence.pk is not None] - if not targets: - return - stats = track_stats_for_occurrences([occurrence.pk for occurrence in targets]) - for occurrence in targets: - _apply_stats(occurrence, stats.get(occurrence.pk)) - Occurrence.objects.bulk_update(targets, TRACK_STAT_FIELDS) - - -def refresh_track_stats_for_ids(occurrence_ids: Iterable[int]) -> int: - """Refresh stored stats by id, in batches of ``REFRESH_BATCH_SIZE``; returns the count. - - Writes through ``bulk_update`` so ``updated_at`` is left alone: the list sorts by it - by default, and a backfill must not reorder every occurrence in a project. - """ - ids = list(occurrence_ids) - for start in range(0, len(ids), REFRESH_BATCH_SIZE): - refresh_track_stats(*(Occurrence(pk=pk) for pk in ids[start : start + REFRESH_BATCH_SIZE])) - return len(ids) diff --git a/ami/main/models_future/tracks.py b/ami/main/models_future/tracks.py index 7f3dda946..6dd5b7b77 100644 --- a/ami/main/models_future/tracks.py +++ b/ami/main/models_future/tracks.py @@ -14,7 +14,6 @@ from django.db.models import F from ami.main.models import Detection, Event, Identification, Occurrence, SourceImage, update_occurrence_determination -from ami.main.models_future.track_stats import refresh_track_stats # Order of detections within an occurrence: capture time, then capture, then detection. # The split and the tracks export both use it, so they agree on what "next" is. @@ -83,7 +82,6 @@ def split_at_session_boundaries(occurrence: Occurrence) -> list[Occurrence]: _copy_identifications(occurrence, pieces) for piece in [occurrence, *pieces]: update_occurrence_determination(piece, save=True) - refresh_track_stats(occurrence, *pieces) return pieces diff --git a/ami/main/tests.py b/ami/main/tests.py index 070cd70e4..083341db3 100644 --- a/ami/main/tests.py +++ b/ami/main/tests.py @@ -1,6 +1,5 @@ import copy import datetime -import io import logging import typing from io import BytesIO @@ -8611,7 +8610,6 @@ def test_an_occurrence_across_a_new_boundary_is_split_into_one_per_session(self) self.assertEqual(piece.determination_id, self.taxon.pk) self.assertEqual(first.occurrences_count, 1) self.assertEqual(second.occurrences_count, 1) - self.assertEqual(piece.track_motion, 0.0) def test_the_link_between_the_pieces_is_kept(self): _, _, detections, _ = self._split_one_occurrence() @@ -8664,108 +8662,3 @@ def test_identifications_are_copied_to_every_piece(self): ) self.assertEqual(occurrence.determination_id, self.other_taxon.pk) self.assertEqual(piece.determination_id, self.other_taxon.pk) - - -class TrackStatsTestCase(TestCase): - """The stored occurrence statistics follow their documented definitions and refresh in bulk.""" - - FRAME_WIDTH = 300 - FRAME_HEIGHT = 400 # a 300x400 image has a diagonal of exactly 500 - - def setUp(self) -> None: - self.project, self.deployment = setup_test_project(reuse=False) - create_taxa(project=self.project) - create_captures(deployment=self.deployment, num_nights=1, images_per_night=3, interval_minutes=1) - SourceImage.objects.filter(deployment=self.deployment).update(width=self.FRAME_WIDTH, height=self.FRAME_HEIGHT) - self.captures = list(SourceImage.objects.filter(deployment=self.deployment).order_by("timestamp")) - self.event = self.captures[0].event - taxa = list(Taxon.objects.filter(projects=self.project).order_by("pk")[:3]) - assert len(taxa) == 3, "Fixture must provide three taxa to disagree with" - self.taxon_a, self.taxon_b, self.taxon_c = taxa - - # Centres (5,5) -> (35,45) -> (65,85): two steps of 50 px over a 500 px diagonal. - # Areas 100, 100, 400. Terminal labels A, A, B plus a non-terminal C to ignore. - self.multi = self._make_occurrence( - [ - ([0, 0, 10, 10], [(self.taxon_a, 0.9, True)]), - ([30, 40, 40, 50], [(self.taxon_a, 0.85, True)]), - ([55, 75, 75, 95], [(self.taxon_b, 0.8, True), (self.taxon_c, 0.5, False)]), - ] - ) - self.single = self._make_occurrence([([10, 10, 20, 20], [(self.taxon_a, 0.7, True)])]) - - def _make_occurrence(self, frames: list[tuple[list[int], list[tuple[Taxon, float, bool]]]]) -> Occurrence: - occurrence = Occurrence.objects.create(event=self.event, deployment=self.deployment, project=self.project) - for capture, (bbox, labels) in zip(self.captures, frames): - detection = Detection.objects.create( - source_image=capture, timestamp=capture.timestamp, bbox=bbox, occurrence=occurrence - ) - for taxon, score, terminal in labels: - detection.classifications.create( - taxon=taxon, score=score, timestamp=capture.timestamp, terminal=terminal - ) - # Pin the determination the agreement is measured against, independent of how save() picks one. - Occurrence.objects.filter(pk=occurrence.pk).update(determination=self.taxon_a, determination_score=0.9) - return occurrence - - def _stored(self, occurrence: Occurrence) -> dict: - from ami.main.models_future.track_stats import TRACK_STAT_FIELDS - - return Occurrence.objects.filter(pk=occurrence.pk).values(*TRACK_STAT_FIELDS).get() - - def test_stats_follow_the_documented_definitions(self): - from ami.main.models_future.track_stats import refresh_track_stats - - refresh_track_stats(self.multi, self.single) - - stored = self._stored(self.multi) - self.assertAlmostEqual(stored["track_motion"], 100 / 500, places=4) - self.assertAlmostEqual(stored["track_size_ratio"], 4.0, places=4) - self.assertEqual(stored["track_distinct_taxa"], 2, "A non-terminal classification must not count as a taxon") - self.assertAlmostEqual(stored["track_id_agreement"], 2 / 3, places=4) - self.assertEqual( - self._stored(self.single), - {"track_motion": 0.0, "track_size_ratio": 1.0, "track_distinct_taxa": 1, "track_id_agreement": 1.0}, - ) - - def test_stats_are_null_until_stored_and_the_refresh_updates_the_instance(self): - from ami.main.models_future.track_stats import refresh_track_stats - - self.assertEqual(set(self._stored(self.multi).values()), {None}) - refresh_track_stats(self.multi) - self.assertEqual(self.multi.track_motion, 0.2) - - def test_refresh_is_three_queries_however_many_occurrences(self): - from cachalot.api import cachalot_disabled - - from ami.main.models_future.track_stats import refresh_track_stats - - with cachalot_disabled(), CaptureQueriesContext(connection) as ctx: - refresh_track_stats(self.multi, self.single) - self.assertEqual(len(ctx.captured_queries), 3) - - def test_an_occurrence_with_no_detections_is_set_back_to_null(self): - from ami.main.models_future.track_stats import refresh_track_stats - - refresh_track_stats(self.multi) - self.multi.detections.update(occurrence=None) - refresh_track_stats(self.multi) - self.assertEqual(set(self._stored(self.multi).values()), {None}) - - def test_the_backfill_command_stores_stats_for_multi_detection_occurrences(self): - from django.core.management import call_command - - out = io.StringIO() - call_command("backfill_track_stats", project=self.project.pk, stdout=out) - self.assertIn("1 multi-detection occurrences", out.getvalue()) - self.assertEqual(self._stored(self.multi)["track_motion"], 0.2) - self.assertIsNone(self._stored(self.single)["track_motion"], "Single detections are skipped by default") - - call_command("backfill_track_stats", project=self.project.pk, only_multi_detection=False, stdout=out) - self.assertEqual(self._stored(self.single)["track_motion"], 0.0) - - def test_the_backfill_command_refuses_an_unknown_project(self): - from django.core.management import CommandError, call_command - - with self.assertRaises(CommandError): - call_command("backfill_track_stats", project=0) diff --git a/ami/ml/post_processing/tests/test_tracking_task.py b/ami/ml/post_processing/tests/test_tracking_task.py index 5e59ce40f..2e92927d6 100644 --- a/ami/ml/post_processing/tests/test_tracking_task.py +++ b/ami/ml/post_processing/tests/test_tracking_task.py @@ -149,7 +149,7 @@ def occurrence_sizes(self, event: Event) -> list[int]: class TestTrackingRun(_TrackingCase): - def test_a_still_insect_is_folded_into_one_occurrence_and_stats_are_stored(self): + def test_a_still_insect_is_folded_into_one_occurrence(self): captures = create_session(self.deployment, [[BOX], [BOX], [BOX]], self.taxa[0]) event = captures[0].event self.assertEqual(Occurrence.objects.filter(event=event).count(), 3) @@ -157,9 +157,6 @@ def test_a_still_insect_is_folded_into_one_occurrence_and_stats_are_stored(self) self.run_task(event) self.assertEqual(self.occurrence_sizes(event), [3]) - occurrence = Occurrence.objects.get(event=event) - self.assertIsNotNone(occurrence.track_motion) - self.assertIsNotNone(occurrence.track_size_ratio) self.assertEqual(Detection.objects.filter(source_image__event=event, next_detection__isnull=False).count(), 2) event.refresh_from_db() self.assertEqual(event.occurrences_count, 1) diff --git a/ami/ml/post_processing/tracking_task.py b/ami/ml/post_processing/tracking_task.py index 13d8947c4..76e5fe15e 100644 --- a/ami/ml/post_processing/tracking_task.py +++ b/ami/ml/post_processing/tracking_task.py @@ -20,7 +20,6 @@ update_calculated_fields_for_sessions_and_stations, update_occurrence_determination, ) -from ami.main.models_future.track_stats import refresh_track_stats_for_ids from ami.main.models_future.tracks import lock_sessions from ami.ml.models import Algorithm from ami.ml.post_processing.base import BasePostProcessingTask @@ -272,7 +271,7 @@ def assign_occurrences_from_detection_chains( session boundary starts a new chain on each side. Identifications move onto the keeper before the occurrences that held them are deleted, because deleting an occurrence deletes its identifications. With ``record_as`` set, a merge that changes the keeper's determination leaves a classification by - that algorithm. Statistics are stored for every occurrence the chains settle on. + that algorithm. """ image_ids = [image.pk for image in source_images] detections = list( @@ -285,7 +284,6 @@ def assign_occurrences_from_detection_chains( has_previous = {det.next_detection_id for det in detections if det.next_detection_id in by_id} visited: set[int] = set() - settled: set[int] = set() created = merged = identifications_moved = determinations_recorded = 0 existing = Occurrence.objects.filter(detections__source_image_id__in=image_ids).distinct().count() @@ -300,9 +298,8 @@ def assign_occurrences_from_detection_chains( current = by_id.get(current.next_detection_id) if current.next_detection_id else None old_occ_ids = {d.occurrence_id for d in chain if d.occurrence_id} - # Coherent chains need no change but still get their statistics stored below. + # A chain already held by exactly one occurrence needs no change. if len(old_occ_ids) == 1 and all(d.occurrence_id is not None for d in chain): - settled.update(old_occ_ids) continue keeper: Occurrence | None = next((d.occurrence for d in chain if d.occurrence_id), None) @@ -331,16 +328,12 @@ def assign_occurrences_from_detection_chains( if record_as is not None and keeper.determination_id != previous_determination_id: if record_tracking_determination(keeper, record_as) is not None: determinations_recorded += 1 - settled.add(keeper.pk) - - # Stored once every determination is settled, since id_agreement is measured against it. - stats_stored = refresh_track_stats_for_ids(settled) new_count = Occurrence.objects.filter(detections__source_image_id__in=image_ids).distinct().count() logger.info( f"Created {created} occurrences and merged {merged} across {len(image_ids)} captures " f"(occurrences before: {existing}, after: {new_count}). Moved {identifications_moved} identification(s), " - f"recorded {determinations_recorded} determination change(s), stored statistics for {stats_stored}." + f"recorded {determinations_recorded} determination change(s)." ) return { "occurrences_before": existing, From 249c3c03dfa52b220b135764daca98c9bec2073f Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Mon, 5 Oct 2026 16:58:40 -0700 Subject: [PATCH 08/28] refactor(tracking): give tracking one home under ami/ml with pure code kept apart from Django Tracking code was spread over ami/ml/post_processing/tracking_task.py and ami/main/models_future/tracks.py. It now lives in the package ami/ml/post_processing/tracking/. The settings (config.py) and the matching rules (matching.py) import nothing from Django: matching takes plain (id, bbox) pairs and capture times and returns (id, next_id, cost) links, so the rules are tested with SimpleTestCase and no database. task.py holds the database orchestration and sessions.py holds session locking and the split of occurrences at session boundaries, which regrouping imports lazily. Behaviour, the task key and the job parameters are unchanged. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01QT59KePky4u4nbCsTAggtc --- ami/main/admin.py | 2 +- ami/main/models.py | 2 +- ami/ml/post_processing/__init__.py | 2 +- ami/ml/post_processing/admin/tracking_form.py | 2 +- ami/ml/post_processing/registry.py | 2 +- .../tests/test_tracking_admin.py | 2 +- .../tests/test_tracking_matching.py | 115 +++++++++ .../tests/test_tracking_task.py | 125 +--------- ami/ml/post_processing/tracking/__init__.py | 9 + ami/ml/post_processing/tracking/config.py | 124 +++++++++ ami/ml/post_processing/tracking/matching.py | 110 ++++++++ .../post_processing/tracking/sessions.py} | 3 +- .../{tracking_task.py => tracking/task.py} | 236 ++---------------- 13 files changed, 391 insertions(+), 343 deletions(-) create mode 100644 ami/ml/post_processing/tests/test_tracking_matching.py create mode 100644 ami/ml/post_processing/tracking/__init__.py create mode 100644 ami/ml/post_processing/tracking/config.py create mode 100644 ami/ml/post_processing/tracking/matching.py rename ami/{main/models_future/tracks.py => ml/post_processing/tracking/sessions.py} (96%) rename ami/ml/post_processing/{tracking_task.py => tracking/task.py} (66%) diff --git a/ami/main/admin.py b/ami/main/admin.py index b73fd1ea5..2c4aa2543 100644 --- a/ami/main/admin.py +++ b/ami/main/admin.py @@ -20,7 +20,7 @@ from ami.ml.post_processing.admin.tracking_form import TrackingActionForm from ami.ml.post_processing.class_masking import ClassMaskingTask from ami.ml.post_processing.small_size_filter import SmallSizeFilterTask -from ami.ml.post_processing.tracking_task import TrackingTask +from ami.ml.post_processing.tracking import TrackingTask from ami.ml.tasks import remove_duplicate_classifications from .models import ( diff --git a/ami/main/models.py b/ami/main/models.py index eaf6ed15a..b214d111a 100644 --- a/ami/main/models.py +++ b/ami/main/models.py @@ -1771,7 +1771,7 @@ def _split_occurrences_at_session_boundaries(deployment: Deployment, job: "Job | An occurrence is expected to belong to one session, so a regroup that draws a session boundary through it leaves one piece per session. Returns how many occurrences were split. """ - from ami.main.models_future.tracks import split_at_session_boundaries + from ami.ml.post_processing.tracking.sessions import split_at_session_boundaries spanning_ids = list( Detection.objects.valid() diff --git a/ami/ml/post_processing/__init__.py b/ami/ml/post_processing/__init__.py index 56e3e2fec..2f45726b1 100644 --- a/ami/ml/post_processing/__init__.py +++ b/ami/ml/post_processing/__init__.py @@ -1,3 +1,3 @@ from . import class_masking # noqa: F401 from . import small_size_filter # noqa: F401 -from . import tracking_task # noqa: F401 +from . import tracking # noqa: F401 diff --git a/ami/ml/post_processing/admin/tracking_form.py b/ami/ml/post_processing/admin/tracking_form.py index 27f257a5f..a1c39dc43 100644 --- a/ami/ml/post_processing/admin/tracking_form.py +++ b/ami/ml/post_processing/admin/tracking_form.py @@ -1,7 +1,7 @@ from __future__ import annotations from ami.ml.post_processing.admin.forms import SchemaActionForm -from ami.ml.post_processing.tracking_task import TrackingConfig +from ami.ml.post_processing.tracking import TrackingConfig class TrackingActionForm(SchemaActionForm): diff --git a/ami/ml/post_processing/registry.py b/ami/ml/post_processing/registry.py index 3ad5923bc..0b9156f6e 100644 --- a/ami/ml/post_processing/registry.py +++ b/ami/ml/post_processing/registry.py @@ -1,7 +1,7 @@ # Registry of available post-processing tasks from ami.ml.post_processing.class_masking import ClassMaskingTask from ami.ml.post_processing.small_size_filter import SmallSizeFilterTask -from ami.ml.post_processing.tracking_task import TrackingTask +from ami.ml.post_processing.tracking import TrackingTask POSTPROCESSING_TASKS = { SmallSizeFilterTask.key: SmallSizeFilterTask, diff --git a/ami/ml/post_processing/tests/test_tracking_admin.py b/ami/ml/post_processing/tests/test_tracking_admin.py index 1565bc13b..b1d9f69e9 100644 --- a/ami/ml/post_processing/tests/test_tracking_admin.py +++ b/ami/ml/post_processing/tests/test_tracking_admin.py @@ -9,7 +9,7 @@ from ami.jobs.models import Job from ami.main.models import Project, SourceImageCollection -from ami.ml.post_processing.tracking_task import TrackingConfig +from ami.ml.post_processing.tracking import TrackingConfig from ami.tests.fixtures.main import create_captures, setup_test_project from ami.users.models import User diff --git a/ami/ml/post_processing/tests/test_tracking_matching.py b/ami/ml/post_processing/tests/test_tracking_matching.py new file mode 100644 index 000000000..78f7b988a --- /dev/null +++ b/ami/ml/post_processing/tests/test_tracking_matching.py @@ -0,0 +1,115 @@ +import datetime + +import pydantic +from django.test import SimpleTestCase + +from ami.ml.post_processing.tracking.config import TrackingConfig +from ami.ml.post_processing.tracking.matching import captures_too_far_apart, pair_cost, select_links + +BOX = [100, 100, 200, 200] +DIAG = 1000 * 2**0.5 + + +def _config(**kwargs) -> TrackingConfig: + return TrackingConfig(event_ids=[1], **kwargs) + + +class TestTrackingConfig(SimpleTestCase): + def test_defaults_are_the_plain_sum_baseline_with_every_limit_off(self): + config = _config() + self.assertEqual( + (config.cost_threshold, config.iou_weight, config.size_weight, config.distance_weight), + (1.0, 1.0, 1.0, 1.0), + ) + for name in ("min_iou", "min_size_ratio", "max_distance", "max_capture_interval_seconds"): + self.assertIsNone(getattr(config, name), name) + self.assertTrue(config.skip_if_human_identifications) + self.assertTrue(config.require_fresh_event) + + def test_exactly_one_scope(self): + with self.assertRaises(pydantic.ValidationError): + TrackingConfig() + with self.assertRaises(pydantic.ValidationError): + TrackingConfig(source_image_collection_id=1, event_ids=[1]) + TrackingConfig(source_image_collection_id=1) + + def test_values_outside_their_range_are_rejected(self): + for bad in ( + {"min_iou": 1.5}, + {"min_size_ratio": -0.1}, + {"max_distance": -1}, + {"max_capture_interval_seconds": 0}, + {"iou_weight": -1}, + {"cost_threshold": -1}, + {"unknown_option": 1}, + ): + with self.subTest(bad), self.assertRaises(pydantic.ValidationError): + _config(**bad) + + def test_every_tunable_has_a_title_and_help_text(self): + for name, field in TrackingConfig.__fields__.items(): + if name in ("source_image_collection_id", "event_ids"): + continue + self.assertTrue(field.field_info.title, name) + self.assertTrue(field.field_info.description, name) + + +class TestPairCost(SimpleTestCase): + def test_default_cost_is_the_plain_sum_of_the_three_terms(self): + shifted = [150, 100, 250, 200] + # IoU with the +1 pixel convention: overlap 51x101, union 2*101*101 - 51*101. + expected = (1 - 51 * 101 / (2 * 101 * 101 - 51 * 101)) + 0.0 + 50 / DIAG + self.assertAlmostEqual(pair_cost(BOX, shifted, DIAG, _config()), expected) + + def test_identical_boxes_cost_nothing(self): + self.assertAlmostEqual(pair_cost(BOX, BOX, DIAG, _config()), 0.0) + + def test_weights_scale_their_term(self): + shifted = [150, 100, 250, 200] + base = pair_cost(BOX, shifted, DIAG, _config()) + no_overlap_term = pair_cost(BOX, shifted, DIAG, _config(iou_weight=0)) + self.assertAlmostEqual(base - no_overlap_term, 1 - 51 * 101 / (2 * 101 * 101 - 51 * 101)) + self.assertAlmostEqual(pair_cost(BOX, shifted, DIAG, _config(distance_weight=2)), base + 50 / DIAG) + smaller = [100, 100, 150, 150] + self.assertAlmostEqual( + pair_cost(BOX, smaller, DIAG, _config(size_weight=0)), + pair_cost(BOX, smaller, DIAG, _config(size_weight=3)) - 3 * (1 - 51 * 51 / (101 * 101)), + ) + + def test_each_enabled_limit_rejects_a_pair_that_fails_it(self): + shifted = [150, 100, 250, 200] # IoU about 0.33, same size, centres 50 px apart + smaller = [100, 100, 150, 150] # smaller box: size ratio about 0.25 + self.assertIsNotNone(pair_cost(BOX, shifted, DIAG, _config(min_iou=0.3))) + self.assertIsNone(pair_cost(BOX, shifted, DIAG, _config(min_iou=0.5))) + self.assertIsNotNone(pair_cost(BOX, smaller, DIAG, _config(min_size_ratio=0.2))) + self.assertIsNone(pair_cost(BOX, smaller, DIAG, _config(min_size_ratio=0.5))) + self.assertIsNotNone(pair_cost(BOX, shifted, DIAG, _config(max_distance=0.05))) + self.assertIsNone(pair_cost(BOX, shifted, DIAG, _config(max_distance=0.03))) + + def test_a_pair_failing_a_limit_is_not_a_candidate_even_at_a_huge_cutoff(self): + links = select_links([(1, BOX)], [(2, [150, 100, 250, 200])], DIAG, _config(min_iou=0.9, cost_threshold=99)) + self.assertEqual(links, []) + + def test_links_are_one_to_one_and_lowest_cost_wins(self): + left, right = [100, 100, 200, 200], [700, 700, 800, 800] + links = select_links( + [(1, left), (2, right)], [(3, [705, 700, 805, 800]), (4, [110, 100, 210, 200])], DIAG, _config() + ) + self.assertEqual([(a, b) for a, b, _ in links], [(2, 3), (1, 4)]) + self.assertLessEqual(links[0][2], links[1][2]) + + def test_a_detection_is_linked_at_most_once(self): + links = select_links([(1, BOX), (2, BOX)], [(3, BOX)], DIAG, _config()) + self.assertEqual([(a, b) for a, b, _ in links], [(1, 3)], "Ties break on the lower ids") + + +class TestIntervalLimit(SimpleTestCase): + def test_interval_limit(self): + t0 = datetime.datetime(2026, 7, 1, 22, 0, 0) + near, far = t0 + datetime.timedelta(seconds=20), t0 + datetime.timedelta(seconds=45) + first = t0 + self.assertFalse(captures_too_far_apart(first, far, _config())) + config = _config(max_capture_interval_seconds=30) + self.assertFalse(captures_too_far_apart(first, near, config)) + self.assertTrue(captures_too_far_apart(first, far, config)) + self.assertTrue(captures_too_far_apart(first, None, config)) diff --git a/ami/ml/post_processing/tests/test_tracking_task.py b/ami/ml/post_processing/tests/test_tracking_task.py index 2e92927d6..962c3872b 100644 --- a/ami/ml/post_processing/tests/test_tracking_task.py +++ b/ami/ml/post_processing/tests/test_tracking_task.py @@ -2,21 +2,14 @@ import logging import typing -import pydantic -from django.test import SimpleTestCase, TestCase +from django.test import TestCase from ami.jobs.models import Job from ami.main.models import Classification, Detection, Event, Identification, Occurrence, Taxon from ami.ml.models import Algorithm from ami.ml.post_processing.registry import get_postprocessing_task -from ami.ml.post_processing.tracking_task import ( - TrackingConfig, - TrackingTask, - assign_occurrences_from_detection_chains, - captures_too_far_apart, - pair_cost, - select_links, -) +from ami.ml.post_processing.tracking import TrackingTask +from ami.ml.post_processing.tracking.task import assign_occurrences_from_detection_chains from ami.tests.fixtures.main import create_taxa, setup_test_project from ami.tests.fixtures.tracking import add_detection, create_session from ami.users.tests.factories import UserFactory @@ -24,113 +17,6 @@ logger = logging.getLogger(__name__) BOX = [100, 100, 200, 200] -DIAG = 1000 * 2**0.5 - - -def _config(**kwargs) -> TrackingConfig: - return TrackingConfig(event_ids=[1], **kwargs) - - -class TestTrackingConfig(SimpleTestCase): - def test_defaults_are_the_plain_sum_baseline_with_every_limit_off(self): - config = _config() - self.assertEqual( - (config.cost_threshold, config.iou_weight, config.size_weight, config.distance_weight), - (1.0, 1.0, 1.0, 1.0), - ) - for name in ("min_iou", "min_size_ratio", "max_distance", "max_capture_interval_seconds"): - self.assertIsNone(getattr(config, name), name) - self.assertTrue(config.skip_if_human_identifications) - self.assertTrue(config.require_fresh_event) - - def test_exactly_one_scope(self): - with self.assertRaises(pydantic.ValidationError): - TrackingConfig() - with self.assertRaises(pydantic.ValidationError): - TrackingConfig(source_image_collection_id=1, event_ids=[1]) - TrackingConfig(source_image_collection_id=1) - - def test_values_outside_their_range_are_rejected(self): - for bad in ( - {"min_iou": 1.5}, - {"min_size_ratio": -0.1}, - {"max_distance": -1}, - {"max_capture_interval_seconds": 0}, - {"iou_weight": -1}, - {"cost_threshold": -1}, - {"unknown_option": 1}, - ): - with self.subTest(bad), self.assertRaises(pydantic.ValidationError): - _config(**bad) - - def test_every_tunable_has_a_title_and_help_text(self): - for name, field in TrackingConfig.__fields__.items(): - if name in ("source_image_collection_id", "event_ids"): - continue - self.assertTrue(field.field_info.title, name) - self.assertTrue(field.field_info.description, name) - - def test_task_is_registered(self): - self.assertIs(get_postprocessing_task("tracking"), TrackingTask) - - -class TestPairCost(SimpleTestCase): - def test_default_cost_is_the_plain_sum_of_the_three_terms(self): - shifted = [150, 100, 250, 200] - # IoU with the +1 pixel convention: overlap 51x101, union 2*101*101 - 51*101. - expected = (1 - 51 * 101 / (2 * 101 * 101 - 51 * 101)) + 0.0 + 50 / DIAG - self.assertAlmostEqual(pair_cost(BOX, shifted, DIAG, _config()), expected) - - def test_identical_boxes_cost_nothing(self): - self.assertAlmostEqual(pair_cost(BOX, BOX, DIAG, _config()), 0.0) - - def test_weights_scale_their_term(self): - shifted = [150, 100, 250, 200] - base = pair_cost(BOX, shifted, DIAG, _config()) - no_overlap_term = pair_cost(BOX, shifted, DIAG, _config(iou_weight=0)) - self.assertAlmostEqual(base - no_overlap_term, 1 - 51 * 101 / (2 * 101 * 101 - 51 * 101)) - self.assertAlmostEqual(pair_cost(BOX, shifted, DIAG, _config(distance_weight=2)), base + 50 / DIAG) - smaller = [100, 100, 150, 150] - self.assertAlmostEqual( - pair_cost(BOX, smaller, DIAG, _config(size_weight=0)), - pair_cost(BOX, smaller, DIAG, _config(size_weight=3)) - 3 * (1 - 51 * 51 / (101 * 101)), - ) - - def test_each_enabled_limit_rejects_a_pair_that_fails_it(self): - shifted = [150, 100, 250, 200] # IoU about 0.33, same size, centres 50 px apart - smaller = [100, 100, 150, 150] # smaller box: size ratio about 0.25 - self.assertIsNotNone(pair_cost(BOX, shifted, DIAG, _config(min_iou=0.3))) - self.assertIsNone(pair_cost(BOX, shifted, DIAG, _config(min_iou=0.5))) - self.assertIsNotNone(pair_cost(BOX, smaller, DIAG, _config(min_size_ratio=0.2))) - self.assertIsNone(pair_cost(BOX, smaller, DIAG, _config(min_size_ratio=0.5))) - self.assertIsNotNone(pair_cost(BOX, shifted, DIAG, _config(max_distance=0.05))) - self.assertIsNone(pair_cost(BOX, shifted, DIAG, _config(max_distance=0.03))) - - def test_a_pair_failing_a_limit_is_not_a_candidate_even_at_a_huge_cutoff(self): - class Det: - def __init__(self, pk, bbox): - self.pk, self.bbox = pk, bbox - - links = select_links( - [Det(1, BOX)], [Det(2, [150, 100, 250, 200])], DIAG, _config(min_iou=0.9, cost_threshold=99) - ) - self.assertEqual(links, []) - - -class TestIntervalLimit(SimpleTestCase): - def test_interval_limit(self): - class Capture: - def __init__(self, timestamp): - self.timestamp = timestamp - - t0 = datetime.datetime(2026, 7, 1, 22, 0, 0) - near, far = Capture(t0 + datetime.timedelta(seconds=20)), Capture(t0 + datetime.timedelta(seconds=45)) - first = Capture(t0) - self.assertFalse(captures_too_far_apart(first, far, _config())) - config = _config(max_capture_interval_seconds=30) - self.assertFalse(captures_too_far_apart(first, near, config)) - self.assertTrue(captures_too_far_apart(first, far, config)) - self.assertTrue(captures_too_far_apart(first, Capture(None), config)) class _TrackingCase(TestCase): @@ -148,6 +34,11 @@ def occurrence_sizes(self, event: Event) -> list[int]: return sorted(o.detections.count() for o in Occurrence.objects.filter(event=event)) +class TestRegistration(TestCase): + def test_task_is_registered(self): + self.assertIs(get_postprocessing_task("tracking"), TrackingTask) + + class TestTrackingRun(_TrackingCase): def test_a_still_insect_is_folded_into_one_occurrence(self): captures = create_session(self.deployment, [[BOX], [BOX], [BOX]], self.taxa[0]) diff --git a/ami/ml/post_processing/tracking/__init__.py b/ami/ml/post_processing/tracking/__init__.py new file mode 100644 index 000000000..9596380e7 --- /dev/null +++ b/ami/ml/post_processing/tracking/__init__.py @@ -0,0 +1,9 @@ +"""Occurrence tracking: linking the detections of one insect across consecutive captures. + +``config`` and ``matching`` are pure Python; ``task`` and ``sessions`` use the database. +""" + +from .config import TrackingConfig +from .task import TrackingTask + +__all__ = ["TrackingConfig", "TrackingTask"] diff --git a/ami/ml/post_processing/tracking/config.py b/ami/ml/post_processing/tracking/config.py new file mode 100644 index 000000000..3adb8b74f --- /dev/null +++ b/ami/ml/post_processing/tracking/config.py @@ -0,0 +1,124 @@ +"""Settings for a tracking run: which sessions to track and how detections are matched. + +This module uses pydantic only and has no Django imports, so it can be used and tested without a database. +""" + +import pydantic + +COST_NOTE = ( + "The default is a starting point that is still being tuned by experiment. " + "It suits captures taken about 20 seconds apart." +) + + +class TrackingConfig(pydantic.BaseModel): + """Scope and tunables for a tracking run. + + Scope: exactly one of ``source_image_collection_id`` or ``event_ids`` says + which sessions to track. A capture set is the bulk path; an explicit event + list is what the Events admin page sends. + + The matching cost between two detections in consecutive captures is + ``iou_weight * (1 - IoU) + size_weight * (1 - size ratio) + distance_weight * (distance / diagonal)``. + Two detections are linked only when the cost is below ``cost_threshold`` and + every enabled limit passes. The field titles and descriptions are the help text + shown on the admin form. + """ + + source_image_collection_id: int | None = None + event_ids: list[int] = [] + + cost_threshold: float = pydantic.Field( + 1.0, + title="Cost cutoff", + ge=0, + description=( + "Two detections in neighbouring captures are only linked when their matching cost is below this " + "value. The cost adds up how little the boxes overlap, how different their sizes are, and how far " + "apart their centres are. Lower values link fewer detections and make fewer mistakes. " + COST_NOTE + ), + ) + iou_weight: float = pydantic.Field( + 1.0, + title="Overlap weight", + ge=0, + description=("How strongly poor overlap between two boxes raises the cost. 0 ignores overlap. " + COST_NOTE), + ) + size_weight: float = pydantic.Field( + 1.0, + title="Size weight", + ge=0, + description=("How strongly a difference in box area raises the cost. 0 ignores size. " + COST_NOTE), + ) + distance_weight: float = pydantic.Field( + 1.0, + title="Distance weight", + ge=0, + description=( + "How strongly the distance between box centres, measured as a share of the image diagonal, " + "raises the cost. 0 ignores distance. " + COST_NOTE + ), + ) + + min_iou: float | None = pydantic.Field( + None, + title="Minimum overlap", + ge=0, + le=1, + description=( + "Never link two detections whose boxes overlap by less than this (0 to 1, where 1 is identical " + "boxes). Leave blank for no limit." + ), + ) + min_size_ratio: float | None = pydantic.Field( + None, + title="Minimum size ratio", + ge=0, + le=1, + description=( + "Never link two detections when the smaller box has less than this share of the area of the larger " + "one (0 to 1). Leave blank for no limit." + ), + ) + max_distance: float | None = pydantic.Field( + None, + title="Maximum distance", + ge=0, + description=( + "Never link two detections whose box centres are further apart than this share of the image " + "diagonal (for example 0.1 is ten percent). Leave blank for no limit." + ), + ) + max_capture_interval_seconds: float | None = pydantic.Field( + None, + title="Maximum time between captures", + gt=0, + description=( + "Never link detections in two neighbouring captures taken further apart than this many seconds. " + "Leave blank for no limit." + ), + ) + + skip_if_human_identifications: bool = pydantic.Field( + True, + title="Skip sessions with human identifications", + description="Leave a session alone when someone has already identified one of its occurrences.", + ) + require_fresh_event: bool = pydantic.Field( + True, + title="Only track sessions that have not been tracked", + description=( + "Skip a session when any of its occurrences already holds more than one detection. Turned off, a " + "run can add links and merges to such a session, but it never undoes earlier ones." + ), + ) + + @pydantic.root_validator(skip_on_failure=True) + def _exactly_one_scope(cls, values: dict) -> dict: + scopes = [values.get("source_image_collection_id"), values.get("event_ids") or None] + if sum(s is not None for s in scopes) != 1: + raise ValueError("Provide exactly one of source_image_collection_id or event_ids") + return values + + class Config: + extra = "forbid" diff --git a/ami/ml/post_processing/tracking/matching.py b/ami/ml/post_processing/tracking/matching.py new file mode 100644 index 000000000..0d5bdbd79 --- /dev/null +++ b/ami/ml/post_processing/tracking/matching.py @@ -0,0 +1,110 @@ +"""Pure matching rules that decide which detections in neighbouring captures are the same insect. + +Nothing here touches Django or the database. Detections are passed as ``(id, bbox)`` pairs and +links come back as ``(id, next_id, cost)`` tuples, so the rules can be tested without models. +""" + +import datetime +import math +from collections.abc import Sequence + +from .config import TrackingConfig + +BBox = Sequence[float] +Link = tuple[int, int, float] + + +def iou(bb1, bb2) -> float: + xA = max(bb1[0], bb2[0]) + yA = max(bb1[1], bb2[1]) + xB = min(bb1[2], bb2[2]) + yB = min(bb1[3], bb2[3]) + inter = max(0, xB - xA + 1) * max(0, yB - yA + 1) + area1 = (bb1[2] - bb1[0] + 1) * (bb1[3] - bb1[1] + 1) + area2 = (bb2[2] - bb2[0] + 1) * (bb2[3] - bb2[1] + 1) + union = area1 + area2 - inter + return inter / union if union > 0 else 0.0 + + +def box_ratio(bb1, bb2) -> float: + area1 = (bb1[2] - bb1[0] + 1) * (bb1[3] - bb1[1] + 1) + area2 = (bb2[2] - bb2[0] + 1) * (bb2[3] - bb2[1] + 1) + return min(area1, area2) / max(area1, area2) + + +def distance_ratio(bb1, bb2, img_diag: float) -> float: + cx1 = (bb1[0] + bb1[2]) / 2 + cy1 = (bb1[1] + bb1[3]) / 2 + cx2 = (bb2[0] + bb2[2]) / 2 + cy2 = (bb2[1] + bb2[3]) / 2 + dist = math.sqrt((cx2 - cx1) ** 2 + (cy2 - cy1) ** 2) + return dist / img_diag if img_diag > 0 else 1.0 + + +def image_diagonal(width: int, height: int) -> int: + return int(math.ceil(math.sqrt(width**2 + height**2))) + + +def pair_cost(bb1, bb2, diag: float, config: TrackingConfig) -> float | None: + """Matching cost between two detections; lower means more likely the same insect. + + Returns None when the pair fails an enabled limit, so it is never a candidate however + low its cost. With default weights the cost is the plain sum of the three terms. + """ + overlap = iou(bb1, bb2) + size_ratio = box_ratio(bb1, bb2) + distance = distance_ratio(bb1, bb2, diag) + if config.min_iou is not None and overlap < config.min_iou: + return None + if config.min_size_ratio is not None and size_ratio < config.min_size_ratio: + return None + if config.max_distance is not None and distance > config.max_distance: + return None + return ( + config.iou_weight * (1 - overlap) + config.size_weight * (1 - size_ratio) + config.distance_weight * distance + ) + + +def captures_too_far_apart( + first: datetime.datetime | None, second: datetime.datetime | None, config: TrackingConfig +) -> bool: + """Is the gap between two capture times over the interval limit? A missing timestamp counts as too far.""" + if config.max_capture_interval_seconds is None: + return False + if first is None or second is None: + return True + return abs((second - first).total_seconds()) > config.max_capture_interval_seconds + + +def select_links( + current_detections: Sequence[tuple[int, BBox]], + next_detections: Sequence[tuple[int, BBox]], + diag: float, + config: TrackingConfig, +) -> list[Link]: + """The links to make between two adjacent captures, lowest cost first. + + Each detection is an ``(id, bbox)`` pair. A pair is a candidate when it passes every enabled limit + and its cost is below the cutoff. Candidates are taken lowest cost first, and each detection is + linked at most once on either side. + """ + candidates: list[Link] = [] + for det_id, det_box in current_detections: + for next_id, next_box in next_detections: + cost = pair_cost(det_box, next_box, diag, config) + if cost is not None and cost < config.cost_threshold: + candidates.append((det_id, next_id, cost)) + + # Secondary keys keep tied costs deterministic across runs. + candidates.sort(key=lambda x: (x[2], x[0], x[1])) + + claimed_current: set[int] = set() + claimed_next: set[int] = set() + links: list[Link] = [] + for det_id, next_id, cost in candidates: + if det_id in claimed_current or next_id in claimed_next: + continue + claimed_current.add(det_id) + claimed_next.add(next_id) + links.append((det_id, next_id, cost)) + return links diff --git a/ami/main/models_future/tracks.py b/ami/ml/post_processing/tracking/sessions.py similarity index 96% rename from ami/main/models_future/tracks.py rename to ami/ml/post_processing/tracking/sessions.py index 6dd5b7b77..f2f6403ec 100644 --- a/ami/main/models_future/tracks.py +++ b/ami/ml/post_processing/tracking/sessions.py @@ -1,4 +1,4 @@ -"""Operations that keep tracked occurrences consistent with session boundaries. +"""Session locking and splitting of occurrences at session boundaries, for tracking and regrouping. Tracking links the detections of one insect through ``Detection.next_detection`` and attaches the chain to a single occurrence. Chains never cross a session boundary, but regrouping @@ -16,7 +16,6 @@ from ami.main.models import Detection, Event, Identification, Occurrence, SourceImage, update_occurrence_determination # Order of detections within an occurrence: capture time, then capture, then detection. -# The split and the tracks export both use it, so they agree on what "next" is. CAPTURE_ORDER = (F("source_image__timestamp").asc(nulls_last=True), "source_image_id", "pk") diff --git a/ami/ml/post_processing/tracking_task.py b/ami/ml/post_processing/tracking/task.py similarity index 66% rename from ami/ml/post_processing/tracking_task.py rename to ami/ml/post_processing/tracking/task.py index 76e5fe15e..3c967b80b 100644 --- a/ami/ml/post_processing/tracking_task.py +++ b/ami/ml/post_processing/tracking/task.py @@ -1,10 +1,14 @@ +"""The tracking post-processing task: links detections across consecutive captures and merges each chain. + +The matching rules live in ``matching.py`` and the settings in ``config.py``, both free of Django; +this module reads and writes the database around them. +""" + import collections import logging -import math import typing from collections.abc import Iterable, Iterator, Sequence -import pydantic from django.db import transaction from django.db.models import Count, Exists, OuterRef from django.utils import timezone @@ -20,187 +24,12 @@ update_calculated_fields_for_sessions_and_stations, update_occurrence_determination, ) -from ami.main.models_future.tracks import lock_sessions from ami.ml.models import Algorithm from ami.ml.post_processing.base import BasePostProcessingTask -COST_NOTE = ( - "The default is a starting point that is still being tuned by experiment. " - "It suits captures taken about 20 seconds apart." -) - - -class TrackingConfig(pydantic.BaseModel): - """Scope and tunables for a tracking run. - - Scope: exactly one of ``source_image_collection_id`` or ``event_ids`` says - which sessions to track. A capture set is the bulk path; an explicit event - list is what the Events admin page sends. - - The matching cost between two detections in consecutive captures is - ``iou_weight * (1 - IoU) + size_weight * (1 - size ratio) + distance_weight * (distance / diagonal)``. - Two detections are linked only when the cost is below ``cost_threshold`` and - every enabled limit passes. The field titles and descriptions are the help text - shown on the admin form. - """ - - source_image_collection_id: int | None = None - event_ids: list[int] = [] - - cost_threshold: float = pydantic.Field( - 1.0, - title="Cost cutoff", - ge=0, - description=( - "Two detections in neighbouring captures are only linked when their matching cost is below this " - "value. The cost adds up how little the boxes overlap, how different their sizes are, and how far " - "apart their centres are. Lower values link fewer detections and make fewer mistakes. " + COST_NOTE - ), - ) - iou_weight: float = pydantic.Field( - 1.0, - title="Overlap weight", - ge=0, - description=("How strongly poor overlap between two boxes raises the cost. 0 ignores overlap. " + COST_NOTE), - ) - size_weight: float = pydantic.Field( - 1.0, - title="Size weight", - ge=0, - description=("How strongly a difference in box area raises the cost. 0 ignores size. " + COST_NOTE), - ) - distance_weight: float = pydantic.Field( - 1.0, - title="Distance weight", - ge=0, - description=( - "How strongly the distance between box centres, measured as a share of the image diagonal, " - "raises the cost. 0 ignores distance. " + COST_NOTE - ), - ) - - min_iou: float | None = pydantic.Field( - None, - title="Minimum overlap", - ge=0, - le=1, - description=( - "Never link two detections whose boxes overlap by less than this (0 to 1, where 1 is identical " - "boxes). Leave blank for no limit." - ), - ) - min_size_ratio: float | None = pydantic.Field( - None, - title="Minimum size ratio", - ge=0, - le=1, - description=( - "Never link two detections when the smaller box has less than this share of the area of the larger " - "one (0 to 1). Leave blank for no limit." - ), - ) - max_distance: float | None = pydantic.Field( - None, - title="Maximum distance", - ge=0, - description=( - "Never link two detections whose box centres are further apart than this share of the image " - "diagonal (for example 0.1 is ten percent). Leave blank for no limit." - ), - ) - max_capture_interval_seconds: float | None = pydantic.Field( - None, - title="Maximum time between captures", - gt=0, - description=( - "Never link detections in two neighbouring captures taken further apart than this many seconds. " - "Leave blank for no limit." - ), - ) - - skip_if_human_identifications: bool = pydantic.Field( - True, - title="Skip sessions with human identifications", - description="Leave a session alone when someone has already identified one of its occurrences.", - ) - require_fresh_event: bool = pydantic.Field( - True, - title="Only track sessions that have not been tracked", - description=( - "Skip a session when any of its occurrences already holds more than one detection. Turned off, a " - "run can add links and merges to such a session, but it never undoes earlier ones." - ), - ) - - @pydantic.root_validator(skip_on_failure=True) - def _exactly_one_scope(cls, values: dict) -> dict: - scopes = [values.get("source_image_collection_id"), values.get("event_ids") or None] - if sum(s is not None for s in scopes) != 1: - raise ValueError("Provide exactly one of source_image_collection_id or event_ids") - return values - - class Config: - extra = "forbid" - - -def iou(bb1, bb2) -> float: - xA = max(bb1[0], bb2[0]) - yA = max(bb1[1], bb2[1]) - xB = min(bb1[2], bb2[2]) - yB = min(bb1[3], bb2[3]) - inter = max(0, xB - xA + 1) * max(0, yB - yA + 1) - area1 = (bb1[2] - bb1[0] + 1) * (bb1[3] - bb1[1] + 1) - area2 = (bb2[2] - bb2[0] + 1) * (bb2[3] - bb2[1] + 1) - union = area1 + area2 - inter - return inter / union if union > 0 else 0.0 - - -def box_ratio(bb1, bb2) -> float: - area1 = (bb1[2] - bb1[0] + 1) * (bb1[3] - bb1[1] + 1) - area2 = (bb2[2] - bb2[0] + 1) * (bb2[3] - bb2[1] + 1) - return min(area1, area2) / max(area1, area2) - - -def distance_ratio(bb1, bb2, img_diag: float) -> float: - cx1 = (bb1[0] + bb1[2]) / 2 - cy1 = (bb1[1] + bb1[3]) / 2 - cx2 = (bb2[0] + bb2[2]) / 2 - cy2 = (bb2[1] + bb2[3]) / 2 - dist = math.sqrt((cx2 - cx1) ** 2 + (cy2 - cy1) ** 2) - return dist / img_diag if img_diag > 0 else 1.0 - - -def image_diagonal(width: int, height: int) -> int: - return int(math.ceil(math.sqrt(width**2 + height**2))) - - -def pair_cost(bb1, bb2, diag: float, config: TrackingConfig) -> float | None: - """Matching cost between two detections; lower means more likely the same insect. - - Returns None when the pair fails an enabled limit, so it is never a candidate however - low its cost. With default weights the cost is the plain sum of the three terms. - """ - overlap = iou(bb1, bb2) - size_ratio = box_ratio(bb1, bb2) - distance = distance_ratio(bb1, bb2, diag) - if config.min_iou is not None and overlap < config.min_iou: - return None - if config.min_size_ratio is not None and size_ratio < config.min_size_ratio: - return None - if config.max_distance is not None and distance > config.max_distance: - return None - return ( - config.iou_weight * (1 - overlap) + config.size_weight * (1 - size_ratio) + config.distance_weight * distance - ) - - -def captures_too_far_apart(first: SourceImage, second: SourceImage, config: TrackingConfig) -> bool: - """Is the gap between two captures over the interval limit? A missing timestamp counts as too far.""" - if config.max_capture_interval_seconds is None: - return False - if first.timestamp is None or second.timestamp is None: - return True - return abs((second.timestamp - first.timestamp).total_seconds()) > config.max_capture_interval_seconds +from .config import TrackingConfig +from .matching import captures_too_far_apart, image_diagonal, select_links +from .sessions import lock_sessions def event_is_fresh(event: Event) -> tuple[bool, str]: @@ -345,39 +174,6 @@ def assign_occurrences_from_detection_chains( } -def select_links( - current_detections: Sequence[Detection], - next_detections: Sequence[Detection], - diag: float, - config: TrackingConfig, -) -> list[tuple[Detection, Detection, float]]: - """The links to make between two adjacent captures, lowest cost first; nothing is saved. - - A pair is a candidate when it passes every enabled limit and its cost is below the cutoff. - Candidates are taken lowest cost first, and each detection is linked at most once on either side. - """ - candidates: list[tuple[Detection, Detection, float]] = [] - for det in current_detections: - for nxt in next_detections: - cost = pair_cost(det.bbox, nxt.bbox, diag, config) - if cost is not None and cost < config.cost_threshold: - candidates.append((det, nxt, cost)) - - # Secondary keys keep tied costs deterministic across runs. - candidates.sort(key=lambda x: (x[2], x[0].pk, x[1].pk)) - - claimed_current: set[int] = set() - claimed_next: set[int] = set() - links: list[tuple[Detection, Detection, float]] = [] - for det, nxt, cost in candidates: - if det.pk in claimed_current or nxt.pk in claimed_next: - continue - claimed_current.add(det.pk) - claimed_next.add(nxt.pk) - links.append((det, nxt, cost)) - return links - - def save_links(links: Iterable[tuple[Detection, Detection, float]], logger: logging.Logger) -> None: """Store each link as ``next_detection``, first detaching any other detection that points at the target.""" links = list(links) @@ -401,19 +197,22 @@ def iter_transition_links( transitions = len(source_images) - 1 for i in range(transitions): cur, nxt = source_images[i], source_images[i + 1] - if captures_too_far_apart(cur, nxt, config): + if captures_too_far_apart(cur.timestamp, nxt.timestamp, config): yield [] continue if not cur.width or not cur.height: logger.warning(f"Capture {cur.pk} has no dimensions; not comparing it with the next capture.") yield None continue - yield select_links( - list(cur.detections.valid()), - list(nxt.detections.valid()), + current = {det.pk: det for det in cur.detections.valid()} + following = {det.pk: det for det in nxt.detections.valid()} + links = select_links( + [(pk, det.bbox) for pk, det in current.items()], + [(pk, det.bbox) for pk, det in following.items()], image_diagonal(cur.width, cur.height), config, ) + yield [(current[from_id], following[to_id], cost) for from_id, to_id, cost in links] def assign_occurrences_by_tracking_images( @@ -432,7 +231,8 @@ def assign_occurrences_by_tracking_images( transitions = len(source_images) - 1 links = skipped_for_dimensions = 0 skipped_for_interval = sum( - captures_too_far_apart(source_images[i], source_images[i + 1], config) for i in range(transitions) + captures_too_far_apart(source_images[i].timestamp, source_images[i + 1].timestamp, config) + for i in range(transitions) ) # Per-session atomic boundary: a crash mid-session rolls back this session only. with transaction.atomic(): From d503ed32a1fa899a52a7d6634fb4fb88943fc798 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Mon, 5 Oct 2026 17:19:47 -0700 Subject: [PATCH 09/28] chore(migrations): renumber the next_detection migration after the algorithm result migrations The branch now sits on the post-processing results branch, which adds main/0096 to main/0099. The detection link migration follows them as main/0100. Co-Authored-By: Claude Opus 5.5 --- ...ction_next_detection.py => 0100_detection_next_detection.py} | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) rename ami/main/migrations/{0096_detection_next_detection.py => 0100_detection_next_detection.py} (91%) diff --git a/ami/main/migrations/0096_detection_next_detection.py b/ami/main/migrations/0100_detection_next_detection.py similarity index 91% rename from ami/main/migrations/0096_detection_next_detection.py rename to ami/main/migrations/0100_detection_next_detection.py index f20634eab..7994ac6f5 100644 --- a/ami/main/migrations/0096_detection_next_detection.py +++ b/ami/main/migrations/0100_detection_next_detection.py @@ -6,7 +6,7 @@ class Migration(migrations.Migration): dependencies = [ - ("main", "0095_grant_sync_deployment_to_mldatamanager"), + ("main", "0099_classification_algorithm_result_index"), ] operations = [ From 4d55301785100a5a6d0a3e71459aff4f939e05ca Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Mon, 5 Oct 2026 17:25:16 -0700 Subject: [PATCH 10/28] feat(tracking): record a tracking result for each occurrence a run links A tracking run now leaves one algorithm result per occurrence it built from two or more detections or by merging occurrences. The result holds the figures that were previously planned as occurrence columns: the number of detections, how far the insect moved relative to the image diagonal (also stored as the result's value so lists can sort on it), how much its box changed size, how many taxa its classifications name, how much of them agree with the determination after the run, and which occurrences were folded in. The figures are computed by a new pure module, tracking/stats.py, from the chains already in memory plus one query for their terminal classifications. Results of occurrences absorbed by a merge move onto the keeper before the absorbed occurrences are deleted, so their history is no longer cascaded away. The classification that records a changed determination now carries the job and points at the occurrence's tracking result. The job also reports how many occurrences were recorded. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01QT59KePky4u4nbCsTAggtc --- .../tests/test_tracking_stats.py | 65 ++++++++++ .../tests/test_tracking_task.py | 99 ++++++++++++++- ami/ml/post_processing/tracking/stats.py | 81 ++++++++++++ ami/ml/post_processing/tracking/task.py | 115 ++++++++++++++++-- ami/ml/results/schemas.py | 24 +++- 5 files changed, 370 insertions(+), 14 deletions(-) create mode 100644 ami/ml/post_processing/tests/test_tracking_stats.py create mode 100644 ami/ml/post_processing/tracking/stats.py diff --git a/ami/ml/post_processing/tests/test_tracking_stats.py b/ami/ml/post_processing/tests/test_tracking_stats.py new file mode 100644 index 000000000..81bf805de --- /dev/null +++ b/ami/ml/post_processing/tests/test_tracking_stats.py @@ -0,0 +1,65 @@ +from django.test import SimpleTestCase + +from ami.ml.post_processing.tracking import stats + + +class TestMotion(SimpleTestCase): + def test_a_still_insect_has_no_motion(self): + self.assertEqual(stats.motion([[0, 0, 10, 10]] * 3, diagonal=100.0), 0.0) + + def test_motion_sums_the_steps_between_consecutive_centres(self): + boxes = [[0, 0, 10, 10], [30, 0, 40, 10], [30, 40, 40, 50]] + self.assertEqual(stats.motion(boxes, diagonal=100.0), 0.7) + + def test_a_single_box_has_no_motion(self): + self.assertEqual(stats.motion([[0, 0, 10, 10]], diagonal=100.0), 0.0) + + +class TestFrameDiagonal(SimpleTestCase): + def test_uses_the_largest_capture_size(self): + self.assertEqual(stats.frame_diagonal([(300, 100), (600, 800)], []), 1000.0) + + def test_falls_back_to_the_farthest_box_corner_without_dimensions(self): + self.assertEqual(stats.frame_diagonal([(None, None)], [[0, 0, 30, 10], [0, 0, 10, 40]]), 50.0) + + def test_is_one_when_nothing_is_known(self): + self.assertEqual(stats.frame_diagonal([], []), 1.0) + + +class TestSizeRatio(SimpleTestCase): + def test_is_largest_area_over_smallest(self): + self.assertEqual(stats.size_ratio([[0, 0, 10, 10], [0, 0, 20, 20], [0, 0, 5, 20]]), 4.0) + + def test_areas_are_floored_at_one(self): + self.assertEqual(stats.size_ratio([[0, 0, 0, 0], [0, 0, 10, 10]]), 100.0) + + def test_is_one_without_boxes(self): + self.assertEqual(stats.size_ratio([]), 1.0) + + +class TestLabels(SimpleTestCase): + def test_distinct_taxa_ignores_labels_without_a_taxon(self): + self.assertEqual(stats.distinct_taxa([1, 2, 2, None]), 2) + + def test_agreement_is_the_share_naming_the_determination(self): + self.assertEqual(stats.id_agreement([1, 1, 2, 3], 1), 0.5) + + def test_agreement_is_none_without_labels(self): + self.assertIsNone(stats.id_agreement([], 1)) + + def test_missing_taxa_never_agree_with_a_missing_determination(self): + self.assertEqual(stats.id_agreement([None, 1], None), 0.0) + + +class TestOccurrenceFigures(SimpleTestCase): + def test_collects_every_figure(self): + figures = stats.occurrence_figures( + boxes=[[0, 0, 10, 10], [30, 0, 40, 10]], + sizes=[(600, 800), (600, 800)], + labels=[1, 2], + determination_id=2, + ) + self.assertEqual( + figures, + stats.OccurrenceFigures(detection_count=2, motion=0.03, size_ratio=1.0, distinct_taxa=2, id_agreement=0.5), + ) diff --git a/ami/ml/post_processing/tests/test_tracking_task.py b/ami/ml/post_processing/tests/test_tracking_task.py index 962c3872b..93f40a8b3 100644 --- a/ami/ml/post_processing/tests/test_tracking_task.py +++ b/ami/ml/post_processing/tests/test_tracking_task.py @@ -6,7 +6,7 @@ from ami.jobs.models import Job from ami.main.models import Classification, Detection, Event, Identification, Occurrence, Taxon -from ami.ml.models import Algorithm +from ami.ml.models import Algorithm, AlgorithmResult from ami.ml.post_processing.registry import get_postprocessing_task from ami.ml.post_processing.tracking import TrackingTask from ami.ml.post_processing.tracking.task import assign_occurrences_from_detection_chains @@ -213,6 +213,103 @@ def test_a_run_records_the_determination_under_the_task_algorithm(self): self.assertEqual(Classification.objects.filter(algorithm=task.algorithm).count(), 1) +class TestTrackingResults(_TrackingCase): + def make_job(self) -> Job: + job = Job.objects.create(name="t", project=self.project, job_type_key="post_processing") + job.progress.add_stage("Post-processing", key="post_processing") + job.save() + return job + + def test_each_linked_occurrence_gets_a_result_with_its_figures(self): + moving, still, lone = [100, 100, 200, 200], [700, 700, 800, 800], [400, 800, 420, 820] + captures = create_session( + self.deployment, + [[moving, still], [[110, 100, 210, 200], still], [[120, 100, 220, 200], still, lone]], + self.taxa[0], + ) + event = captures[0].event + job = self.make_job() + + task = self.run_task(event, job=job) + + results = {r.occurrence_id: r for r in AlgorithmResult.objects.filter(kind="tracking")} + self.assertEqual(len(results), 2) + self.assertNotIn(captures[2].detections.get(bbox=lone).occurrence_id, results) + by_motion = sorted(results.values(), key=lambda r: r.value) + still_result, moving_result = by_motion + self.assertEqual(still_result.value, 0.0) + self.assertEqual(moving_result.value, 0.0141) + for result in by_motion: + self.assertEqual( + (result.job_id, result.algorithm_id, result.is_current), (job.pk, task.algorithm.pk, True) + ) + self.assertEqual(result.project_id, self.project.pk) + self.assertEqual(result.data["detection_count"], 3) + self.assertEqual(result.data["size_ratio"], 1.0) + self.assertEqual(result.data["distinct_taxa"], 1) + self.assertEqual(result.data["id_agreement"], 1.0) + self.assertEqual(result.data["determination_after_id"], self.taxa[0].pk) + self.assertEqual(len(result.data["merged_occurrence_ids"]), 2) + job.refresh_from_db() + params = {p.name: p.value for p in job.progress.get_stage("post_processing").params} + self.assertEqual(params["Occurrences recorded"], 2) + + def test_figures_follow_the_labels_and_boxes_of_the_merged_detections(self): + captures = create_session(self.deployment, [[[100, 100, 400, 400]], [[100, 100, 390, 400]]], self.taxa[0]) + second = captures[1].detections.get() + Classification.objects.filter(detection=second).update(taxon=self.taxa[1], score=0.6) + second.occurrence.save() + + self.run_task(captures[0].event) + + result = AlgorithmResult.objects.get(kind="tracking") + self.assertEqual(result.data["distinct_taxa"], 2) + self.assertEqual(result.data["size_ratio"], round(300 * 300 / (290 * 300), 4)) + self.assertEqual(result.data["id_agreement"], 0.5) + self.assertEqual(result.data["determination_before_id"], self.taxa[0].pk) + + def test_the_determination_classification_points_at_the_job_and_the_result(self): + captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) + for capture, taxon, score in zip(captures, self.taxa, (0.3, 0.9)): + Classification.objects.filter(detection__source_image=capture).update(taxon=taxon, score=score) + capture.detections.get().occurrence.save() + job = self.make_job() + + task = self.run_task(captures[0].event, job=job) + + record = Classification.objects.get(algorithm=task.algorithm) + result = AlgorithmResult.objects.get(kind="tracking") + self.assertEqual((record.job_id, record.algorithm_result_id), (job.pk, result.pk)) + self.assertEqual(result.data["determination_after_id"], self.taxa[1].pk) + self.assertNotEqual(result.data["determination_before_id"], result.data["determination_after_id"]) + # The tracking classification repeats the winner, so it is not counted as a label. + self.assertEqual(result.data["id_agreement"], 0.5) + + def test_results_of_merged_occurrences_move_onto_the_keeper(self): + captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) + first, second = (c.detections.get().occurrence for c in captures) + filter_algorithm = Algorithm.objects.create(name="Test filter", key="test-filter") + earlier = AlgorithmResult.objects.record( + occurrence=second, algorithm=filter_algorithm, kind="size_filter", data={"relative_size": 0.01} + ) + + self.run_task(captures[0].event) + + earlier.refresh_from_db() + self.assertEqual((earlier.occurrence_id, earlier.is_current), (first.pk, True)) + self.assertFalse(Occurrence.objects.filter(pk=second.pk).exists()) + tracking = AlgorithmResult.objects.get(kind="tracking") + self.assertEqual(tracking.occurrence_id, first.pk) + self.assertEqual(tracking.data["merged_occurrence_ids"], [second.pk]) + + def test_a_session_without_links_records_no_results(self): + captures = create_session(self.deployment, [[[0, 0, 50, 50]], [[900, 900, 950, 950]]], self.taxa[0]) + + self.run_task(captures[0].event) + + self.assertFalse(AlgorithmResult.objects.exists()) + + class TestTrackingJobMetrics(_TrackingCase): def test_result_line_says_what_was_tracked(self): captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) diff --git a/ami/ml/post_processing/tracking/stats.py b/ami/ml/post_processing/tracking/stats.py new file mode 100644 index 000000000..41db06584 --- /dev/null +++ b/ami/ml/post_processing/tracking/stats.py @@ -0,0 +1,81 @@ +"""Figures that describe one occurrence's path, computed from its detections' boxes and labels. + +Nothing here touches Django or the database. Boxes are ``(x1, y1, x2, y2)`` in capture order, and +labels are the taxon ids of the occurrence's terminal classifications. +""" + +import math +from collections.abc import Sequence +from dataclasses import dataclass + +BBox = Sequence[float] + +_MIN_AREA = 1.0 +_ROUND_TO = 4 + + +@dataclass(frozen=True) +class OccurrenceFigures: + detection_count: int + motion: float + size_ratio: float + distinct_taxa: int + id_agreement: float | None + + +def frame_diagonal(sizes: Sequence[tuple[int | None, int | None]], boxes: Sequence[BBox]) -> float: + """The length motion is divided by, in the pixel units of the boxes. + + It is the diagonal of the largest capture width and height seen. Captures without dimensions fall back + to the farthest box corner, and to 1.0 when there is no box either. + """ + widths = [width for width, _ in sizes if width] + heights = [height for _, height in sizes if height] + if widths and heights: + return math.hypot(max(widths), max(heights)) + far_x = max((box[2] for box in boxes), default=0.0) + far_y = max((box[3] for box in boxes), default=0.0) + return math.hypot(far_x, far_y) or 1.0 + + +def motion(boxes: Sequence[BBox], diagonal: float) -> float: + """Path length between consecutive box centres, as a fraction of ``diagonal``.""" + centres = [((box[0] + box[2]) / 2, (box[1] + box[3]) / 2) for box in boxes] + path = sum(math.dist(a, b) for a, b in zip(centres, centres[1:])) + return round(path / diagonal, _ROUND_TO) + + +def size_ratio(boxes: Sequence[BBox]) -> float: + """Largest box area over the smallest, with areas floored at 1 so a degenerate box cannot divide by zero.""" + areas = [max(abs((box[2] - box[0]) * (box[3] - box[1])), _MIN_AREA) for box in boxes] + if not areas: + return 1.0 + return round(max(areas) / min(areas), _ROUND_TO) + + +def distinct_taxa(labels: Sequence[int | None]) -> int: + """How many different taxa the labels name; a label without a taxon counts for none.""" + return len({label for label in labels if label is not None}) + + +def id_agreement(labels: Sequence[int | None], determination_id: int | None) -> float | None: + """The share of labels naming the determination, or None when there are no labels.""" + if not labels: + return None + return round(sum(label == determination_id and label is not None for label in labels) / len(labels), _ROUND_TO) + + +def occurrence_figures( + boxes: Sequence[BBox], + sizes: Sequence[tuple[int | None, int | None]], + labels: Sequence[int | None], + determination_id: int | None, +) -> OccurrenceFigures: + """All the figures for one occurrence: ``sizes`` holds the width and height of each detection's capture.""" + return OccurrenceFigures( + detection_count=len(boxes), + motion=motion(boxes, frame_diagonal(sizes, boxes)), + size_ratio=size_ratio(boxes), + distinct_taxa=distinct_taxa(labels), + id_agreement=id_agreement(labels, determination_id), + ) diff --git a/ami/ml/post_processing/tracking/task.py b/ami/ml/post_processing/tracking/task.py index 3c967b80b..222b864bb 100644 --- a/ami/ml/post_processing/tracking/task.py +++ b/ami/ml/post_processing/tracking/task.py @@ -5,6 +5,7 @@ """ import collections +import dataclasses import logging import typing from collections.abc import Iterable, Iterator, Sequence @@ -24,12 +25,26 @@ update_calculated_fields_for_sessions_and_stations, update_occurrence_determination, ) -from ami.ml.models import Algorithm +from ami.ml.models import Algorithm, AlgorithmResult from ami.ml.post_processing.base import BasePostProcessingTask from .config import TrackingConfig from .matching import captures_too_far_apart, image_diagonal, select_links from .sessions import lock_sessions +from .stats import occurrence_figures + +if typing.TYPE_CHECKING: + from ami.jobs.models import Job + + +@dataclasses.dataclass +class LinkedOccurrence: + """An occurrence the run built from a chain of two or more detections, or by merging occurrences.""" + + keeper: Occurrence + detections: list[Detection] + determination_before_id: int | None + merged_ids: list[int] def event_is_fresh(event: Event) -> tuple[bool, str]: @@ -68,12 +83,18 @@ def processed_captures(event: Event) -> list[SourceImage]: ) -def record_tracking_determination(occurrence: Occurrence, algorithm: Algorithm) -> Classification | None: +def record_tracking_determination( + occurrence: Occurrence, + algorithm: Algorithm, + job: "Job | None" = None, + result: AlgorithmResult | None = None, +) -> Classification | None: """Leave a terminal classification by the tracking algorithm after a merge changed the determination. The row carries the winning prediction and points back at it through ``applied_to``, the way class masking does, so the history shows what tracking decided. Nothing is written when the - winner already came from the tracking algorithm. + winner already came from the tracking algorithm. The row points at the run's ``job`` and at the + occurrence's tracking ``result`` so the history shows it under that result. """ winner = occurrence.best_prediction if winner is None or winner.detection_id is None or winner.taxon_id is None: @@ -88,19 +109,71 @@ class masking does, so the history shows what tracking decided. Nothing is writt algorithm=algorithm, timestamp=timezone.now(), applied_to=winner, + job=job, + algorithm_result=result, ) +def record_tracking_results( + linked: Sequence[LinkedOccurrence], algorithm: Algorithm, job: "Job | None" +) -> dict[int, AlgorithmResult]: + """Write one tracking result per linked occurrence, keyed by occurrence id. + + The figures are computed in memory from the chains' detections plus one query for their terminal + classifications. Classifications by the tracking algorithm itself are left out, since they repeat + the winner and would count as an extra vote. Call after the determinations have settled. + """ + if not linked: + return {} + detection_ids = [d.pk for item in linked for d in item.detections] + labels: dict[int, list[int | None]] = collections.defaultdict(list) + for detection_id, taxon_id in ( + Classification.objects.filter(detection_id__in=detection_ids, terminal=True) + .exclude(algorithm=algorithm) + .values_list("detection_id", "taxon_id") + ): + labels[detection_id].append(taxon_id) + + results = [] + for item in linked: + figures = occurrence_figures( + boxes=[d.bbox for d in item.detections], + sizes=[(d.source_image.width, d.source_image.height) for d in item.detections], + labels=[label for d in item.detections for label in labels.get(d.pk, [])], + determination_id=item.keeper.determination_id, + ) + results.append( + AlgorithmResult( + occurrence=item.keeper, + algorithm=algorithm, + job=job, + kind=AlgorithmResult.Kind.TRACKING, + value=figures.motion, + data={ + **dataclasses.asdict(figures), + "determination_before_id": item.determination_before_id, + "determination_after_id": item.keeper.determination_id, + "merged_occurrence_ids": item.merged_ids, + }, + ) + ) + return {result.occurrence_id: result for result in AlgorithmResult.objects.record_many(results)} + + def assign_occurrences_from_detection_chains( - source_images: Sequence[SourceImage], logger: logging.Logger, record_as: Algorithm | None = None + source_images: Sequence[SourceImage], + logger: logging.Logger, + record_as: Algorithm | None = None, + job: "Job | None" = None, ) -> dict[str, int]: """Fold each chain of linked detections into one occurrence, keeping the first existing one. A chain never leaves the given captures, which belong to one session, so a link that crosses a session boundary starts a new chain on each side. Identifications move onto the keeper before the occurrences that held them are deleted, because deleting an occurrence deletes its identifications. - With ``record_as`` set, a merge that changes the keeper's determination leaves a classification by - that algorithm. + Results of the absorbed occurrences move onto the keeper first. With ``record_as`` set, every keeper + built from two or more detections or from a merge gets a tracking result, and a merge that changes + its determination leaves a classification by that algorithm, both attributed to ``job``. """ image_ids = [image.pk for image in source_images] detections = list( @@ -114,6 +187,8 @@ def assign_occurrences_from_detection_chains( visited: set[int] = set() created = merged = identifications_moved = determinations_recorded = 0 + linked: dict[int, LinkedOccurrence] = {} + deleted: set[int] = set() existing = Occurrence.objects.filter(detections__source_image_id__in=image_ids).distinct().count() for det in detections: @@ -148,21 +223,33 @@ def assign_occurrences_from_detection_chains( doomed = old_occ_ids - {keeper.pk} if doomed: identifications_moved += Identification.objects.filter(occurrence_id__in=doomed).update(occurrence=keeper) + # Deleting an occurrence deletes its results, so they move onto the keeper first. + AlgorithmResult.objects.move_to_occurrence(keeper, doomed) Occurrence.objects.filter(pk__in=doomed).delete() + deleted |= doomed merged += len(doomed) # Only the determination is written, so a column this run did not change keeps its stored value. if update_occurrence_determination(keeper, save=False): keeper.save(update_determination=False, update_fields=["determination", "determination_score"]) - if record_as is not None and keeper.determination_id != previous_determination_id: - if record_tracking_determination(keeper, record_as) is not None: - determinations_recorded += 1 + if record_as is not None and (len(chain) > 1 or doomed): + linked[keeper.pk] = LinkedOccurrence(keeper, chain, previous_determination_id, sorted(doomed)) + + results: dict[int, AlgorithmResult] = {} + if record_as is not None: + # A keeper that a later chain absorbed no longer exists. + final = [item for pk, item in linked.items() if pk not in deleted] + results = record_tracking_results(final, record_as, job) + for item in final: + if item.keeper.determination_id != item.determination_before_id: + if record_tracking_determination(item.keeper, record_as, job, results.get(item.keeper.pk)) is not None: + determinations_recorded += 1 new_count = Occurrence.objects.filter(detections__source_image_id__in=image_ids).distinct().count() logger.info( f"Created {created} occurrences and merged {merged} across {len(image_ids)} captures " f"(occurrences before: {existing}, after: {new_count}). Moved {identifications_moved} identification(s), " - f"recorded {determinations_recorded} determination change(s)." + f"recorded {len(results)} result(s) and {determinations_recorded} determination change(s)." ) return { "occurrences_before": existing, @@ -171,6 +258,7 @@ def assign_occurrences_from_detection_chains( "occurrences_merged": merged, "identifications_moved": identifications_moved, "determinations_recorded": determinations_recorded, + "results_recorded": len(results), } @@ -221,6 +309,7 @@ def assign_occurrences_by_tracking_images( config: TrackingConfig, progress_cb: typing.Callable[[float], None] | None = None, record_as: Algorithm | None = None, + job: "Job | None" = None, ) -> dict[str, int]: """Link the detections of one session's processed captures and fold the chains into occurrences.""" source_images = processed_captures(event) @@ -244,7 +333,7 @@ def assign_occurrences_by_tracking_images( links += len(proposed) if progress_cb: progress_cb((i + 1) / transitions) - counters = assign_occurrences_from_detection_chains(source_images, logger, record_as=record_as) + counters = assign_occurrences_from_detection_chains(source_images, logger, record_as=record_as, job=job) counters["links_created"] = links counters["transitions_too_far_apart"] = skipped_for_interval @@ -351,6 +440,7 @@ def _stage_progress(p: float, _idx=idx, _total=total) -> None: logger=self.logger, config=self.config, record_as=self.algorithm, + job=self.job, progress_cb=_stage_progress, ) if not counters: @@ -359,7 +449,7 @@ def _stage_progress(p: float, _idx=idx, _total=total) -> None: continue totals["tracked"] += 1 tracked_event_ids.append(event.pk) - for key in ("links_created", "occurrences_merged", "transitions_too_far_apart"): + for key in ("links_created", "occurrences_merged", "results_recorded", "transitions_too_far_apart"): totals[key] += counters.get(key, 0) # Merging occurrences changes the session and station counts, which no save refreshes. @@ -371,6 +461,7 @@ def _stage_progress(p: float, _idx=idx, _total=total) -> None: "Sessions skipped": totals["skipped"], "Detection links created": totals["links_created"], "Occurrences merged": totals["occurrences_merged"], + "Occurrences recorded": totals["results_recorded"], } if self.config.max_capture_interval_seconds is not None: metrics["Capture pairs too far apart to compare"] = totals["transitions_too_far_apart"] diff --git a/ami/ml/results/schemas.py b/ami/ml/results/schemas.py index f4be31b83..7bc4211df 100644 --- a/ami/ml/results/schemas.py +++ b/ami/ml/results/schemas.py @@ -68,7 +68,29 @@ class SizeFilterResultData(DeterminationSnapshot): relative_size: float -ALGORITHM_RESULT_DATA_MODELS: tuple[type[AlgorithmResultData], ...] = (ClassMaskingResultData, SizeFilterResultData) +class TrackingResultData(DeterminationSnapshot): + """Figures from the detections the run linked into the occurrence, in capture order.""" + + kind: ClassVar[str] = "tracking" + + detection_count: int + # Path length between consecutive detection centres, as a fraction of the image diagonal; 0 for a still insect. + motion: float + # The largest box area over the smallest, with areas floored at 1; 1 when the box never changed size. + size_ratio: float + # Distinct taxa among the terminal classifications of the occurrence's detections, from every algorithm. + distinct_taxa: int + # The share of those classifications naming the determination after the run; None when there are none. + id_agreement: float | None = None + # Occurrences the run folded into this one. They are deleted, so these are plain ids, not references. + merged_occurrence_ids: list[int] = [] + + +ALGORITHM_RESULT_DATA_MODELS: tuple[type[AlgorithmResultData], ...] = ( + ClassMaskingResultData, + SizeFilterResultData, + TrackingResultData, +) ALGORITHM_RESULT_DATA_SCHEMAS: dict[str, type[AlgorithmResultData]] = { model.kind: model for model in ALGORITHM_RESULT_DATA_MODELS From a3f5a5109b5ca3baaf67e4c508989edd6cffba68 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Mon, 5 Oct 2026 17:26:54 -0700 Subject: [PATCH 11/28] feat(ui): show the tracking result on an occurrence's history An occurrence that a tracking run linked now shows a card in its history with the figures the run recorded: the number of detections, how far the insect moved relative to the image diagonal, how much its box changed size, how many taxa its classifications name, how much of them agree with the determination, and how many occurrences were merged into it. The card sits beside the class masking and size filter cards and reuses their layout. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01QT59KePky4u4nbCsTAggtc --- .../models/occurrence-history.test.ts | 41 ++++++++++++++- .../models/occurrence-history.ts | 25 +++++++++- .../identification-card/algorithm-result.tsx | 50 +++++++++++++++++-- ui/src/utils/language.ts | 20 ++++++++ 4 files changed, 128 insertions(+), 8 deletions(-) diff --git a/ui/src/data-services/models/occurrence-history.test.ts b/ui/src/data-services/models/occurrence-history.test.ts index be8ba5440..6e9defaf4 100644 --- a/ui/src/data-services/models/occurrence-history.test.ts +++ b/ui/src/data-services/models/occurrence-history.test.ts @@ -111,6 +111,31 @@ const classMasking: ServerOccurrenceHistoryEntry = { value: 0.38, } +const tracking: ServerOccurrenceHistoryEntry = { + ...base, + algorithm: { id: 8, key: 'tracking', name: 'Occurrence tracking' }, + classifications: [], + data: { + detection_count: 3, + determination_after_id: 3, + determination_before_id: 3, + distinct_taxa: 1, + extra: {}, + id_agreement: null, + merged_occurrence_ids: [11, 12], + motion: 0.0141, + size_ratio: 1.15, + }, + data_references: {}, + determination_after: NOCTUA, + determination_before: NOCTUA, + id: 6, + is_current: true, + kind: 'tracking', + score: 0.0141, + type: 'algorithm_result', +} + const ownIdentification: HumanIdentification = { comment: '', createdAt: '2026-04-29T22:00:00', @@ -150,6 +175,18 @@ describe('getTimelineItems', () => { ]) }) + test('builds a card for a tracking result', () => { + const items = getTimelineItems({ + entries: [tracking], + identifications: [], + predictions: [], + }) + + expect(items).toMatchObject([ + { type: 'algorithm_result', entry: { kind: 'tracking' } }, + ]) + }) + test("reuses the occurrence's own records, which carry the viewer's permissions", () => { const prediction = ownPrediction('6', 7) const items = getTimelineItems({ @@ -209,10 +246,10 @@ describe('getTimelineItems', () => { }) test('drops results of a kind it has no card for, and predictions without a taxon', () => { - // A kind the server added before the UI has a card for it, e.g. tracking. + // A kind the server added before the UI has a card for it, e.g. a rank roll-up. const unknown = { ...classMasking, - kind: 'tracking', + kind: 'rank_rollup', } as unknown as ServerOccurrenceHistoryEntry expect( diff --git a/ui/src/data-services/models/occurrence-history.ts b/ui/src/data-services/models/occurrence-history.ts index 45d7fc660..3f603e2a8 100644 --- a/ui/src/data-services/models/occurrence-history.ts +++ b/ui/src/data-services/models/occurrence-history.ts @@ -81,6 +81,20 @@ export interface SizeFilterResultData extends ServerDeterminationSnapshot { relative_size: number } +export interface TrackingResultData extends ServerDeterminationSnapshot { + /** The share of those classifications naming the determination after the run; null when there are none. */ + id_agreement: number | null + detection_count: number + /** Distinct taxa among the terminal classifications of the occurrence's detections. */ + distinct_taxa: number + /** Occurrences the run folded into this one; they no longer exist. */ + merged_occurrence_ids: number[] + /** The path length between detection centres as a fraction of the image diagonal. */ + motion: number + /** The largest box area over the smallest. */ + size_ratio: number +} + export interface ServerIdentificationDetails { agreed_with_identification_id: number | null agreed_with_prediction_id: number | null @@ -131,9 +145,14 @@ export type SizeFilterResultEntry = ServerResultEntry< 'size_filter', SizeFilterResultData > +export type TrackingResultEntry = ServerResultEntry< + 'tracking', + TrackingResultData +> export type AlgorithmResultEntry = | ClassMaskingResultEntry | SizeFilterResultEntry + | TrackingResultEntry export interface IdentificationEntry extends ServerHistoryEntryBase { details: ServerIdentificationDetails @@ -168,7 +187,11 @@ export type TimelineItem = | { type: 'algorithm_result'; id: string; entry: AlgorithmResultEntry } /** The result kinds this UI has a card for; results of any other kind are skipped. */ -const ALGORITHM_RESULT_KINDS: string[] = ['class_masking', 'size_filter'] +const ALGORITHM_RESULT_KINDS: string[] = [ + 'class_masking', + 'size_filter', + 'tracking', +] export const convertHistoryTaxon = (taxon: ServerHistoryTaxon) => new Taxon({ ...taxon, id: `${taxon.id}`, cover_image_url: null }) diff --git a/ui/src/pages/occurrence-details/identification-card/algorithm-result.tsx b/ui/src/pages/occurrence-details/identification-card/algorithm-result.tsx index f262af5c1..a47355bc4 100644 --- a/ui/src/pages/occurrence-details/identification-card/algorithm-result.tsx +++ b/ui/src/pages/occurrence-details/identification-card/algorithm-result.tsx @@ -6,7 +6,7 @@ import { ServerHistoryTaxon, } from 'data-services/models/occurrence-history' import { OccurrenceDetails as Occurrence } from 'data-services/models/occurrence-details' -import { FilterIcon, RulerIcon } from 'lucide-react' +import { FilterIcon, RouteIcon, RulerIcon } from 'lucide-react' import { BasicTooltip, IdentificationCard, @@ -32,6 +32,7 @@ import { const KINDS = { class_masking: { icon: FilterIcon, label: STRING.HISTORY_CLASS_MASKING }, size_filter: { icon: RulerIcon, label: STRING.HISTORY_SIZE_FILTER }, + tracking: { icon: RouteIcon, label: STRING.HISTORY_TRACKING }, } /** What to call a result's kind, e.g. "Class masking", for the card and for predictions it superseded. */ @@ -153,11 +154,50 @@ export const AlgorithmResult = ({ }) break } + case 'tracking': { + const { data } = entry + stats.push( + { + label: translate(STRING.HISTORY_TRACKING_DETECTIONS), + value: data.detection_count, + }, + { + label: translate(STRING.HISTORY_TRACKING_MOVEMENT), + value: translate(STRING.HISTORY_TRACKING_MOVEMENT_VALUE, { + distance: formatPercent(data.motion), + }), + }, + { + label: translate(STRING.HISTORY_TRACKING_SIZE_CHANGE), + value: translate(STRING.HISTORY_TRACKING_SIZE_CHANGE_VALUE, { + ratio: `${Math.round(data.size_ratio * 100) / 100}`, + }), + }, + { + label: translate(STRING.HISTORY_TRACKING_TAXA), + value: data.distinct_taxa, + }, + { + label: translate(STRING.HISTORY_TRACKING_AGREEMENT), + value: + data.id_agreement !== null + ? formatPercent(data.id_agreement) + : translate(STRING.VALUE_NOT_AVAILABLE), + }, + { + label: translate(STRING.HISTORY_TRACKING_MERGED), + value: data.merged_occurrence_ids.length, + } + ) + break + } + } + if (entry.kind !== 'tracking') { + stats.push({ + label: translate(STRING.HISTORY_DETECTIONS_AFFECTED), + value: new Set(entry.classifications.map((c) => c.detection_id)).size, + }) } - stats.push({ - label: translate(STRING.HISTORY_DETECTIONS_AFFECTED), - value: new Set(entry.classifications.map((c) => c.detection_id)).size, - }) getJobSettings(entry.job).forEach(({ label, value, ref }) => { stats.push({ label, diff --git a/ui/src/utils/language.ts b/ui/src/utils/language.ts index dce93adc7..3a6d03e16 100644 --- a/ui/src/utils/language.ts +++ b/ui/src/utils/language.ts @@ -337,6 +337,15 @@ export enum STRING { HISTORY_RECORD_ID, HISTORY_SIZE_FILTER, HISTORY_SUPERSEDED_BY, + HISTORY_TRACKING, + HISTORY_TRACKING_AGREEMENT, + HISTORY_TRACKING_DETECTIONS, + HISTORY_TRACKING_MERGED, + HISTORY_TRACKING_MOVEMENT, + HISTORY_TRACKING_MOVEMENT_VALUE, + HISTORY_TRACKING_SIZE_CHANGE, + HISTORY_TRACKING_SIZE_CHANGE_VALUE, + HISTORY_TRACKING_TAXA, ID_APPLIED, INFO, INTERMEDIATE_CLASSIFICATION, @@ -794,6 +803,17 @@ const ENGLISH_STRINGS: { [key in STRING]: string } = { [STRING.HISTORY_RECORD_ID]: '#{{id}}', [STRING.HISTORY_SIZE_FILTER]: 'Size filter', [STRING.HISTORY_SUPERSEDED_BY]: 'Superseded by {{name}}', + [STRING.HISTORY_TRACKING]: 'Occurrence tracking', + [STRING.HISTORY_TRACKING_AGREEMENT]: 'Agreement', + [STRING.HISTORY_TRACKING_DETECTIONS]: 'Detections', + [STRING.HISTORY_TRACKING_MERGED]: 'Occurrences merged in', + [STRING.HISTORY_TRACKING_MOVEMENT]: 'Movement', + [STRING.HISTORY_TRACKING_MOVEMENT_VALUE]: + '{{distance}} of the image diagonal', + [STRING.HISTORY_TRACKING_SIZE_CHANGE]: 'Size change', + [STRING.HISTORY_TRACKING_SIZE_CHANGE_VALUE]: + '{{ratio}}× from smallest to largest', + [STRING.HISTORY_TRACKING_TAXA]: 'Taxa', [STRING.ID_APPLIED]: 'ID applied', [STRING.INFO]: 'Info', [STRING.INTERMEDIATE_CLASSIFICATION]: 'Intermediate classification', From 1245868d1cd85967af0a484b043794119916de71 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Mon, 5 Oct 2026 22:50:09 -0700 Subject: [PATCH 12/28] fix(migrations): add the detection next-detection link without blocking reads of the detection table Django adds a one-to-one column with its unique index and foreign key in a single statement, which holds a lock on the large detection table that blocks reads while the index is built. The column is now added on its own in 0100, a catalogue-only change, and 0101 builds the unique index concurrently, attaches it as the unique constraint, and adds the foreign key as NOT VALID before validating it. The constraint names are the ones Django generates, so later AlterField migrations still find them, and makemigrations --check stays clean. A database that already applied the earlier version of 0100 has the column and its constraints, so 0101 would fail there; the earlier version was never merged or deployed. Co-Authored-By: Claude Sonnet 5.5 Claude-Session: https://claude.ai/code/session_01QT59KePky4u4nbCsTAggtc --- .../0100_detection_next_detection.py | 43 +++++++++---- ...01_detection_next_detection_constraints.py | 64 +++++++++++++++++++ 2 files changed, 94 insertions(+), 13 deletions(-) create mode 100644 ami/main/migrations/0101_detection_next_detection_constraints.py diff --git a/ami/main/migrations/0100_detection_next_detection.py b/ami/main/migrations/0100_detection_next_detection.py index 7994ac6f5..825935ee0 100644 --- a/ami/main/migrations/0100_detection_next_detection.py +++ b/ami/main/migrations/0100_detection_next_detection.py @@ -1,25 +1,42 @@ -# Generated by Django 4.2.10 on 2026-10-04 12:24 - import django.db.models.deletion from django.db import migrations, models class Migration(migrations.Migration): + """Add the column that links a detection to the one that follows it in a tracking sequence. + + Django would add this column together with a unique constraint and a foreign key in one + statement, which takes a lock on the large detection table that blocks reads while the + unique index is built. Here the column is added on its own: a nullable column without a + default is a catalogue change that does not scan the table. Migration 0101 adds the unique + constraint and the foreign key afterwards without blocking readers or writers. + """ + dependencies = [ ("main", "0099_classification_algorithm_result_index"), ] operations = [ - migrations.AddField( - model_name="detection", - name="next_detection", - field=models.OneToOneField( - blank=True, - help_text="The detection that follows this one in the tracking sequence.", - null=True, - on_delete=django.db.models.deletion.SET_NULL, - related_name="previous_detection", - to="main.detection", - ), + migrations.SeparateDatabaseAndState( + state_operations=[ + migrations.AddField( + model_name="detection", + name="next_detection", + field=models.OneToOneField( + blank=True, + help_text="The detection that follows this one in the tracking sequence.", + null=True, + on_delete=django.db.models.deletion.SET_NULL, + related_name="previous_detection", + to="main.detection", + ), + ), + ], + database_operations=[ + migrations.RunSQL( + sql='ALTER TABLE "main_detection" ADD COLUMN "next_detection_id" bigint NULL;', + reverse_sql='ALTER TABLE "main_detection" DROP COLUMN "next_detection_id";', + ), + ], ), ] diff --git a/ami/main/migrations/0101_detection_next_detection_constraints.py b/ami/main/migrations/0101_detection_next_detection_constraints.py new file mode 100644 index 000000000..b0673903b --- /dev/null +++ b/ami/main/migrations/0101_detection_next_detection_constraints.py @@ -0,0 +1,64 @@ +from django.db import migrations + + +class Migration(migrations.Migration): + """Give ``Detection.next_detection`` the unique constraint and foreign key that Django would have created. + + Both are built so that neither blocks reads or writes on the large detection table. The + unique index is built CONCURRENTLY, which needs a non-atomic migration, and is then attached + as a constraint, which is a catalogue change. The foreign key is added NOT VALID, so existing + rows are not checked while a strong lock is held, and validated afterwards, which takes only a + light lock. See 0093 for why the statement timeout is cleared and restored around the build. + + The constraint names match the ones Django generates, so later AlterField migrations find them. + If the index build is interrupted it leaves an invalid index of the same name; drop it before retrying. + """ + + atomic = False + + dependencies = [ + ("main", "0100_detection_next_detection"), + ] + + operations = [ + migrations.RunSQL( + sql="SET statement_timeout = 0;", + reverse_sql=migrations.RunSQL.noop, + ), + migrations.RunSQL( + sql=( + 'CREATE UNIQUE INDEX CONCURRENTLY "main_detection_next_detection_id_key" ' + 'ON "main_detection" ("next_detection_id");' + ), + reverse_sql=migrations.RunSQL.noop, + ), + migrations.RunSQL( + sql=( + 'ALTER TABLE "main_detection" ADD CONSTRAINT "main_detection_next_detection_id_key" ' + 'UNIQUE USING INDEX "main_detection_next_detection_id_key";' + ), + reverse_sql='ALTER TABLE "main_detection" DROP CONSTRAINT "main_detection_next_detection_id_key";', + ), + migrations.RunSQL( + sql=( + 'ALTER TABLE "main_detection" ADD CONSTRAINT "main_detection_next_detection_id_f0201e13_fk_main_detection_id" ' + 'FOREIGN KEY ("next_detection_id") REFERENCES "main_detection" ("id") ' + "DEFERRABLE INITIALLY DEFERRED NOT VALID;" + ), + reverse_sql=( + 'ALTER TABLE "main_detection" ' + 'DROP CONSTRAINT "main_detection_next_detection_id_f0201e13_fk_main_detection_id";' + ), + ), + migrations.RunSQL( + sql=( + 'ALTER TABLE "main_detection" VALIDATE CONSTRAINT ' + '"main_detection_next_detection_id_f0201e13_fk_main_detection_id";' + ), + reverse_sql=migrations.RunSQL.noop, + ), + migrations.RunSQL( + sql="RESET statement_timeout;", + reverse_sql=migrations.RunSQL.noop, + ), + ] From c486d83ad0b6ce4b077e2a674a25742d5d3f0c08 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Mon, 5 Oct 2026 22:50:20 -0700 Subject: [PATCH 13/28] fix(regroup): look for occurrences split by a regroup only among the sessions it touched, under the sessions' locks The check that finds occurrences spanning two sessions after a regroup grouped every detection of the deployment's occurrences, which scanned the whole detection table on every capture sync. It now starts from the captures of the sessions the regroup touched, looks up the occurrences on them with literal id lists, and only then checks those occurrences for several sessions. The sessions are locked before the split, so a tracking run on one of them finishes first or waits. The tracking result stays on the earliest piece, which a test now pins. Co-Authored-By: Claude Sonnet 5.5 Claude-Session: https://claude.ai/code/session_01QT59KePky4u4nbCsTAggtc --- ami/main/models.py | 68 ++++++++++++++------- ami/main/tests.py | 20 ++++++ ami/ml/post_processing/tracking/sessions.py | 4 +- 3 files changed, 68 insertions(+), 24 deletions(-) diff --git a/ami/main/models.py b/ami/main/models.py index b214d111a..d59bd221d 100644 --- a/ami/main/models.py +++ b/ami/main/models.py @@ -1670,7 +1670,7 @@ def _group_images_into_events_locked( f"Done grouping {len(image_timestamps)} captures into {len(events)} events " f"for deployment {deployment}" ) - occurrences_split_count = _split_occurrences_at_session_boundaries(deployment, job) + occurrences_split_count = _split_occurrences_at_session_boundaries(job, touched_event_pks) # Realign Occurrence.event_id with each occurrence's detections' current # source_image.event_id. Occurrences are bound to an event once at creation @@ -1765,32 +1765,56 @@ def _group_images_into_events_locked( return events -def _split_occurrences_at_session_boundaries(deployment: Deployment, job: "Job | None") -> int: - """Split every occurrence in the deployment whose detections now span several sessions. +def _split_occurrences_at_session_boundaries(job: "Job | None", event_pks: set[int]) -> int: + """Split every occurrence whose detections now span several sessions, among the sessions a regroup touched. An occurrence is expected to belong to one session, so a regroup that draws a session - boundary through it leaves one piece per session. Returns how many occurrences were split. + boundary through it leaves one piece per session. Only an occurrence with a detection on a capture of + ``event_pks`` can have been cut, so the search starts from those captures. The ids are looked up + in steps with literal id lists: left as a subquery, Postgres scans the whole detection table + (measured on a copy of production data). Returns how many occurrences were split. """ - from ami.ml.post_processing.tracking.sessions import split_at_session_boundaries - - spanning_ids = list( - Detection.objects.valid() - .filter(occurrence__deployment=deployment) - .values("occurrence_id") - .annotate(sessions=models.Count("source_image__event", distinct=True)) - .filter(sessions__gt=1) - .values_list("occurrence_id", flat=True) - ) + from ami.ml.post_processing.tracking.sessions import lock_sessions, split_at_session_boundaries + + def find_spanning_ids() -> list[int]: + capture_ids = list(SourceImage.objects.filter(event_id__in=event_pks).values_list("pk", flat=True)) + touched_occurrence_ids = list( + Detection.objects.filter(source_image_id__in=capture_ids, occurrence__isnull=False) + .values_list("occurrence_id", flat=True) + .distinct() + ) + return list( + Detection.objects.valid() + .filter(occurrence_id__in=touched_occurrence_ids) + .values("occurrence_id") + .annotate(sessions=models.Count("source_image__event", distinct=True)) + .filter(sessions__gt=1) + .values_list("occurrence_id", flat=True) + ) + + candidate_ids = find_spanning_ids() + if not candidate_ids: + return 0 split_count = 0 - for occurrence in Occurrence.objects.filter(pk__in=spanning_ids).order_by("pk"): - pieces = split_at_session_boundaries(occurrence) - if not pieces: - continue - split_count += 1 - (job.logger if job else logger).info( - f"Split occurrence {occurrence.pk} at a session boundary; " - f"new occurrence(s) {[piece.pk for piece in pieces]} hold the later sessions." + # Holding the sessions' locks while the occurrences are found and split makes a tracking run + # on one of these sessions finish first, or wait for the split. + with transaction.atomic(): + lock_sessions( + list( + SourceImage.objects.filter(detections__occurrence_id__in=candidate_ids) + .values_list("event_id", flat=True) + .distinct() + ) ) + for occurrence in Occurrence.objects.filter(pk__in=find_spanning_ids()).order_by("pk"): + pieces = split_at_session_boundaries(occurrence) + if not pieces: + continue + split_count += 1 + (job.logger if job else logger).info( + f"Split occurrence {occurrence.pk} at a session boundary; " + f"new occurrence(s) {[piece.pk for piece in pieces]} hold the later sessions." + ) return split_count diff --git a/ami/main/tests.py b/ami/main/tests.py index 083341db3..ef73e5334 100644 --- a/ami/main/tests.py +++ b/ami/main/tests.py @@ -8662,3 +8662,23 @@ def test_identifications_are_copied_to_every_piece(self): ) self.assertEqual(occurrence.determination_id, self.other_taxon.pk) self.assertEqual(piece.determination_id, self.other_taxon.pk) + + def test_a_tracking_result_stays_on_the_earliest_piece(self): + from ami.ml.models.algorithm import Algorithm + + self._group(gap_hours=6) + occurrence, _ = self._make_occurrence(self.captures) + algorithm = Algorithm.objects.create(name="Test tracking", key="test-tracking") + result = AlgorithmResult.objects.record( + occurrence=occurrence, + algorithm=algorithm, + kind="tracking", + data={"detection_count": 6, "motion": 0.0, "path_length": 0.0, "size_change": 1.0, "distinct_taxa": 1}, + ) + + self._group(gap_hours=2) + + piece = Occurrence.objects.exclude(pk=occurrence.pk).get(deployment=self.deployment) + result.refresh_from_db() + self.assertEqual((result.occurrence_id, result.is_current), (occurrence.pk, True)) + self.assertFalse(AlgorithmResult.objects.filter(occurrence=piece).exists()) diff --git a/ami/ml/post_processing/tracking/sessions.py b/ami/ml/post_processing/tracking/sessions.py index f2f6403ec..fd906bd47 100644 --- a/ami/ml/post_processing/tracking/sessions.py +++ b/ami/ml/post_processing/tracking/sessions.py @@ -53,8 +53,8 @@ def split_at_session_boundaries(occurrence: Occurrence) -> list[Occurrence]: The piece in the earliest session keeps this occurrence and its identifications; each later piece is a new occurrence holding copies of them. The link between the last detection of one piece and the first of the next is kept, since tracking stops at session - boundaries and so never walks across it. Returns the new occurrences in time order, or an - empty list when nothing was split. + boundaries and so never walks across it. A tracking result stays on the earliest piece. + Returns the new occurrences in time order, or an empty list when nothing was split. """ sessions = SourceImage.objects.filter(detections__occurrence=occurrence).values_list("event_id", flat=True) lock_sessions([occurrence.event_id, *sessions]) From 940749e0d06af9c19f1a4c1f5079b703bbfc52c9 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Mon, 5 Oct 2026 22:50:35 -0700 Subject: [PATCH 14/28] refactor(tracking): drop the tracking classification, record link costs, and make a run survive a failing session A merged occurrence no longer gets a terminal classification from the tracking algorithm. That row could never be re-scored by class masking, so it pinned the determination to the unmasked taxon. The tracking result already records the determination before and after, and the determination is recomputed from the source classifications as before; a test runs class masking after tracking to pin that. The tracking result now holds each link's matching cost in chain order, the mean movement per step (with the total path length beside it), the size change renamed from size_ratio so it no longer clashes with the config's minimum size ratio, and a label agreement that counts only machine labels, leaving out post-processing classifications. A run now saves progress between sessions, outside their transactions, so the job row is not locked for a whole session. A session that fails is rolled back, logged and counted, the run continues, the counts of the tracked sessions are refreshed, and the run raises at the end so the job is marked failed. A chain's detections are reassigned with one update instead of one save each. The capture-set scope is documented as tracking every processed capture of the sessions the set touches. Tests: the guard-off test now observes a change, two redundant tests are merged, and new tests cover mid-run failure, capture-set scope, cross-project sessions, progress timing and the query count per capture. Co-Authored-By: Claude Sonnet 5.5 Claude-Session: https://claude.ai/code/session_01QT59KePky4u4nbCsTAggtc --- .../tests/test_tracking_matching.py | 4 +- .../tests/test_tracking_stats.py | 26 +- .../tests/test_tracking_task.py | 253 ++++++++++++++---- ami/ml/post_processing/tracking/config.py | 10 +- ami/ml/post_processing/tracking/stats.py | 46 +++- ami/ml/post_processing/tracking/task.py | 162 ++++++----- ami/ml/results/schemas.py | 15 +- 7 files changed, 346 insertions(+), 170 deletions(-) diff --git a/ami/ml/post_processing/tests/test_tracking_matching.py b/ami/ml/post_processing/tests/test_tracking_matching.py index 78f7b988a..b74605bb9 100644 --- a/ami/ml/post_processing/tests/test_tracking_matching.py +++ b/ami/ml/post_processing/tests/test_tracking_matching.py @@ -98,9 +98,9 @@ def test_links_are_one_to_one_and_lowest_cost_wins(self): self.assertEqual([(a, b) for a, b, _ in links], [(2, 3), (1, 4)]) self.assertLessEqual(links[0][2], links[1][2]) - def test_a_detection_is_linked_at_most_once(self): + # Two detections competing for one target: only one links, and a tie goes to the lower id. links = select_links([(1, BOX), (2, BOX)], [(3, BOX)], DIAG, _config()) - self.assertEqual([(a, b) for a, b, _ in links], [(1, 3)], "Ties break on the lower ids") + self.assertEqual([(a, b) for a, b, _ in links], [(1, 3)]) class TestIntervalLimit(SimpleTestCase): diff --git a/ami/ml/post_processing/tests/test_tracking_stats.py b/ami/ml/post_processing/tests/test_tracking_stats.py index 81bf805de..f102a754f 100644 --- a/ami/ml/post_processing/tests/test_tracking_stats.py +++ b/ami/ml/post_processing/tests/test_tracking_stats.py @@ -7,9 +7,13 @@ class TestMotion(SimpleTestCase): def test_a_still_insect_has_no_motion(self): self.assertEqual(stats.motion([[0, 0, 10, 10]] * 3, diagonal=100.0), 0.0) - def test_motion_sums_the_steps_between_consecutive_centres(self): + def test_motion_is_the_mean_step_between_consecutive_centres(self): boxes = [[0, 0, 10, 10], [30, 0, 40, 10], [30, 40, 40, 50]] - self.assertEqual(stats.motion(boxes, diagonal=100.0), 0.7) + self.assertEqual(stats.motion(boxes, diagonal=100.0), 0.35) + + def test_path_length_is_the_total_of_the_steps(self): + boxes = [[0, 0, 10, 10], [30, 0, 40, 10], [30, 40, 40, 50]] + self.assertEqual(stats.path_length(boxes, diagonal=100.0), 0.7) def test_a_single_box_has_no_motion(self): self.assertEqual(stats.motion([[0, 0, 10, 10]], diagonal=100.0), 0.0) @@ -26,15 +30,15 @@ def test_is_one_when_nothing_is_known(self): self.assertEqual(stats.frame_diagonal([], []), 1.0) -class TestSizeRatio(SimpleTestCase): +class TestSizeChange(SimpleTestCase): def test_is_largest_area_over_smallest(self): - self.assertEqual(stats.size_ratio([[0, 0, 10, 10], [0, 0, 20, 20], [0, 0, 5, 20]]), 4.0) + self.assertEqual(stats.size_change([[0, 0, 10, 10], [0, 0, 20, 20], [0, 0, 5, 20]]), 4.0) def test_areas_are_floored_at_one(self): - self.assertEqual(stats.size_ratio([[0, 0, 0, 0], [0, 0, 10, 10]]), 100.0) + self.assertEqual(stats.size_change([[0, 0, 0, 0], [0, 0, 10, 10]]), 100.0) def test_is_one_without_boxes(self): - self.assertEqual(stats.size_ratio([]), 1.0) + self.assertEqual(stats.size_change([]), 1.0) class TestLabels(SimpleTestCase): @@ -42,13 +46,13 @@ def test_distinct_taxa_ignores_labels_without_a_taxon(self): self.assertEqual(stats.distinct_taxa([1, 2, 2, None]), 2) def test_agreement_is_the_share_naming_the_determination(self): - self.assertEqual(stats.id_agreement([1, 1, 2, 3], 1), 0.5) + self.assertEqual(stats.label_agreement([1, 1, 2, 3], 1), 0.5) def test_agreement_is_none_without_labels(self): - self.assertIsNone(stats.id_agreement([], 1)) + self.assertIsNone(stats.label_agreement([], 1)) def test_missing_taxa_never_agree_with_a_missing_determination(self): - self.assertEqual(stats.id_agreement([None, 1], None), 0.0) + self.assertEqual(stats.label_agreement([None, 1], None), 0.0) class TestOccurrenceFigures(SimpleTestCase): @@ -61,5 +65,7 @@ def test_collects_every_figure(self): ) self.assertEqual( figures, - stats.OccurrenceFigures(detection_count=2, motion=0.03, size_ratio=1.0, distinct_taxa=2, id_agreement=0.5), + stats.OccurrenceFigures( + detection_count=2, motion=0.03, path_length=0.03, size_change=1.0, distinct_taxa=2, label_agreement=0.5 + ), ) diff --git a/ami/ml/post_processing/tests/test_tracking_task.py b/ami/ml/post_processing/tests/test_tracking_task.py index 93f40a8b3..abdad1050 100644 --- a/ami/ml/post_processing/tests/test_tracking_task.py +++ b/ami/ml/post_processing/tests/test_tracking_task.py @@ -1,15 +1,27 @@ import datetime import logging import typing +from unittest import mock +from django.db import connection from django.test import TestCase +from django.test.utils import CaptureQueriesContext from ami.jobs.models import Job -from ami.main.models import Classification, Detection, Event, Identification, Occurrence, Taxon +from ami.main.models import ( + Classification, + Detection, + Event, + Identification, + Occurrence, + SourceImageCollection, + Taxon, + update_occurrence_determination, +) from ami.ml.models import Algorithm, AlgorithmResult from ami.ml.post_processing.registry import get_postprocessing_task from ami.ml.post_processing.tracking import TrackingTask -from ami.ml.post_processing.tracking.task import assign_occurrences_from_detection_chains +from ami.ml.post_processing.tracking import task as task_module from ami.tests.fixtures.main import create_taxa, setup_test_project from ami.tests.fixtures.tracking import add_detection, create_session from ami.users.tests.factories import UserFactory @@ -52,14 +64,6 @@ def test_a_still_insect_is_folded_into_one_occurrence(self): event.refresh_from_db() self.assertEqual(event.occurrences_count, 1) - def test_two_insects_are_matched_one_to_one_by_lowest_cost(self): - left, right = [100, 100, 200, 200], [700, 700, 800, 800] - captures = create_session( - self.deployment, [[left, right], [[110, 100, 210, 200], [705, 700, 805, 800]]], self.taxa[0] - ) - self.run_task(captures[0].event) - self.assertEqual(self.occurrence_sizes(captures[0].event), [2, 2]) - def test_an_unprocessed_capture_between_processed_ones_does_not_break_the_chain(self): captures = create_session(self.deployment, [[BOX], None, [BOX]], self.taxa[0]) event = captures[0].event @@ -126,18 +130,18 @@ def test_a_session_with_a_single_processed_capture_is_skipped_with_a_reason(self class TestGuards(_TrackingCase): def test_a_session_that_was_already_tracked_is_skipped_unless_the_guard_is_off(self): - captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) + # The third capture is processed after the first run, with an insect that follows the others. + captures = create_session(self.deployment, [[BOX], [BOX], None], self.taxa[0]) event = captures[0].event self.run_task(event) - add_detection(captures[1], [500, 500, 560, 560], self.taxa[0]) + add_detection(captures[2], BOX, self.taxa[0]) self.assertEqual(self.occurrence_sizes(event), [1, 2]) self.run_task(event) self.assertEqual(self.occurrence_sizes(event), [1, 2]) - self.assertEqual(Detection.objects.filter(source_image__event=event, next_detection__isnull=False).count(), 1) self.run_task(event, require_fresh_event=False) - self.assertEqual(self.occurrence_sizes(event), [1, 2]) + self.assertEqual(self.occurrence_sizes(event), [3]) def test_human_identifications_skip_the_session_unless_the_guard_is_off(self): captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) @@ -177,40 +181,33 @@ def _two_detection_chain(self, first_score: float, second_score: float): first.save(update_fields=["next_detection"]) return captures, first, second - def test_a_changed_determination_is_recorded_with_the_winning_prediction_as_applied_to(self): + def test_class_masking_after_tracking_still_changes_the_determination(self): + """A later re-scoring replaces a merged occurrence's determination, because tracking adds no classification.""" captures, first, second = self._two_detection_chain(first_score=0.3, second_score=0.9) - tracking_algorithm = Algorithm.objects.create(name="Test tracking", key="test-tracking") - - assign_occurrences_from_detection_chains(captures, logger, record_as=tracking_algorithm) - - keeper = Occurrence.objects.get(pk=first.occurrence_id) - self.assertEqual(keeper.determination, self.taxa[1]) - record = Classification.objects.get(algorithm=tracking_algorithm) - self.assertEqual((record.detection_id, record.taxon, record.terminal), (second.pk, self.taxa[1], True)) - self.assertEqual((record.applied_to.detection_id, record.applied_to.taxon), (second.pk, self.taxa[1])) - self.assertNotEqual(record.applied_to.algorithm_id, tracking_algorithm.pk) - - assign_occurrences_from_detection_chains(captures, logger, record_as=tracking_algorithm) - self.assertEqual(Classification.objects.filter(algorithm=tracking_algorithm).count(), 1) - - def test_an_unchanged_determination_records_nothing(self): - captures, first, _ = self._two_detection_chain(first_score=0.9, second_score=0.3) - tracking_algorithm = Algorithm.objects.create(name="Test tracking", key="test-tracking") - - assign_occurrences_from_detection_chains(captures, logger, record_as=tracking_algorithm) - - self.assertEqual(Occurrence.objects.get(pk=first.occurrence_id).determination, self.taxa[0]) - self.assertFalse(Classification.objects.filter(algorithm=tracking_algorithm).exists()) - - def test_a_run_records_the_determination_under_the_task_algorithm(self): - captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) - for capture, taxon, score in zip(captures, self.taxa, (0.3, 0.9)): - Classification.objects.filter(detection__source_image=capture).update(taxon=taxon, score=score) - capture.detections.get().occurrence.save() - - task = self.run_task(captures[0].event) + self.run_task(captures[0].event) + occurrence = Occurrence.objects.get(pk=first.occurrence_id) + self.assertEqual(occurrence.determination, self.taxa[1]) + self.assertFalse(Classification.objects.filter(algorithm__key="tracking").exists()) + + # What class masking does: demote the source rows and add a terminal row naming another taxon. + masker = Algorithm.objects.create(name="Test masking", key="test-masking") + for source in Classification.objects.filter(detection__occurrence=occurrence, terminal=True): + source.terminal = False + source.save(update_fields=["terminal"]) + Classification.objects.create( + detection=source.detection, + taxon=self.taxa[2], + score=0.5, + terminal=True, + algorithm=masker, + applied_to=source, + timestamp=source.timestamp, + ) + occurrence = Occurrence.objects.get(pk=occurrence.pk) + update_occurrence_determination(occurrence, save=True) - self.assertEqual(Classification.objects.filter(algorithm=task.algorithm).count(), 1) + occurrence.refresh_from_db() + self.assertEqual(occurrence.determination, self.taxa[2]) class TestTrackingResults(_TrackingCase): @@ -238,16 +235,20 @@ def test_each_linked_occurrence_gets_a_result_with_its_figures(self): by_motion = sorted(results.values(), key=lambda r: r.value) still_result, moving_result = by_motion self.assertEqual(still_result.value, 0.0) - self.assertEqual(moving_result.value, 0.0141) + self.assertEqual(still_result.data["path_length"], 0.0) + # The box centre moves 10 px per step on a 1,000 x 1,000 capture. + self.assertEqual(moving_result.value, 0.0071) + self.assertEqual(moving_result.data["path_length"], 0.0141) for result in by_motion: self.assertEqual( (result.job_id, result.algorithm_id, result.is_current), (job.pk, task.algorithm.pk, True) ) self.assertEqual(result.project_id, self.project.pk) self.assertEqual(result.data["detection_count"], 3) - self.assertEqual(result.data["size_ratio"], 1.0) + self.assertEqual(result.data["size_change"], 1.0) self.assertEqual(result.data["distinct_taxa"], 1) - self.assertEqual(result.data["id_agreement"], 1.0) + self.assertEqual(result.data["label_agreement"], 1.0) + self.assertEqual(len(result.data["link_costs"]), 2) self.assertEqual(result.data["determination_after_id"], self.taxa[0].pk) self.assertEqual(len(result.data["merged_occurrence_ids"]), 2) job.refresh_from_db() @@ -264,11 +265,11 @@ def test_figures_follow_the_labels_and_boxes_of_the_merged_detections(self): result = AlgorithmResult.objects.get(kind="tracking") self.assertEqual(result.data["distinct_taxa"], 2) - self.assertEqual(result.data["size_ratio"], round(300 * 300 / (290 * 300), 4)) - self.assertEqual(result.data["id_agreement"], 0.5) + self.assertEqual(result.data["size_change"], round(300 * 300 / (290 * 300), 4)) + self.assertEqual(result.data["label_agreement"], 0.5) self.assertEqual(result.data["determination_before_id"], self.taxa[0].pk) - def test_the_determination_classification_points_at_the_job_and_the_result(self): + def test_the_result_records_the_determination_before_and_after_without_a_classification(self): captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) for capture, taxon, score in zip(captures, self.taxa, (0.3, 0.9)): Classification.objects.filter(detection__source_image=capture).update(taxon=taxon, score=score) @@ -277,13 +278,46 @@ def test_the_determination_classification_points_at_the_job_and_the_result(self) task = self.run_task(captures[0].event, job=job) - record = Classification.objects.get(algorithm=task.algorithm) result = AlgorithmResult.objects.get(kind="tracking") - self.assertEqual((record.job_id, record.algorithm_result_id), (job.pk, result.pk)) self.assertEqual(result.data["determination_after_id"], self.taxa[1].pk) - self.assertNotEqual(result.data["determination_before_id"], result.data["determination_after_id"]) - # The tracking classification repeats the winner, so it is not counted as a label. - self.assertEqual(result.data["id_agreement"], 0.5) + self.assertEqual(result.data["determination_before_id"], self.taxa[0].pk) + self.assertFalse(Classification.objects.filter(algorithm=task.algorithm).exists()) + + def test_machine_labels_leave_out_post_processing_classifications(self): + """A size filter's terminal row on a frame is not another vote on the taxon.""" + captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) + size_filter = Algorithm.objects.create( + name="Test size filter", key="test-size-filter", task_type="post_processing" + ) + second = captures[1].detections.get() + Classification.objects.create( + detection=second, + taxon=self.taxa[1], + score=1.0, + terminal=True, + algorithm=size_filter, + timestamp=second.timestamp, + ) + + self.run_task(captures[0].event) + + result = AlgorithmResult.objects.get(kind="tracking") + self.assertEqual(result.data["distinct_taxa"], 1) + # The filter's row wins the determination, and no machine label names that taxon. + self.assertEqual(result.data["determination_after_id"], self.taxa[1].pk) + self.assertEqual(result.data["label_agreement"], 0.0) + + def test_each_links_cost_is_recorded_in_chain_order(self): + captures = create_session( + self.deployment, [[BOX], [[110, 100, 210, 200]], [[130, 100, 230, 200]]], self.taxa[0] + ) + + self.run_task(captures[0].event) + + costs = AlgorithmResult.objects.get(kind="tracking").data["link_costs"] + self.assertEqual(len(costs), 2) + self.assertTrue(all(isinstance(cost, float) and cost > 0 for cost in costs)) + self.assertLess(costs[0], costs[1]) def test_results_of_merged_occurrences_move_onto_the_keeper(self): captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) @@ -324,3 +358,108 @@ def test_result_line_says_what_was_tracked(self): self.assertEqual(params["Sessions tracked"], 1) self.assertEqual(params["Detection links created"], 1) self.assertEqual(params["Result"], "Tracked 1 session(s).") + + +class TestTrackingScope(_TrackingCase): + def test_a_capture_set_tracks_every_processed_capture_of_the_sessions_it_touches(self): + captures = create_session(self.deployment, [[BOX], [BOX], [BOX]], self.taxa[0]) + collection = SourceImageCollection.objects.create(name="Sampled", project=self.project) + collection.images.add(captures[0]) + + task = TrackingTask(job=None, logger=logger, source_image_collection_id=collection.pk) + task.run() + + self.assertEqual(self.occurrence_sizes(captures[0].event), [3]) + + def test_a_session_of_another_project_is_not_tracked_for_a_job(self): + captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) + other_project, _ = setup_test_project(reuse=False) + job = Job.objects.create(name="t", project=other_project, job_type_key="post_processing") + job.progress.add_stage("Post-processing", key="post_processing") + job.save() + + self.run_task(captures[0].event, job=job) + + self.assertEqual(self.occurrence_sizes(captures[0].event), [1, 1]) + + +class TestTrackingFailures(_TrackingCase): + def make_job(self) -> Job: + job = Job.objects.create(name="t", project=self.project, job_type_key="post_processing") + job.progress.add_stage("Post-processing", key="post_processing") + job.save() + return job + + def test_a_failing_session_is_counted_and_the_others_are_still_tracked(self): + day = datetime.datetime(2026, 7, 1, 22, 0, 0) + first = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0], start=day) + second = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0], start=day + datetime.timedelta(hours=6)) + failing, working = first[0].event, second[0].event + self.assertNotEqual(failing.pk, working.pk) + job = self.make_job() + real = task_module.assign_occurrences_by_tracking_images + + def fail_for_the_first_session(event, *args, **kwargs): + if event.pk == failing.pk: + # Written before the failure, so the rollback of this session is visible. + Occurrence.objects.filter(event=event).update(determination_score=0.123) + raise ValueError("boom") + return real(event, *args, **kwargs) + + with mock.patch.object(task_module, "assign_occurrences_by_tracking_images", fail_for_the_first_session): + with self.assertRaisesMessage(RuntimeError, "Tracking failed for 1 of 2 session(s)"): + TrackingTask(job=job, logger=logger, event_ids=[failing.pk, working.pk]).run() + + self.assertEqual(self.occurrence_sizes(working), [2]) + self.assertEqual(self.occurrence_sizes(failing), [1, 1]) + self.assertFalse(Occurrence.objects.filter(event=failing, determination_score=0.123).exists()) + working.refresh_from_db() + self.assertEqual(working.occurrences_count, 1) + job.refresh_from_db() + params = {p.name: p.value for p in job.progress.get_stage("post_processing").params} + self.assertEqual((params["Sessions tracked"], params["Sessions failed"]), (1, 1)) + self.assertIn("failed for 1 of 2", params["Result"]) + + def test_progress_is_saved_between_sessions_not_inside_their_transactions(self): + day = datetime.datetime(2026, 7, 1, 22, 0, 0) + first = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0], start=day) + second = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0], start=day + datetime.timedelta(hours=6)) + job = self.make_job() + depths = [] + + def record_depth(progress: float) -> None: + depths.append((round(progress, 2), len(connection.savepoint_ids))) + + task = TrackingTask(job=job, logger=logger, event_ids=[first[0].event_id, second[0].event_id]) + baseline = len(connection.savepoint_ids) + with mock.patch.object(task, "update_progress", record_depth): + task.run() + + self.assertEqual(depths, [(0.5, baseline), (1.0, baseline), (1.0, baseline)]) + + +class TestTrackingQueries(_TrackingCase): + def count_queries(self, boxes_per_capture: list[list[list[int]] | None], start: datetime.datetime) -> int: + from cachalot.api import cachalot_disabled + + captures = create_session(self.deployment, boxes_per_capture, self.taxa[0], start=start) + disabled = cachalot_disabled() + disabled.__enter__() + try: + with CaptureQueriesContext(connection) as queries: + self.run_task(captures[0].event) + finally: + # cachalot_disabled() does not restore itself when the block raises. + disabled.__exit__(None, None, None) + self.assertEqual(self.occurrence_sizes(captures[0].event), [len(boxes_per_capture)]) + return len(queries) + + def test_queries_grow_by_a_constant_per_capture_not_per_detection(self): + """Each extra capture adds a fixed number of queries; reassigning a chain's detections is one update.""" + self.count_queries([[BOX]] * 3, datetime.datetime(2026, 6, 30)) # warms per-process caches + counts = [ + self.count_queries([[BOX]] * n, datetime.datetime(2026, 7, day)) for day, n in ((1, 3), (2, 6), (3, 9)) + ] + self.assertEqual(counts[2] - counts[1], counts[1] - counts[0]) + per_capture = (counts[1] - counts[0]) // 3 + self.assertLessEqual(per_capture, 4) diff --git a/ami/ml/post_processing/tracking/config.py b/ami/ml/post_processing/tracking/config.py index 3adb8b74f..fc672f445 100644 --- a/ami/ml/post_processing/tracking/config.py +++ b/ami/ml/post_processing/tracking/config.py @@ -25,7 +25,15 @@ class TrackingConfig(pydantic.BaseModel): shown on the admin form. """ - source_image_collection_id: int | None = None + source_image_collection_id: int | None = pydantic.Field( + None, + title="Capture set", + description=( + "Track the sessions that contain captures from this set. Every processed capture of those " + "sessions is tracked, not only the captures in the set, because a chain needs the captures " + "between its detections." + ), + ) event_ids: list[int] = [] cost_threshold: float = pydantic.Field( diff --git a/ami/ml/post_processing/tracking/stats.py b/ami/ml/post_processing/tracking/stats.py index 41db06584..2461cc23b 100644 --- a/ami/ml/post_processing/tracking/stats.py +++ b/ami/ml/post_processing/tracking/stats.py @@ -18,9 +18,10 @@ class OccurrenceFigures: detection_count: int motion: float - size_ratio: float + path_length: float + size_change: float distinct_taxa: int - id_agreement: float | None + label_agreement: float | None def frame_diagonal(sizes: Sequence[tuple[int | None, int | None]], boxes: Sequence[BBox]) -> float: @@ -38,15 +39,29 @@ def frame_diagonal(sizes: Sequence[tuple[int | None, int | None]], boxes: Sequen return math.hypot(far_x, far_y) or 1.0 -def motion(boxes: Sequence[BBox], diagonal: float) -> float: - """Path length between consecutive box centres, as a fraction of ``diagonal``.""" +def _raw_path_length(boxes: Sequence[BBox]) -> float: centres = [((box[0] + box[2]) / 2, (box[1] + box[3]) / 2) for box in boxes] - path = sum(math.dist(a, b) for a, b in zip(centres, centres[1:])) - return round(path / diagonal, _ROUND_TO) + return sum(math.dist(a, b) for a, b in zip(centres, centres[1:])) + + +def path_length(boxes: Sequence[BBox], diagonal: float) -> float: + """Total distance between consecutive box centres, as a fraction of ``diagonal``.""" + return round(_raw_path_length(boxes) / diagonal, _ROUND_TO) + + +def motion(boxes: Sequence[BBox], diagonal: float) -> float: + """Mean distance per step between consecutive box centres, as a fraction of ``diagonal``; 0 for one box.""" + steps = len(boxes) - 1 + if steps < 1: + return 0.0 + return round(_raw_path_length(boxes) / steps / diagonal, _ROUND_TO) + +def size_change(boxes: Sequence[BBox]) -> float: + """Largest box area over the smallest (at least 1); areas are floored at 1 so a degenerate box cannot divide by 0. -def size_ratio(boxes: Sequence[BBox]) -> float: - """Largest box area over the smallest, with areas floored at 1 so a degenerate box cannot divide by zero.""" + The config's ``min_size_ratio`` is the opposite way round (smaller over larger, at most 1). + """ areas = [max(abs((box[2] - box[0]) * (box[3] - box[1])), _MIN_AREA) for box in boxes] if not areas: return 1.0 @@ -58,8 +73,11 @@ def distinct_taxa(labels: Sequence[int | None]) -> int: return len({label for label in labels if label is not None}) -def id_agreement(labels: Sequence[int | None], determination_id: int | None) -> float | None: - """The share of labels naming the determination, or None when there are no labels.""" +def label_agreement(labels: Sequence[int | None], determination_id: int | None) -> float | None: + """The share of labels naming the determination, or None when there are no labels. + + Labels are machine classifications, not human identifications. + """ if not labels: return None return round(sum(label == determination_id and label is not None for label in labels) / len(labels), _ROUND_TO) @@ -72,10 +90,12 @@ def occurrence_figures( determination_id: int | None, ) -> OccurrenceFigures: """All the figures for one occurrence: ``sizes`` holds the width and height of each detection's capture.""" + diagonal = frame_diagonal(sizes, boxes) return OccurrenceFigures( detection_count=len(boxes), - motion=motion(boxes, frame_diagonal(sizes, boxes)), - size_ratio=size_ratio(boxes), + motion=motion(boxes, diagonal), + path_length=path_length(boxes, diagonal), + size_change=size_change(boxes), distinct_taxa=distinct_taxa(labels), - id_agreement=id_agreement(labels, determination_id), + label_agreement=label_agreement(labels, determination_id), ) diff --git a/ami/ml/post_processing/tracking/task.py b/ami/ml/post_processing/tracking/task.py index 222b864bb..714c93a63 100644 --- a/ami/ml/post_processing/tracking/task.py +++ b/ami/ml/post_processing/tracking/task.py @@ -12,7 +12,6 @@ from django.db import transaction from django.db.models import Count, Exists, OuterRef -from django.utils import timezone from ami.main.models import ( Classification, @@ -26,6 +25,7 @@ update_occurrence_determination, ) from ami.ml.models import Algorithm, AlgorithmResult +from ami.ml.models.algorithm import AlgorithmTaskType from ami.ml.post_processing.base import BasePostProcessingTask from .config import TrackingConfig @@ -45,6 +45,7 @@ class LinkedOccurrence: detections: list[Detection] determination_before_id: int | None merged_ids: list[int] + link_costs: list[float] def event_is_fresh(event: Event) -> tuple[bool, str]: @@ -83,45 +84,15 @@ def processed_captures(event: Event) -> list[SourceImage]: ) -def record_tracking_determination( - occurrence: Occurrence, - algorithm: Algorithm, - job: "Job | None" = None, - result: AlgorithmResult | None = None, -) -> Classification | None: - """Leave a terminal classification by the tracking algorithm after a merge changed the determination. - - The row carries the winning prediction and points back at it through ``applied_to``, the way - class masking does, so the history shows what tracking decided. Nothing is written when the - winner already came from the tracking algorithm. The row points at the run's ``job`` and at the - occurrence's tracking ``result`` so the history shows it under that result. - """ - winner = occurrence.best_prediction - if winner is None or winner.detection_id is None or winner.taxon_id is None: - return None - if winner.algorithm_id == algorithm.pk: - return None - return Classification.objects.create( - detection=winner.detection, - taxon=winner.taxon, - score=winner.score, - terminal=True, - algorithm=algorithm, - timestamp=timezone.now(), - applied_to=winner, - job=job, - algorithm_result=result, - ) - - def record_tracking_results( linked: Sequence[LinkedOccurrence], algorithm: Algorithm, job: "Job | None" ) -> dict[int, AlgorithmResult]: """Write one tracking result per linked occurrence, keyed by occurrence id. The figures are computed in memory from the chains' detections plus one query for their terminal - classifications. Classifications by the tracking algorithm itself are left out, since they repeat - the winner and would count as an extra vote. Call after the determinations have settled. + classifications. Only machine labels count: classifications by post-processing algorithms (size filter, + class masking) are left out, and a classification with no algorithm counts as a source label. + Call after the determinations have settled. """ if not linked: return {} @@ -129,7 +100,7 @@ def record_tracking_results( labels: dict[int, list[int | None]] = collections.defaultdict(list) for detection_id, taxon_id in ( Classification.objects.filter(detection_id__in=detection_ids, terminal=True) - .exclude(algorithm=algorithm) + .exclude(algorithm__task_type=AlgorithmTaskType.POST_PROCESSING.value) .values_list("detection_id", "taxon_id") ): labels[detection_id].append(taxon_id) @@ -151,6 +122,7 @@ def record_tracking_results( value=figures.motion, data={ **dataclasses.asdict(figures), + "link_costs": item.link_costs, "determination_before_id": item.determination_before_id, "determination_after_id": item.keeper.determination_id, "merged_occurrence_ids": item.merged_ids, @@ -165,6 +137,7 @@ def assign_occurrences_from_detection_chains( logger: logging.Logger, record_as: Algorithm | None = None, job: "Job | None" = None, + link_costs: dict[int, float] | None = None, ) -> dict[str, int]: """Fold each chain of linked detections into one occurrence, keeping the first existing one. @@ -172,8 +145,9 @@ def assign_occurrences_from_detection_chains( session boundary starts a new chain on each side. Identifications move onto the keeper before the occurrences that held them are deleted, because deleting an occurrence deletes its identifications. Results of the absorbed occurrences move onto the keeper first. With ``record_as`` set, every keeper - built from two or more detections or from a merge gets a tracking result, and a merge that changes - its determination leaves a classification by that algorithm, both attributed to ``job``. + built from two or more detections or from a merge gets a tracking result attributed to ``job``, which + records the determination before and after, so no classification is written for it. ``link_costs`` maps a + detection to the cost of its link to the next one, for the links this run made. """ image_ids = [image.pk for image in source_images] detections = list( @@ -186,7 +160,7 @@ def assign_occurrences_from_detection_chains( has_previous = {det.next_detection_id for det in detections if det.next_detection_id in by_id} visited: set[int] = set() - created = merged = identifications_moved = determinations_recorded = 0 + created = merged = identifications_moved = 0 linked: dict[int, LinkedOccurrence] = {} deleted: set[int] = set() existing = Occurrence.objects.filter(detections__source_image_id__in=image_ids).distinct().count() @@ -215,10 +189,11 @@ def assign_occurrences_from_detection_chains( ) created += 1 - for d in chain: - if d.occurrence_id != keeper.pk: + moving = [d for d in chain if d.occurrence_id != keeper.pk] + if moving: + Detection.objects.filter(pk__in=[d.pk for d in moving]).update(occurrence=keeper) + for d in moving: d.occurrence = keeper - d.save(update_fields=["occurrence"]) doomed = old_occ_ids - {keeper.pk} if doomed: @@ -233,23 +208,20 @@ def assign_occurrences_from_detection_chains( if update_occurrence_determination(keeper, save=False): keeper.save(update_determination=False, update_fields=["determination", "determination_score"]) if record_as is not None and (len(chain) > 1 or doomed): - linked[keeper.pk] = LinkedOccurrence(keeper, chain, previous_determination_id, sorted(doomed)) + costs = [round(link_costs[d.pk], 4) for d in chain[:-1] if d.pk in link_costs] + linked[keeper.pk] = LinkedOccurrence(keeper, chain, previous_determination_id, sorted(doomed), costs) results: dict[int, AlgorithmResult] = {} if record_as is not None: # A keeper that a later chain absorbed no longer exists. final = [item for pk, item in linked.items() if pk not in deleted] results = record_tracking_results(final, record_as, job) - for item in final: - if item.keeper.determination_id != item.determination_before_id: - if record_tracking_determination(item.keeper, record_as, job, results.get(item.keeper.pk)) is not None: - determinations_recorded += 1 new_count = Occurrence.objects.filter(detections__source_image_id__in=image_ids).distinct().count() logger.info( f"Created {created} occurrences and merged {merged} across {len(image_ids)} captures " f"(occurrences before: {existing}, after: {new_count}). Moved {identifications_moved} identification(s), " - f"recorded {len(results)} result(s) and {determinations_recorded} determination change(s)." + f"recorded {len(results)} result(s)." ) return { "occurrences_before": existing, @@ -257,7 +229,6 @@ def assign_occurrences_from_detection_chains( "occurrences_created": created, "occurrences_merged": merged, "identifications_moved": identifications_moved, - "determinations_recorded": determinations_recorded, "results_recorded": len(results), } @@ -307,33 +278,37 @@ def assign_occurrences_by_tracking_images( event: Event, logger: logging.Logger, config: TrackingConfig, - progress_cb: typing.Callable[[float], None] | None = None, record_as: Algorithm | None = None, job: "Job | None" = None, ) -> dict[str, int]: - """Link the detections of one session's processed captures and fold the chains into occurrences.""" + """Link the detections of one session's processed captures and fold the chains into occurrences. + + Runs in one transaction, so a failure leaves the session as it was. The caller reports progress + after it returns, since saving the job inside the transaction would keep its row locked. + """ source_images = processed_captures(event) if len(source_images) < 2: logger.warning(f"Session {event.pk}: fewer than two processed captures ({len(source_images)}).") return {} - transitions = len(source_images) - 1 links = skipped_for_dimensions = 0 + costs: dict[int, float] = {} skipped_for_interval = sum( captures_too_far_apart(source_images[i].timestamp, source_images[i + 1].timestamp, config) - for i in range(transitions) + for i in range(len(source_images) - 1) ) # Per-session atomic boundary: a crash mid-session rolls back this session only. with transaction.atomic(): - for i, proposed in enumerate(iter_transition_links(source_images, config, logger)): + for proposed in iter_transition_links(source_images, config, logger): if proposed is None: skipped_for_dimensions += 1 else: save_links(proposed, logger) links += len(proposed) - if progress_cb: - progress_cb((i + 1) / transitions) - counters = assign_occurrences_from_detection_chains(source_images, logger, record_as=record_as, job=job) + costs.update({det.pk: cost for det, _, cost in proposed}) + counters = assign_occurrences_from_detection_chains( + source_images, logger, record_as=record_as, job=job, link_costs=costs + ) counters["links_created"] = links counters["transitions_too_far_apart"] = skipped_for_interval @@ -408,7 +383,34 @@ def _skip_reason(self, event: Event) -> str | None: return "it has human identifications" return None + def _track_session(self, event: Event) -> tuple[dict[str, int] | None, str | None]: + """Track one session in its own transaction, returning ``(counters, None)`` or ``(None, skip reason)``. + + The checks and the writes share one lock on the session, so an edit made since the job + started is seen and an edit made during the run waits. Nothing here saves the job, so its row + is not locked while the session is processed. + """ + with transaction.atomic(): + lock_sessions([event.pk]) + reason = self._skip_reason(event) + if reason is not None: + return None, reason + counters = assign_occurrences_by_tracking_images( + event=event, logger=self.logger, config=self.config, record_as=self.algorithm, job=self.job + ) + if not counters: + return None, "it has fewer than two processed captures" + return counters, None + def run(self) -> None: + """Track every session in scope, one transaction per session. + + A capture-set scope tracks every processed capture of the sessions the set touches, not only the + captures in the set, because a chain needs the captures between its detections. A session that + fails is rolled back, logged and counted, and the run goes on with the next one. After the counts + of the sessions that were tracked are refreshed and the metrics reported, the run raises if any + session failed, so the job is marked failed. + """ self.logger.info(f"Tracking starting with config: {self.config.dict()}") events = self._resolve_events() @@ -417,40 +419,28 @@ def run(self) -> None: totals: collections.Counter[str] = collections.Counter() tracked_event_ids: list[int] = [] + failed_event_ids: list[int] = [] # Why each session was skipped, so a run that tracks nothing can say so. skip_reasons: collections.Counter[str] = collections.Counter() for idx, event in enumerate(events, start=1): self.logger.info(f"Tracking session {idx}/{total} (id={event.pk})") - # The checks and the writes share one lock on the session, so an edit made since the - # job started is seen and an edit made during the run waits. - with transaction.atomic(): - lock_sessions([event.pk]) - reason = self._skip_reason(event) - if reason is not None: - totals["skipped"] += 1 - skip_reasons[reason] += 1 - continue - - def _stage_progress(p: float, _idx=idx, _total=total) -> None: - self.update_progress(((_idx - 1) + p) / _total) - - counters = assign_occurrences_by_tracking_images( - event=event, - logger=self.logger, - config=self.config, - record_as=self.algorithm, - job=self.job, - progress_cb=_stage_progress, - ) - if not counters: - totals["skipped"] += 1 - skip_reasons["it has fewer than two processed captures"] += 1 - continue + try: + counters, reason = self._track_session(event) + except Exception: + self.logger.exception(f"Tracking failed for session {event.pk}; its changes were rolled back.") + failed_event_ids.append(event.pk) + counters, reason = None, None + if counters is not None: totals["tracked"] += 1 tracked_event_ids.append(event.pk) for key in ("links_created", "occurrences_merged", "results_recorded", "transitions_too_far_apart"): totals[key] += counters.get(key, 0) + elif reason is not None: + totals["skipped"] += 1 + skip_reasons[reason] += 1 + # Saved between sessions, after the transaction has committed, so the bar moves per session. + self.update_progress(idx / total) # Merging occurrences changes the session and station counts, which no save refreshes. # This already runs in a background job, so the station refresh stays inline. @@ -459,6 +449,7 @@ def _stage_progress(p: float, _idx=idx, _total=total) -> None: metrics: dict[str, typing.Any] = { "Sessions tracked": totals["tracked"], "Sessions skipped": totals["skipped"], + "Sessions failed": len(failed_event_ids), "Detection links created": totals["links_created"], "Occurrences merged": totals["occurrences_merged"], "Occurrences recorded": totals["results_recorded"], @@ -467,7 +458,12 @@ def _stage_progress(p: float, _idx=idx, _total=total) -> None: metrics["Capture pairs too far apart to compare"] = totals["transitions_too_far_apart"] # The job still succeeds when every session is skipped, so this line is written on every run: # a retry keeps text params, and a stale line would contradict the counts. - if totals["tracked"]: + if failed_event_ids: + metrics["Result"] = ( + f"Tracking failed for {len(failed_event_ids)} of {total} session(s) " + f"(ids {failed_event_ids}); {totals['tracked']} were tracked." + ) + elif totals["tracked"]: metrics["Result"] = f"Tracked {totals['tracked']} session(s)." elif skip_reasons: metrics["Result"] = nothing_tracked_summary(skip_reasons) @@ -477,3 +473,5 @@ def _stage_progress(p: float, _idx=idx, _total=total) -> None: self.report_stage_metrics(metrics) self.update_progress(1.0) self.logger.info(f"Tracking finished: {dict(totals)}") + if failed_event_ids: + raise RuntimeError(metrics["Result"]) diff --git a/ami/ml/results/schemas.py b/ami/ml/results/schemas.py index 7bc4211df..005c2ea28 100644 --- a/ami/ml/results/schemas.py +++ b/ami/ml/results/schemas.py @@ -74,14 +74,19 @@ class TrackingResultData(DeterminationSnapshot): kind: ClassVar[str] = "tracking" detection_count: int - # Path length between consecutive detection centres, as a fraction of the image diagonal; 0 for a still insect. + # Mean distance per step between consecutive detection centres, as a fraction of the image diagonal; + # 0 for one detection. motion: float - # The largest box area over the smallest, with areas floored at 1; 1 when the box never changed size. - size_ratio: float - # Distinct taxa among the terminal classifications of the occurrence's detections, from every algorithm. + # The same distances added up over the whole path, as a fraction of the image diagonal. + path_length: float + # The largest box area over the smallest (at least 1), with areas floored at 1; 1 when the box never changed size. + size_change: float + # Distinct taxa among the machine classifications of the occurrence's detections, leaving out post-processing ones. distinct_taxa: int # The share of those classifications naming the determination after the run; None when there are none. - id_agreement: float | None = None + label_agreement: float | None = None + # The matching cost of each link this run made in the chain, rounded to 4 places, in chain order. + link_costs: list[float] = [] # Occurrences the run folded into this one. They are deleted, so these are plain ids, not references. merged_occurrence_ids: list[int] = [] From c1cb15da5b013ddd4cfdd31cf186456468457336 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Mon, 5 Oct 2026 22:51:27 -0700 Subject: [PATCH 15/28] feat(ui): show the reworked tracking figures on an occurrence's history The tracking card reads the renamed result fields: the mean movement per step beside the total path length, the size change, and the label agreement, which counts machine labels only. The labels are "Movement per step", "Path length", "Size change", "Taxa" and "Label agreement". Co-Authored-By: Claude Sonnet 5.5 Claude-Session: https://claude.ai/code/session_01QT59KePky4u4nbCsTAggtc --- .../models/occurrence-history.test.ts | 8 +++++--- ui/src/data-services/models/occurrence-history.ts | 14 +++++++++----- .../identification-card/algorithm-result.tsx | 14 ++++++++++---- ui/src/utils/language.ts | 8 +++++--- 4 files changed, 29 insertions(+), 15 deletions(-) diff --git a/ui/src/data-services/models/occurrence-history.test.ts b/ui/src/data-services/models/occurrence-history.test.ts index 6e9defaf4..7cd07780a 100644 --- a/ui/src/data-services/models/occurrence-history.test.ts +++ b/ui/src/data-services/models/occurrence-history.test.ts @@ -121,10 +121,12 @@ const tracking: ServerOccurrenceHistoryEntry = { determination_before_id: 3, distinct_taxa: 1, extra: {}, - id_agreement: null, + label_agreement: null, + link_costs: [0.12, 0.2], merged_occurrence_ids: [11, 12], - motion: 0.0141, - size_ratio: 1.15, + motion: 0.0071, + path_length: 0.0141, + size_change: 1.15, }, data_references: {}, determination_after: NOCTUA, diff --git a/ui/src/data-services/models/occurrence-history.ts b/ui/src/data-services/models/occurrence-history.ts index 3f603e2a8..1da8236ee 100644 --- a/ui/src/data-services/models/occurrence-history.ts +++ b/ui/src/data-services/models/occurrence-history.ts @@ -82,17 +82,21 @@ export interface SizeFilterResultData extends ServerDeterminationSnapshot { } export interface TrackingResultData extends ServerDeterminationSnapshot { - /** The share of those classifications naming the determination after the run; null when there are none. */ - id_agreement: number | null detection_count: number - /** Distinct taxa among the terminal classifications of the occurrence's detections. */ + /** Distinct taxa among the machine classifications of the occurrence's detections. */ distinct_taxa: number + /** The share of those classifications naming the determination after the run; null when there are none. */ + label_agreement: number | null + /** The matching cost of each link the run made, in chain order. */ + link_costs: number[] /** Occurrences the run folded into this one; they no longer exist. */ merged_occurrence_ids: number[] - /** The path length between detection centres as a fraction of the image diagonal. */ + /** The mean distance per step between detection centres as a fraction of the image diagonal. */ motion: number + /** The total distance between detection centres as a fraction of the image diagonal. */ + path_length: number /** The largest box area over the smallest. */ - size_ratio: number + size_change: number } export interface ServerIdentificationDetails { diff --git a/ui/src/pages/occurrence-details/identification-card/algorithm-result.tsx b/ui/src/pages/occurrence-details/identification-card/algorithm-result.tsx index a47355bc4..d40ea91fa 100644 --- a/ui/src/pages/occurrence-details/identification-card/algorithm-result.tsx +++ b/ui/src/pages/occurrence-details/identification-card/algorithm-result.tsx @@ -167,10 +167,16 @@ export const AlgorithmResult = ({ distance: formatPercent(data.motion), }), }, + { + label: translate(STRING.HISTORY_TRACKING_PATH_LENGTH), + value: translate(STRING.HISTORY_TRACKING_MOVEMENT_VALUE, { + distance: formatPercent(data.path_length), + }), + }, { label: translate(STRING.HISTORY_TRACKING_SIZE_CHANGE), value: translate(STRING.HISTORY_TRACKING_SIZE_CHANGE_VALUE, { - ratio: `${Math.round(data.size_ratio * 100) / 100}`, + ratio: `${Math.round(data.size_change * 100) / 100}`, }), }, { @@ -178,10 +184,10 @@ export const AlgorithmResult = ({ value: data.distinct_taxa, }, { - label: translate(STRING.HISTORY_TRACKING_AGREEMENT), + label: translate(STRING.HISTORY_TRACKING_LABEL_AGREEMENT), value: - data.id_agreement !== null - ? formatPercent(data.id_agreement) + data.label_agreement !== null + ? formatPercent(data.label_agreement) : translate(STRING.VALUE_NOT_AVAILABLE), }, { diff --git a/ui/src/utils/language.ts b/ui/src/utils/language.ts index 3a6d03e16..88d2e1d69 100644 --- a/ui/src/utils/language.ts +++ b/ui/src/utils/language.ts @@ -338,11 +338,12 @@ export enum STRING { HISTORY_SIZE_FILTER, HISTORY_SUPERSEDED_BY, HISTORY_TRACKING, - HISTORY_TRACKING_AGREEMENT, + HISTORY_TRACKING_LABEL_AGREEMENT, HISTORY_TRACKING_DETECTIONS, HISTORY_TRACKING_MERGED, HISTORY_TRACKING_MOVEMENT, HISTORY_TRACKING_MOVEMENT_VALUE, + HISTORY_TRACKING_PATH_LENGTH, HISTORY_TRACKING_SIZE_CHANGE, HISTORY_TRACKING_SIZE_CHANGE_VALUE, HISTORY_TRACKING_TAXA, @@ -804,12 +805,13 @@ const ENGLISH_STRINGS: { [key in STRING]: string } = { [STRING.HISTORY_SIZE_FILTER]: 'Size filter', [STRING.HISTORY_SUPERSEDED_BY]: 'Superseded by {{name}}', [STRING.HISTORY_TRACKING]: 'Occurrence tracking', - [STRING.HISTORY_TRACKING_AGREEMENT]: 'Agreement', + [STRING.HISTORY_TRACKING_LABEL_AGREEMENT]: 'Label agreement', [STRING.HISTORY_TRACKING_DETECTIONS]: 'Detections', [STRING.HISTORY_TRACKING_MERGED]: 'Occurrences merged in', - [STRING.HISTORY_TRACKING_MOVEMENT]: 'Movement', + [STRING.HISTORY_TRACKING_MOVEMENT]: 'Movement per step', [STRING.HISTORY_TRACKING_MOVEMENT_VALUE]: '{{distance}} of the image diagonal', + [STRING.HISTORY_TRACKING_PATH_LENGTH]: 'Path length', [STRING.HISTORY_TRACKING_SIZE_CHANGE]: 'Size change', [STRING.HISTORY_TRACKING_SIZE_CHANGE_VALUE]: '{{ratio}}× from smallest to largest', From 78437c1243a26891c3859b83b16c45db00f93b2d Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Mon, 5 Oct 2026 23:06:41 -0700 Subject: [PATCH 16/28] fix(tracking): save job progress while a session is being matched, and write the session in one short transaction The stale-job reaper revokes a job whose updated_at stops moving, and a busy session can take minutes. A progress write inside the session transaction is invisible until commit and keeps the job row locked, so a run now works in two phases per session. Matching reads the session's detections and proposes every link outside any transaction, saving progress every few transitions or seconds. The write phase is then a short transaction that locks the session, repeats the guards, refuses to write when the detections' links or occurrences changed since matching (the session is skipped with a logged reason), and saves the links and folds the chains without touching the job. Tests cover progress saved between transitions of one session outside any transaction, and a session that changes while it is matched. Co-Authored-By: Claude Sonnet 5.5 Claude-Session: https://claude.ai/code/session_01QT59KePky4u4nbCsTAggtc --- .../tests/test_tracking_task.py | 59 ++++++- ami/ml/post_processing/tracking/task.py | 163 +++++++++++++----- 2 files changed, 174 insertions(+), 48 deletions(-) diff --git a/ami/ml/post_processing/tests/test_tracking_task.py b/ami/ml/post_processing/tests/test_tracking_task.py index abdad1050..97ceabfe5 100644 --- a/ami/ml/post_processing/tests/test_tracking_task.py +++ b/ami/ml/post_processing/tests/test_tracking_task.py @@ -397,16 +397,16 @@ def test_a_failing_session_is_counted_and_the_others_are_still_tracked(self): failing, working = first[0].event, second[0].event self.assertNotEqual(failing.pk, working.pk) job = self.make_job() - real = task_module.assign_occurrences_by_tracking_images + real = task_module.write_session_plan - def fail_for_the_first_session(event, *args, **kwargs): - if event.pk == failing.pk: + def fail_for_the_first_session(plan, *args, **kwargs): + if plan.source_images[0].event_id == failing.pk: # Written before the failure, so the rollback of this session is visible. - Occurrence.objects.filter(event=event).update(determination_score=0.123) + Occurrence.objects.filter(event=failing).update(determination_score=0.123) raise ValueError("boom") - return real(event, *args, **kwargs) + return real(plan, *args, **kwargs) - with mock.patch.object(task_module, "assign_occurrences_by_tracking_images", fail_for_the_first_session): + with mock.patch.object(task_module, "write_session_plan", fail_for_the_first_session): with self.assertRaisesMessage(RuntimeError, "Tracking failed for 1 of 2 session(s)"): TrackingTask(job=job, logger=logger, event_ids=[failing.pk, working.pk]).run() @@ -438,6 +438,53 @@ def record_depth(progress: float) -> None: self.assertEqual(depths, [(0.5, baseline), (1.0, baseline), (1.0, baseline)]) +class TestTrackingProgressDuringMatching(_TrackingCase): + def make_job(self) -> Job: + job = Job.objects.create(name="t", project=self.project, job_type_key="post_processing") + job.progress.add_stage("Post-processing", key="post_processing") + job.save() + return job + + def test_progress_is_saved_between_transitions_of_one_session_outside_any_transaction(self): + """The stale-job reaper watches ``updated_at``, so a long session must keep saving the job while it matches.""" + captures = create_session(self.deployment, [[BOX]] * 7, self.taxa[0]) + task = TrackingTask(job=self.make_job(), logger=logger, event_ids=[captures[0].event_id]) + baseline = len(connection.savepoint_ids) + seen = [] + + def record(progress: float) -> None: + seen.append((round(progress, 2), len(connection.savepoint_ids))) + + with mock.patch.object(task_module, "PROGRESS_EVERY_TRANSITIONS", 2), mock.patch.object( + task, "update_progress", record + ): + task.run() + + # Six transitions, saved after the 2nd, 4th and 6th, then once per session and once at the end. + self.assertEqual([p for p, _ in seen], [0.33, 0.67, 1.0, 1.0, 1.0]) + self.assertEqual({depth for _, depth in seen}, {baseline}) + + def test_a_session_that_changes_while_it_is_matched_is_skipped_not_written(self): + captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) + event = captures[0].event + job = self.make_job() + real = task_module.plan_session_links + + def plan_then_change(*args, **kwargs): + plan = real(*args, **kwargs) + add_detection(captures[1], [500, 500, 560, 560], self.taxa[0]) + return plan + + with mock.patch.object(task_module, "plan_session_links", plan_then_change): + self.run_task(event, job=job) + + self.assertEqual(self.occurrence_sizes(event), [1, 1, 1]) + self.assertFalse(Detection.objects.filter(source_image__event=event, next_detection__isnull=False).exists()) + job.refresh_from_db() + params = {p.name: p.value for p in job.progress.get_stage("post_processing").params} + self.assertIn("it changed while it was being tracked", params["Result"]) + + class TestTrackingQueries(_TrackingCase): def count_queries(self, boxes_per_capture: list[list[list[int]] | None], start: datetime.datetime) -> int: from cachalot.api import cachalot_disabled diff --git a/ami/ml/post_processing/tracking/task.py b/ami/ml/post_processing/tracking/task.py index 714c93a63..2bfe07ce7 100644 --- a/ami/ml/post_processing/tracking/task.py +++ b/ami/ml/post_processing/tracking/task.py @@ -7,6 +7,7 @@ import collections import dataclasses import logging +import time import typing from collections.abc import Iterable, Iterator, Sequence @@ -37,6 +38,11 @@ from ami.jobs.models import Job +# Progress is saved after this many matched transitions, or after this many seconds, whichever comes first. +PROGRESS_EVERY_TRANSITIONS = 25 +PROGRESS_EVERY_SECONDS = 5.0 + + @dataclasses.dataclass class LinkedOccurrence: """An occurrence the run built from a chain of two or more detections, or by merging occurrences.""" @@ -245,16 +251,32 @@ def save_links(links: Iterable[tuple[Detection, Detection, float]], logger: logg Detection.objects.bulk_update([det for det, _, _ in links], ["next_detection"]) +@dataclasses.dataclass +class SessionPlan: + """The links proposed for one session, computed without writing anything. + + ``snapshot`` records each detection's ``next_detection`` and occurrence as they were read, so the + write phase can tell whether the session changed while the links were being matched. + """ + + source_images: list[SourceImage] + snapshot: dict[int, tuple[int | None, int | None]] + proposals: list[list[tuple[Detection, Detection, float]] | None] + transitions_too_far_apart: int + + def iter_transition_links( - source_images: Sequence[SourceImage], config: TrackingConfig, logger: logging.Logger + source_images: Sequence[SourceImage], + detections_by_capture: dict[int, dict[int, Detection]], + config: TrackingConfig, + logger: logging.Logger, ) -> Iterator[list[tuple[Detection, Detection, float]] | None]: """Yield the proposed links for each pair of consecutive captures, in order, saving nothing. Yields None for a transition that is not compared (the earlier capture has no dimensions) and an empty list for one over the interval limit. """ - transitions = len(source_images) - 1 - for i in range(transitions): + for i in range(len(source_images) - 1): cur, nxt = source_images[i], source_images[i + 1] if captures_too_far_apart(cur.timestamp, nxt.timestamp, config): yield [] @@ -263,8 +285,8 @@ def iter_transition_links( logger.warning(f"Capture {cur.pk} has no dimensions; not comparing it with the next capture.") yield None continue - current = {det.pk: det for det in cur.detections.valid()} - following = {det.pk: det for det in nxt.detections.valid()} + current = detections_by_capture.get(cur.pk, {}) + following = detections_by_capture.get(nxt.pk, {}) links = select_links( [(pk, det.bbox) for pk, det in current.items()], [(pk, det.bbox) for pk, det in following.items()], @@ -274,44 +296,78 @@ def iter_transition_links( yield [(current[from_id], following[to_id], cost) for from_id, to_id, cost in links] -def assign_occurrences_by_tracking_images( +def plan_session_links( event: Event, logger: logging.Logger, config: TrackingConfig, - record_as: Algorithm | None = None, - job: "Job | None" = None, -) -> dict[str, int]: - """Link the detections of one session's processed captures and fold the chains into occurrences. + progress_cb: typing.Callable[[float], None] | None = None, +) -> SessionPlan | None: + """Match the detections of one session's processed captures, writing nothing and holding no lock. - Runs in one transaction, so a failure leaves the session as it was. The caller reports progress - after it returns, since saving the job inside the transaction would keep its row locked. + Returns None when the session has fewer than two processed captures. ``progress_cb`` receives the + share of transitions matched after each one. """ source_images = processed_captures(event) if len(source_images) < 2: logger.warning(f"Session {event.pk}: fewer than two processed captures ({len(source_images)}).") - return {} + return None + detections_by_capture: dict[int, dict[int, Detection]] = collections.defaultdict(dict) + for det in Detection.objects.valid().filter(source_image_id__in=[image.pk for image in source_images]): + detections_by_capture[det.source_image_id][det.pk] = det + snapshot = { + det.pk: (det.next_detection_id, det.occurrence_id) + for detections in detections_by_capture.values() + for det in detections.values() + } + transitions = len(source_images) - 1 + proposals = [] + for i, proposed in enumerate(iter_transition_links(source_images, detections_by_capture, config, logger)): + proposals.append(proposed) + if progress_cb: + progress_cb((i + 1) / transitions) + too_far = sum( + captures_too_far_apart(source_images[i].timestamp, source_images[i + 1].timestamp, config) + for i in range(transitions) + ) + return SessionPlan(source_images, snapshot, proposals, too_far) + +def plan_is_current(plan: SessionPlan) -> bool: + """Whether the session's detections still have the links and occurrences the plan was matched against.""" + current = { + pk: (next_id, occurrence_id) + for pk, next_id, occurrence_id in Detection.objects.valid() + .filter(source_image_id__in=[image.pk for image in plan.source_images]) + .values_list("pk", "next_detection_id", "occurrence_id") + } + return current == plan.snapshot + + +def write_session_plan( + plan: SessionPlan, + logger: logging.Logger, + record_as: Algorithm | None = None, + job: "Job | None" = None, +) -> dict[str, int]: + """Save a plan's links and fold the chains into occurrences. + + Call inside the session's transaction, after checking ``plan_is_current``. It does no progress + writes, since saving the job inside the transaction would keep its row locked. + """ links = skipped_for_dimensions = 0 costs: dict[int, float] = {} - skipped_for_interval = sum( - captures_too_far_apart(source_images[i].timestamp, source_images[i + 1].timestamp, config) - for i in range(len(source_images) - 1) + for proposed in plan.proposals: + if proposed is None: + skipped_for_dimensions += 1 + continue + save_links(proposed, logger) + links += len(proposed) + costs.update({det.pk: cost for det, _, cost in proposed}) + counters = assign_occurrences_from_detection_chains( + plan.source_images, logger, record_as=record_as, job=job, link_costs=costs ) - # Per-session atomic boundary: a crash mid-session rolls back this session only. - with transaction.atomic(): - for proposed in iter_transition_links(source_images, config, logger): - if proposed is None: - skipped_for_dimensions += 1 - else: - save_links(proposed, logger) - links += len(proposed) - costs.update({det.pk: cost for det, _, cost in proposed}) - counters = assign_occurrences_from_detection_chains( - source_images, logger, record_as=record_as, job=job, link_costs=costs - ) - counters["links_created"] = links - counters["transitions_too_far_apart"] = skipped_for_interval + counters["transitions_too_far_apart"] = plan.transitions_too_far_apart counters["transitions_without_dimensions"] = skipped_for_dimensions return counters @@ -383,27 +439,50 @@ def _skip_reason(self, event: Event) -> str | None: return "it has human identifications" return None - def _track_session(self, event: Event) -> tuple[dict[str, int] | None, str | None]: - """Track one session in its own transaction, returning ``(counters, None)`` or ``(None, skip reason)``. + def _track_session(self, event: Event, progress_cb: typing.Callable[[float], None] | None = None): + """Track one session in two phases, returning ``(counters, None)`` or ``(None, skip reason)``. - The checks and the writes share one lock on the session, so an edit made since the job - started is seen and an edit made during the run waits. Nothing here saves the job, so its row - is not locked while the session is processed. + Matching reads only and runs outside any transaction, so the job can keep saving progress + while it works. Writing is one short transaction that locks the session, repeats the guards, + and refuses to write when the session changed since it was matched. """ + if self._skip_reason(event) is not None: + plan = None # the locked check below reports the reason + else: + plan = plan_session_links(event, self.logger, self.config, progress_cb) with transaction.atomic(): lock_sessions([event.pk]) reason = self._skip_reason(event) if reason is not None: return None, reason - counters = assign_occurrences_by_tracking_images( - event=event, logger=self.logger, config=self.config, record_as=self.algorithm, job=self.job - ) - if not counters: - return None, "it has fewer than two processed captures" + if plan is None: + return None, "it has fewer than two processed captures" + if not plan_is_current(plan): + self.logger.warning(f"Skipping session {event.pk}: it changed while its links were being matched.") + return None, "it changed while it was being tracked" + counters = write_session_plan(plan, self.logger, record_as=self.algorithm, job=self.job) return counters, None + def _throttled_progress(self, index: int, total: int) -> typing.Callable[[float], None]: + """A callback that saves progress for session ``index`` of ``total`` every few transitions or seconds.""" + last_saved = time.monotonic() + matched = 0 + + def report(share: float) -> None: + nonlocal last_saved, matched + matched += 1 + if matched % PROGRESS_EVERY_TRANSITIONS == 0 or time.monotonic() - last_saved >= PROGRESS_EVERY_SECONDS: + self.update_progress(((index - 1) + share) / total) + last_saved = time.monotonic() + + return report + def run(self) -> None: - """Track every session in scope, one transaction per session. + """Track every session in scope, matching outside a transaction and writing in one short one per session. + + Matching saves job progress as it goes, because the stale-job reaper revokes a job whose + ``updated_at`` stops moving and a busy session can take minutes. Saving inside the write + transaction would not help: it stays invisible until commit and locks the job row. A capture-set scope tracks every processed capture of the sessions the set touches, not only the captures in the set, because a chain needs the captures between its detections. A session that @@ -426,7 +505,7 @@ def run(self) -> None: for idx, event in enumerate(events, start=1): self.logger.info(f"Tracking session {idx}/{total} (id={event.pk})") try: - counters, reason = self._track_session(event) + counters, reason = self._track_session(event, self._throttled_progress(idx, total)) except Exception: self.logger.exception(f"Tracking failed for session {event.pk}; its changes were rolled back.") failed_event_ids.append(event.pk) From 5283b0f41f1c78bad1476b353715f71c2af4f456 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Tue, 6 Oct 2026 08:04:34 -0700 Subject: [PATCH 17/28] feat(tracking): label the session list setting so the history shows it by name The occurrence history now labels each job setting with the title its task's config schema declares. Every tracking setting had a title except the list of sessions, which would have shown its raw key. Co-Authored-By: Claude Opus 5.5 --- ami/ml/post_processing/tracking/config.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ami/ml/post_processing/tracking/config.py b/ami/ml/post_processing/tracking/config.py index fc672f445..3da27682e 100644 --- a/ami/ml/post_processing/tracking/config.py +++ b/ami/ml/post_processing/tracking/config.py @@ -34,7 +34,7 @@ class TrackingConfig(pydantic.BaseModel): "between its detections." ), ) - event_ids: list[int] = [] + event_ids: list[int] = pydantic.Field([], title="Sessions") cost_threshold: float = pydantic.Field( 1.0, From 2eb150a701302b836305d836c40919d9c564cb33 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Tue, 6 Oct 2026 08:05:27 -0700 Subject: [PATCH 18/28] test(ui): give the tracking history fixture a value and no score Algorithm result entries now carry their headline figure in value, and score is reserved for predictions. The tracking fixture follows that shape. Co-Authored-By: Claude Opus 5.5 --- ui/src/data-services/models/occurrence-history.test.ts | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/ui/src/data-services/models/occurrence-history.test.ts b/ui/src/data-services/models/occurrence-history.test.ts index 7cd07780a..3fdc7c8e9 100644 --- a/ui/src/data-services/models/occurrence-history.test.ts +++ b/ui/src/data-services/models/occurrence-history.test.ts @@ -134,8 +134,9 @@ const tracking: ServerOccurrenceHistoryEntry = { id: 6, is_current: true, kind: 'tracking', - score: 0.0141, + score: null, type: 'algorithm_result', + value: 0.0071, } const ownIdentification: HumanIdentification = { From 26a0d83eab80d86f45e0f8b62cb691d720a2e1dc Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Tue, 6 Oct 2026 08:22:23 -0700 Subject: [PATCH 19/28] fix(tracking): never delete an occurrence that still holds detections outside the chain When a session is tracked again with the fresh-session guard off, or an older occurrence spans two sessions, a chain can take some of an occurrence's detections while others stay behind. That occurrence was deleted anyway, and the detections it still held were left without an occurrence. Only occurrences the chain emptied are now merged away; identifications and results move only off those. Co-Authored-By: Claude Opus 5.5 --- .../tests/test_tracking_task.py | 23 +++++++++++++++++++ ami/ml/post_processing/tracking/task.py | 8 ++++++- 2 files changed, 30 insertions(+), 1 deletion(-) diff --git a/ami/ml/post_processing/tests/test_tracking_task.py b/ami/ml/post_processing/tests/test_tracking_task.py index 97ceabfe5..1ccac0db4 100644 --- a/ami/ml/post_processing/tests/test_tracking_task.py +++ b/ami/ml/post_processing/tests/test_tracking_task.py @@ -143,6 +143,29 @@ def test_a_session_that_was_already_tracked_is_skipped_unless_the_guard_is_off(s self.run_task(event, require_fresh_event=False) self.assertEqual(self.occurrence_sizes(event), [3]) + def test_an_occurrence_that_keeps_detections_outside_the_chain_is_not_deleted(self): + """Re-tracking must never leave a detection without an occurrence. + + An earlier grouping put the first and last detections in one occurrence. The new run links the + first detection to a neighbour and the last to another, so that occurrence is split across two + chains; it keeps the first chain and must not be deleted by the second. + """ + far = [600, 500, 650, 550] + captures = create_session(self.deployment, [[BOX], [BOX, far], [far]], self.taxa[0]) + event = captures[0].event + first = captures[0].detections.get() + last = captures[2].detections.get() + emptied = last.occurrence + last.occurrence = first.occurrence + last.save(update_fields=["occurrence"]) + emptied.delete() + + self.run_task(event, require_fresh_event=False) + + detections = Detection.objects.filter(source_image__event=event) + self.assertFalse(detections.filter(occurrence__isnull=True).exists()) + self.assertEqual(self.occurrence_sizes(event), [2, 2]) + def test_human_identifications_skip_the_session_unless_the_guard_is_off(self): captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) event = captures[0].event diff --git a/ami/ml/post_processing/tracking/task.py b/ami/ml/post_processing/tracking/task.py index 2bfe07ce7..4fb95e704 100644 --- a/ami/ml/post_processing/tracking/task.py +++ b/ami/ml/post_processing/tracking/task.py @@ -201,7 +201,13 @@ def assign_occurrences_from_detection_chains( for d in moving: d.occurrence = keeper - doomed = old_occ_ids - {keeper.pk} + # Only occurrences the chain emptied are merged away. One that still holds detections outside the + # chain (an earlier run's link this run did not repeat, or another session) keeps them and its records. + doomed = set( + Occurrence.objects.filter(pk__in=old_occ_ids - {keeper.pk}) + .exclude(Exists(Detection.objects.filter(occurrence_id=OuterRef("pk")))) + .values_list("pk", flat=True) + ) if doomed: identifications_moved += Identification.objects.filter(occurrence_id__in=doomed).update(occurrence=keeper) # Deleting an occurrence deletes its results, so they move onto the keeper first. From 2353d773505833f08456880dcc947e667880898f Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Tue, 6 Oct 2026 11:50:52 -0700 Subject: [PATCH 20/28] refactor(tracking): follow the result framework's kind and settings declarations The result framework now names each kind on its data model and reads job setting labels and references from the task's config schema. The tracking writer uses TrackingResultData.kind, and the capture set setting is declared as a reference so the history links it. Co-Authored-By: Claude Opus 5.5 --- ami/ml/post_processing/tracking/config.py | 5 ++++- ami/ml/post_processing/tracking/task.py | 3 ++- 2 files changed, 6 insertions(+), 2 deletions(-) diff --git a/ami/ml/post_processing/tracking/config.py b/ami/ml/post_processing/tracking/config.py index 3da27682e..030376fcb 100644 --- a/ami/ml/post_processing/tracking/config.py +++ b/ami/ml/post_processing/tracking/config.py @@ -5,6 +5,8 @@ import pydantic +from ami.ml.results.schemas import reference + COST_NOTE = ( "The default is a starting point that is still being tuned by experiment. " "It suits captures taken about 20 seconds apart." @@ -25,7 +27,8 @@ class TrackingConfig(pydantic.BaseModel): shown on the admin form. """ - source_image_collection_id: int | None = pydantic.Field( + source_image_collection_id: int | None = reference( + "capture_set", None, title="Capture set", description=( diff --git a/ami/ml/post_processing/tracking/task.py b/ami/ml/post_processing/tracking/task.py index 4fb95e704..b3364cb79 100644 --- a/ami/ml/post_processing/tracking/task.py +++ b/ami/ml/post_processing/tracking/task.py @@ -28,6 +28,7 @@ from ami.ml.models import Algorithm, AlgorithmResult from ami.ml.models.algorithm import AlgorithmTaskType from ami.ml.post_processing.base import BasePostProcessingTask +from ami.ml.results.schemas import TrackingResultData from .config import TrackingConfig from .matching import captures_too_far_apart, image_diagonal, select_links @@ -124,7 +125,7 @@ class masking) are left out, and a classification with no algorithm counts as a occurrence=item.keeper, algorithm=algorithm, job=job, - kind=AlgorithmResult.Kind.TRACKING, + kind=TrackingResultData.kind, value=figures.motion, data={ **dataclasses.asdict(figures), From 0072d307f41c6ced564f83bf52174e9c43c896c9 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Tue, 6 Oct 2026 15:20:17 -0700 Subject: [PATCH 21/28] refactor(tracking): declare the tracking result model and drop the current-result checks The result framework now keeps every run's result instead of marking one current, and each post-processing task declares the result models it writes. The tracking task declares TrackingResultData, and the tests no longer assert a current flag. Co-Authored-By: Claude Opus 5.5 --- ami/main/tests.py | 2 +- ami/ml/post_processing/tests/test_tracking_task.py | 6 ++---- ami/ml/post_processing/tracking/task.py | 1 + ui/src/data-services/models/occurrence-history.test.ts | 1 - 4 files changed, 4 insertions(+), 6 deletions(-) diff --git a/ami/main/tests.py b/ami/main/tests.py index ef73e5334..da3a7c17e 100644 --- a/ami/main/tests.py +++ b/ami/main/tests.py @@ -8680,5 +8680,5 @@ def test_a_tracking_result_stays_on_the_earliest_piece(self): piece = Occurrence.objects.exclude(pk=occurrence.pk).get(deployment=self.deployment) result.refresh_from_db() - self.assertEqual((result.occurrence_id, result.is_current), (occurrence.pk, True)) + self.assertEqual(result.occurrence_id, occurrence.pk) self.assertFalse(AlgorithmResult.objects.filter(occurrence=piece).exists()) diff --git a/ami/ml/post_processing/tests/test_tracking_task.py b/ami/ml/post_processing/tests/test_tracking_task.py index 1ccac0db4..5c7194d76 100644 --- a/ami/ml/post_processing/tests/test_tracking_task.py +++ b/ami/ml/post_processing/tests/test_tracking_task.py @@ -263,9 +263,7 @@ def test_each_linked_occurrence_gets_a_result_with_its_figures(self): self.assertEqual(moving_result.value, 0.0071) self.assertEqual(moving_result.data["path_length"], 0.0141) for result in by_motion: - self.assertEqual( - (result.job_id, result.algorithm_id, result.is_current), (job.pk, task.algorithm.pk, True) - ) + self.assertEqual((result.job_id, result.algorithm_id), (job.pk, task.algorithm.pk)) self.assertEqual(result.project_id, self.project.pk) self.assertEqual(result.data["detection_count"], 3) self.assertEqual(result.data["size_change"], 1.0) @@ -353,7 +351,7 @@ def test_results_of_merged_occurrences_move_onto_the_keeper(self): self.run_task(captures[0].event) earlier.refresh_from_db() - self.assertEqual((earlier.occurrence_id, earlier.is_current), (first.pk, True)) + self.assertEqual(earlier.occurrence_id, first.pk) self.assertFalse(Occurrence.objects.filter(pk=second.pk).exists()) tracking = AlgorithmResult.objects.get(kind="tracking") self.assertEqual(tracking.occurrence_id, first.pk) diff --git a/ami/ml/post_processing/tracking/task.py b/ami/ml/post_processing/tracking/task.py index b3364cb79..c53b003db 100644 --- a/ami/ml/post_processing/tracking/task.py +++ b/ami/ml/post_processing/tracking/task.py @@ -396,6 +396,7 @@ class TrackingTask(BasePostProcessingTask): key = "tracking" name = "Occurrence tracking" config_schema = TrackingConfig + result_models = (TrackingResultData,) config: TrackingConfig diff --git a/ui/src/data-services/models/occurrence-history.test.ts b/ui/src/data-services/models/occurrence-history.test.ts index 3fdc7c8e9..9138db514 100644 --- a/ui/src/data-services/models/occurrence-history.test.ts +++ b/ui/src/data-services/models/occurrence-history.test.ts @@ -132,7 +132,6 @@ const tracking: ServerOccurrenceHistoryEntry = { determination_after: NOCTUA, determination_before: NOCTUA, id: 6, - is_current: true, kind: 'tracking', score: null, type: 'algorithm_result', From e989d46b5fa86bc872f8dd4b664d92e07226476f Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Tue, 6 Oct 2026 16:19:11 -0700 Subject: [PATCH 22/28] test(tracking): build tracking test fixtures once per class The tracking task tests and the regroup-splits-occurrences tests built their project, taxa and (for regroup) captures before every test. Each `setup_test_project()` call is about 0.6 to 0.8 seconds, so the repeated setup dominated these classes. The fixtures are now built once per class with `setUpTestData`; every test still starts from the same rows, because each runs in a transaction that is rolled back and Django gives each test its own copy of the class's objects. Tests, assertions and per-test captures are unchanged. The tracking admin tests already used `setUpTestData`. On a local run of the tracking task, admin, matching and stats files plus the regroup class (64 tests), the summed setup and call time went from about 46.6 to 33.1 seconds for the task file and from 5.3 to 2.3 seconds for the regroup class, and the wall time reported by pytest from 55.96 to 38.84 seconds. The full backend suite passes (859 tests, 2 skipped). Co-Authored-By: Claude Opus 5.5 --- ami/main/tests.py | 15 ++++++++------- .../post_processing/tests/test_tracking_task.py | 9 +++++---- 2 files changed, 13 insertions(+), 11 deletions(-) diff --git a/ami/main/tests.py b/ami/main/tests.py index da3a7c17e..b086ac67a 100644 --- a/ami/main/tests.py +++ b/ami/main/tests.py @@ -8548,15 +8548,16 @@ class TestRegroupSplitsOccurrences(TestCase): leaves behind when a later regroup draws a boundary through it. """ - def setUp(self) -> None: - self.project, self.deployment = setup_test_project(reuse=False) - create_taxa(project=self.project) - self.taxon, self.other_taxon = list(Taxon.objects.filter(projects=self.project).order_by("pk")[:2]) - self.user = User.objects.create_user(email="regroup-identifier@insectai.org") # type: ignore[attr-defined] + @classmethod + def setUpTestData(cls) -> None: + cls.project, cls.deployment = setup_test_project(reuse=False) + create_taxa(project=cls.project) + cls.taxon, cls.other_taxon = list(Taxon.objects.filter(projects=cls.project).order_by("pk")[:2]) + cls.user = User.objects.create_user(email="regroup-identifier@insectai.org") # type: ignore[attr-defined] start = datetime.datetime(2024, 6, 1, 22, 0) - self.captures = [ + cls.captures = [ SourceImage.objects.create( - deployment=self.deployment, + deployment=cls.deployment, timestamp=start + datetime.timedelta(minutes=minutes), path=f"test/regroup-split-{i}.jpg", width=640, diff --git a/ami/ml/post_processing/tests/test_tracking_task.py b/ami/ml/post_processing/tests/test_tracking_task.py index 5c7194d76..1e71a0713 100644 --- a/ami/ml/post_processing/tests/test_tracking_task.py +++ b/ami/ml/post_processing/tests/test_tracking_task.py @@ -32,10 +32,11 @@ class _TrackingCase(TestCase): - def setUp(self) -> None: - self.project, self.deployment = setup_test_project(reuse=False) - create_taxa(self.project) - self.taxa = list(Taxon.objects.filter(projects=self.project, rank="SPECIES").order_by("name")) + @classmethod + def setUpTestData(cls) -> None: + cls.project, cls.deployment = setup_test_project(reuse=False) + create_taxa(cls.project) + cls.taxa = list(Taxon.objects.filter(projects=cls.project, rank="SPECIES").order_by("name")) def run_task(self, event: Event, job: Job | None = None, **config) -> TrackingTask: task = TrackingTask(job=job, logger=logger, event_ids=[event.pk], **config) From caadb93aa9647e51bcf9b10fde5d018609a64897 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Tue, 6 Oct 2026 18:02:52 -0700 Subject: [PATCH 23/28] refactor(tracking): let the result framework derive the tracking value from motion The result framework now copies each kind's headline figure into the result's value from the data field the kind names in value_field, on every write path. The tracking kind declares motion as that field, and the tracking task no longer passes the value by hand. The history entry no longer carries data_references, so the tracking fixture in the UI history tests drops it. The history test for a kind that is no longer registered used "tracking" as its example; tracking is registered on this branch, so the test uses an unregistered kind instead. The AlgorithmResult docstring no longer lists tracking as an upcoming kind. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01QT59KePky4u4nbCsTAggtc --- ami/ml/models/algorithm_result.py | 4 ++-- ami/ml/post_processing/tracking/task.py | 1 - ami/ml/results/schemas.py | 1 + ui/src/data-services/models/occurrence-history.test.ts | 1 - 4 files changed, 3 insertions(+), 4 deletions(-) diff --git a/ami/ml/models/algorithm_result.py b/ami/ml/models/algorithm_result.py index 407545d9e..c409882f8 100644 --- a/ami/ml/models/algorithm_result.py +++ b/ami/ml/models/algorithm_result.py @@ -85,8 +85,8 @@ class AlgorithmResult(BaseModel): ``extra`` object inside ``data`` is stored, shown and exported only; nothing reads it for logic, and a value a feature needs becomes a typed field. Every run adds its own result, so running a method twice leaves two results on the occurrence, one per job. Write through - ``AlgorithmResult.objects.record`` or ``record_many``. Tracking and rank roll-ups are the - next kinds expected. See #1431. + ``AlgorithmResult.objects.record`` or ``record_many``. Rank roll-ups are the next kind + expected. See #1431. """ # Copied from the occurrence when the result is written, so per-project diff --git a/ami/ml/post_processing/tracking/task.py b/ami/ml/post_processing/tracking/task.py index c53b003db..8a468edb5 100644 --- a/ami/ml/post_processing/tracking/task.py +++ b/ami/ml/post_processing/tracking/task.py @@ -126,7 +126,6 @@ class masking) are left out, and a classification with no algorithm counts as a algorithm=algorithm, job=job, kind=TrackingResultData.kind, - value=figures.motion, data={ **dataclasses.asdict(figures), "link_costs": item.link_costs, diff --git a/ami/ml/results/schemas.py b/ami/ml/results/schemas.py index 005c2ea28..ed86f71c0 100644 --- a/ami/ml/results/schemas.py +++ b/ami/ml/results/schemas.py @@ -72,6 +72,7 @@ class TrackingResultData(DeterminationSnapshot): """Figures from the detections the run linked into the occurrence, in capture order.""" kind: ClassVar[str] = "tracking" + value_field: ClassVar[str | None] = "motion" detection_count: int # Mean distance per step between consecutive detection centres, as a fraction of the image diagonal; diff --git a/ui/src/data-services/models/occurrence-history.test.ts b/ui/src/data-services/models/occurrence-history.test.ts index 9138db514..a746d7012 100644 --- a/ui/src/data-services/models/occurrence-history.test.ts +++ b/ui/src/data-services/models/occurrence-history.test.ts @@ -128,7 +128,6 @@ const tracking: ServerOccurrenceHistoryEntry = { path_length: 0.0141, size_change: 1.15, }, - data_references: {}, determination_after: NOCTUA, determination_before: NOCTUA, id: 6, From d68a2adfe35f9aaef3664a255a102d7c2a8a3585 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Tue, 6 Oct 2026 18:41:32 -0700 Subject: [PATCH 24/28] perf(main): stop update_occurrence_determination from running the lookups it clears update_occurrence_determination cleared the cached best_identification and best_prediction properties with hasattr() checks before deleting them. On a cached_property, hasattr() computes the property when it is not cached yet, so every call ran the identification lookup twice and the prediction lookup once only to throw the answers away. The properties are now dropped from the instance dictionary directly. The repeated lookups mostly hit the query cache, so the saving is Python time rather than database round trips. Together with skipping the query cache in tracking's recompute, it cut the tracking write on a session of 14,366 detections from 23.7 s to 12.5 s; the share of each change was not measured separately. Pipeline result saving calls the same function once per occurrence and benefits too. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01QT59KePky4u4nbCsTAggtc --- ami/main/models.py | 10 +++------- 1 file changed, 3 insertions(+), 7 deletions(-) diff --git a/ami/main/models.py b/ami/main/models.py index d59bd221d..c777afac8 100644 --- a/ami/main/models.py +++ b/ami/main/models.py @@ -3983,13 +3983,9 @@ def update_occurrence_determination( """ needs_update = False - # Invalidate the cached properties so they will be re-calculated - if hasattr(occurrence, "best_identification"): - del occurrence.best_identification - if hasattr(occurrence, "best_prediction"): - del occurrence.best_prediction - if hasattr(occurrence, "best_identification"): - del occurrence.best_identification + # Clear the cached properties so they are recalculated. ``hasattr`` would run their queries first. + occurrence.__dict__.pop("best_identification", None) + occurrence.__dict__.pop("best_prediction", None) current_determination = ( current_determination From ed114d9a060120ef20a97e3a8cec413db8780ada Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Tue, 6 Oct 2026 18:42:02 -0700 Subject: [PATCH 25/28] feat(tracking): write a session in bulk, record the grouping for undo, and add a preview A tracking run on the busiest session of a production copy (14,366 detections over 663 captures) held its write transaction for 38 s and issued more than 9,000 statements, because every chain was merged with its own updates, lookups and delete. The merge plan is now worked out in memory by a new Django-free module, chains.py, and written table by table: one statement each for the links and the detections' occurrences, one insert for new occurrences, one delete, and a determination recompute per changed occurrence. The same session now holds the transaction for 11.6 s with about 6,050 statements, nearly all of them the determination recompute. A run only adds. It merges every occurrence its links join, whole, so it never takes a detection out of its occurrence, and it never replaces or removes an existing link: a detection that already links on is not a source, and one already linked to is not a target. The freshness guard now asks whether any detection of the session has a link, which uses the new unique index, instead of looking for occurrences with several detections. That check also skipped sessions grouped by other means that were never tracked. Each tracking result now records the grouping before the run: the occurrence's detections in capture order, the occurrence each was in, the identifications moved onto it, and those withdrawn. A reset can use this to restore the earlier grouping. Identifications moved by a merge skip Identification.save, so a user who had identified two merged occurrences is left with only the newest one active, as saving an identification does. The write locks the session's occurrences first. An identification saved on one of them waits for the run, so the human-identification guard sees it, and it can no longer land on an occurrence that is then deleted. That guard now follows the session's detections rather than Occurrence.event. New settings: "Preview only" works out the links and merges and reports the counts on the job without changing anything. "Detector" limits a run to one detection algorithm. Without it, a session with detections from several detectors is skipped with a reason, since two detectors find the same insect twice and their boxes would form parallel chains. Captures without a timestamp are left out of the sequence. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01QT59KePky4u4nbCsTAggtc --- .../tests/test_tracking_chains.py | 40 + .../tests/test_tracking_task.py | 147 +++- ami/ml/post_processing/tracking/chains.py | 81 ++ ami/ml/post_processing/tracking/config.py | 24 +- ami/ml/post_processing/tracking/task.py | 695 ++++++++++-------- ami/ml/results/schemas.py | 7 + .../models/occurrence-history.ts | 6 + 7 files changed, 685 insertions(+), 315 deletions(-) create mode 100644 ami/ml/post_processing/tests/test_tracking_chains.py create mode 100644 ami/ml/post_processing/tracking/chains.py diff --git a/ami/ml/post_processing/tests/test_tracking_chains.py b/ami/ml/post_processing/tests/test_tracking_chains.py new file mode 100644 index 000000000..ac08d1ead --- /dev/null +++ b/ami/ml/post_processing/tests/test_tracking_chains.py @@ -0,0 +1,40 @@ +"""The merge plan a tracking run works out from links and the current grouping, without a database.""" + +from unittest import TestCase + +from ami.ml.post_processing.tracking.chains import merge_groups + + +class TestMergeGroups(TestCase): + def test_a_linked_chain_merges_into_the_first_occurrence(self): + groups = merge_groups([1, 2, 3], {1: 10, 2: 20, 3: 30}, {1: 2, 2: 3}, {1: 0.1, 2: 0.2}) + + self.assertEqual(len(groups), 1) + group = groups[0] + self.assertEqual((group.detection_ids, group.keeper_id), ([1, 2, 3], 10)) + self.assertEqual(group.previous_occurrence_ids, [10, 20, 30]) + self.assertEqual(group.absorbed_ids, [20, 30]) + self.assertEqual(group.link_costs, [0.1, 0.2]) + + def test_a_group_already_held_by_one_occurrence_is_left_out(self): + self.assertEqual(merge_groups([1, 2], {1: 10, 2: 10}, {1: 2}, {}), []) + + def test_detections_sharing_an_occurrence_stay_together(self): + """Linking one detection of an occurrence brings the whole occurrence, so nothing is split.""" + groups = merge_groups([1, 2, 3], {1: 10, 2: 20, 3: 10}, {2: 3}, {2: 0.5}) + + self.assertEqual(len(groups), 1) + self.assertEqual((groups[0].detection_ids, groups[0].keeper_id), ([1, 2, 3], 10)) + + def test_a_group_without_an_occurrence_has_no_keeper(self): + groups = merge_groups([1, 2], {1: None, 2: None}, {1: 2}, {1: 0.3}) + + self.assertEqual([(g.detection_ids, g.keeper_id) for g in groups], [([1, 2], None)]) + + def test_the_keeper_is_the_first_occurrence_in_capture_order(self): + groups = merge_groups([5, 1], {5: 50, 1: 10}, {5: 1}, {5: 0.1}) + + self.assertEqual(groups[0].keeper_id, 50) + + def test_a_link_to_a_detection_outside_the_plan_is_ignored(self): + self.assertEqual(merge_groups([1], {1: 10}, {1: 99}, {}), []) diff --git a/ami/ml/post_processing/tests/test_tracking_task.py b/ami/ml/post_processing/tests/test_tracking_task.py index 1e71a0713..7833a4e70 100644 --- a/ami/ml/post_processing/tests/test_tracking_task.py +++ b/ami/ml/post_processing/tests/test_tracking_task.py @@ -107,8 +107,9 @@ def test_a_chain_stops_at_a_session_boundary(self): last_of_first.next_detection = first_of_second last_of_first.save(update_fields=["next_detection"]) + # The stored link marks both sessions as tracked, so the guard is turned off. for event in (first[0].event, second[0].event): - self.run_task(event) + self.run_task(event, require_fresh_event=False) self.assertEqual(self.occurrence_sizes(first[0].event), [2]) self.assertEqual(self.occurrence_sizes(second[0].event), [2]) @@ -144,12 +145,11 @@ def test_a_session_that_was_already_tracked_is_skipped_unless_the_guard_is_off(s self.run_task(event, require_fresh_event=False) self.assertEqual(self.occurrence_sizes(event), [3]) - def test_an_occurrence_that_keeps_detections_outside_the_chain_is_not_deleted(self): - """Re-tracking must never leave a detection without an occurrence. + def test_a_run_merges_whole_occurrences_and_never_takes_a_detection_out_of_one(self): + """A run only adds: chains joined through an earlier grouping end up in one occurrence. - An earlier grouping put the first and last detections in one occurrence. The new run links the - first detection to a neighbour and the last to another, so that occurrence is split across two - chains; it keeps the first chain and must not be deleted by the second. + An earlier grouping put the first and last detections in one occurrence. The run links the first + detection to a neighbour and the last to another, so both chains and that occurrence become one. """ far = [600, 500, 650, 550] captures = create_session(self.deployment, [[BOX], [BOX, far], [far]], self.taxa[0]) @@ -165,7 +165,21 @@ def test_an_occurrence_that_keeps_detections_outside_the_chain_is_not_deleted(se detections = Detection.objects.filter(source_image__event=event) self.assertFalse(detections.filter(occurrence__isnull=True).exists()) - self.assertEqual(self.occurrence_sizes(event), [2, 2]) + self.assertEqual(self.occurrence_sizes(event), [4]) + + def test_a_session_grouped_earlier_without_links_is_tracked(self): + """Only a stored link marks a session as tracked; an occurrence of several detections does not.""" + captures = create_session(self.deployment, [[BOX], [BOX], [BOX]], self.taxa[0]) + event = captures[0].event + first, second = (c.detections.get() for c in captures[:2]) + emptied = second.occurrence + second.occurrence = first.occurrence + second.save(update_fields=["occurrence"]) + emptied.delete() + + self.run_task(event) + + self.assertEqual(self.occurrence_sizes(event), [3]) def test_human_identifications_skip_the_session_unless_the_guard_is_off(self): captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) @@ -194,6 +208,12 @@ def test_identifications_move_onto_the_keeper_instead_of_being_deleted(self): self.assertEqual(keeper.detections.count(), 3) self.assertEqual(Identification.objects.filter(user=user).count(), 3) self.assertEqual(set(Identification.objects.values_list("occurrence_id", flat=True)), {keeper.pk}) + # As when an identification is saved, the user keeps one active identification: the newest. + newest = Identification.objects.filter(user=user).order_by("-created_at", "-pk").first() + self.assertEqual(list(Identification.objects.filter(user=user, withdrawn=False)), [newest]) + result = AlgorithmResult.objects.get(kind="tracking") + self.assertEqual(len(result.data["moved_identifications"]), 2) + self.assertEqual(len(result.data["withdrawn_identification_ids"]), 2) def _two_detection_chain(self, first_score: float, second_score: float): captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) @@ -208,7 +228,7 @@ def _two_detection_chain(self, first_score: float, second_score: float): def test_class_masking_after_tracking_still_changes_the_determination(self): """A later re-scoring replaces a merged occurrence's determination, because tracking adds no classification.""" captures, first, second = self._two_detection_chain(first_score=0.3, second_score=0.9) - self.run_task(captures[0].event) + self.run_task(captures[0].event, require_fresh_event=False) occurrence = Occurrence.objects.get(pk=first.occurrence_id) self.assertEqual(occurrence.determination, self.taxa[1]) self.assertFalse(Classification.objects.filter(algorithm__key="tracking").exists()) @@ -508,10 +528,11 @@ def plan_then_change(*args, **kwargs): class TestTrackingQueries(_TrackingCase): - def count_queries(self, boxes_per_capture: list[list[list[int]] | None], start: datetime.datetime) -> int: + def count_queries(self, captures_count: int, insects: int, start: datetime.datetime) -> int: from cachalot.api import cachalot_disabled - captures = create_session(self.deployment, boxes_per_capture, self.taxa[0], start=start) + boxes = [[x, 100, x + 100, 200] for x in range(0, 300 * insects, 300)] + captures = create_session(self.deployment, [boxes] * captures_count, self.taxa[0], start=start) disabled = cachalot_disabled() disabled.__enter__() try: @@ -520,15 +541,101 @@ def count_queries(self, boxes_per_capture: list[list[list[int]] | None], start: finally: # cachalot_disabled() does not restore itself when the block raises. disabled.__exit__(None, None, None) - self.assertEqual(self.occurrence_sizes(captures[0].event), [len(boxes_per_capture)]) + self.assertEqual(self.occurrence_sizes(captures[0].event), [captures_count] * insects) return len(queries) - def test_queries_grow_by_a_constant_per_capture_not_per_detection(self): - """Each extra capture adds a fixed number of queries; reassigning a chain's detections is one update.""" - self.count_queries([[BOX]] * 3, datetime.datetime(2026, 6, 30)) # warms per-process caches - counts = [ - self.count_queries([[BOX]] * n, datetime.datetime(2026, 7, day)) for day, n in ((1, 3), (2, 6), (3, 9)) - ] - self.assertEqual(counts[2] - counts[1], counts[1] - counts[0]) - per_capture = (counts[1] - counts[0]) // 3 - self.assertLessEqual(per_capture, 4) + def test_queries_do_not_grow_with_the_number_of_detections(self): + """Links and merges are written in bulk; only the determination recompute costs queries per occurrence. + + Each extra capture adds the same number of queries however many insects it holds, and each extra + insect adds a fixed number, however many captures its chain spans. + """ + self.count_queries(3, 1, datetime.datetime(2026, 6, 30)) # warms per-process caches + one = [self.count_queries(n, 1, datetime.datetime(2026, 7, day)) for day, n in ((1, 3), (2, 6))] + three = [self.count_queries(n, 3, datetime.datetime(2026, 7, day)) for day, n in ((3, 3), (4, 6))] + self.assertEqual(one[1] - one[0], three[1] - three[0]) + self.assertLessEqual((one[1] - one[0]) // 3, 2) + self.assertEqual(three[0] - one[0], three[1] - one[1]) + self.assertLessEqual((three[0] - one[0]) // 2, 3) + + +class TestPreview(_TrackingCase): + def test_a_preview_reports_the_counts_and_changes_nothing(self): + captures = create_session(self.deployment, [[BOX], [BOX], [BOX]], self.taxa[0]) + event = captures[0].event + job = Job.objects.create(name="t", project=self.project, job_type_key="post_processing") + job.progress.add_stage("Post-processing", key="post_processing") + job.save() + + self.run_task(event, job=job, preview_only=True) + + self.assertEqual(self.occurrence_sizes(event), [1, 1, 1]) + self.assertFalse(Detection.objects.filter(source_image__event=event, next_detection__isnull=False).exists()) + self.assertFalse(AlgorithmResult.objects.exists()) + job.refresh_from_db() + params = {p.name: p.value for p in job.progress.get_stage("post_processing").params} + self.assertEqual(params["Detection links that would be created"], 2) + self.assertEqual(params["Occurrences that would be merged away"], 2) + self.assertEqual((params["Occurrences before"], params["Occurrences after the run"]), (3, 1)) + self.assertTrue(params["Result"].startswith("Preview only, nothing was changed.")) + + self.run_task(event, job=job) + self.assertEqual(self.occurrence_sizes(event), [3]) + + +class TestDetectors(_TrackingCase): + def two_detector_session(self) -> tuple[Event, Algorithm, Algorithm]: + captures = create_session(self.deployment, [[BOX], [BOX]], self.taxa[0]) + first, second = (Algorithm.objects.create(name=f"Detector {n}", key=f"detector-{n}") for n in (1, 2)) + Detection.objects.filter(source_image__in=captures).update(detection_algorithm=first) + for capture in captures: + duplicate = add_detection(capture, [102, 102, 202, 202], self.taxa[0]) + duplicate.detection_algorithm = second + duplicate.save(update_fields=["detection_algorithm"]) + return captures[0].event, first, second + + def test_a_session_with_two_detectors_is_skipped_with_a_reason(self): + event, _, _ = self.two_detector_session() + job = Job.objects.create(name="t", project=self.project, job_type_key="post_processing") + job.progress.add_stage("Post-processing", key="post_processing") + job.save() + + self.run_task(event, job=job) + + self.assertEqual(self.occurrence_sizes(event), [1, 1, 1, 1]) + job.refresh_from_db() + params = {p.name: p.value for p in job.progress.get_stage("post_processing").params} + self.assertIn("more than one detector", params["Result"]) + + def test_a_chosen_detector_links_only_its_own_detections(self): + event, first, second = self.two_detector_session() + + self.run_task(event, detection_algorithm_id=second.pk) + + self.assertEqual(self.occurrence_sizes(event), [1, 1, 2]) + linked = Detection.objects.filter(source_image__event=event, next_detection__isnull=False).get() + self.assertEqual(linked.detection_algorithm_id, second.pk) + + +class TestCaptureOrder(_TrackingCase): + def test_a_capture_without_a_timestamp_is_left_out(self): + captures = create_session(self.deployment, [[BOX], [BOX], [BOX]], self.taxa[0]) + type(captures[1]).objects.filter(pk=captures[1].pk).update(timestamp=None) + + self.run_task(captures[0].event) + + self.assertEqual(self.occurrence_sizes(captures[0].event), [1, 2]) + self.assertIsNone(captures[1].detections.get().next_detection_id) + + +class TestUndoRecord(_TrackingCase): + def test_the_result_records_the_grouping_before_the_run(self): + captures = create_session(self.deployment, [[BOX], [BOX], [BOX]], self.taxa[0]) + detections = [c.detections.get() for c in captures] + + self.run_task(captures[0].event) + + result = AlgorithmResult.objects.get(kind="tracking") + self.assertEqual(result.data["detection_ids"], [d.pk for d in detections]) + self.assertEqual(result.data["previous_occurrence_ids"], [d.occurrence_id for d in detections]) + self.assertEqual(result.data["merged_occurrence_ids"], sorted(d.occurrence_id for d in detections[1:])) diff --git a/ami/ml/post_processing/tracking/chains.py b/ami/ml/post_processing/tracking/chains.py new file mode 100644 index 000000000..bc78c4e4b --- /dev/null +++ b/ami/ml/post_processing/tracking/chains.py @@ -0,0 +1,81 @@ +"""Work out which occurrences a tracking run merges, from links and the current grouping, without a database. + +A run only adds: it links detections and merges the occurrences those links join, and it never +moves a detection away from the other detections of its occurrence. So the unit of change is a +group of detections connected by links (earlier ones and this run's) or by sharing an occurrence. +""" + +import dataclasses +from collections.abc import Mapping, Sequence + + +@dataclasses.dataclass +class MergeGroup: + """Detections a run puts in one occurrence, in capture order, and the occurrence each held before.""" + + detection_ids: list[int] + previous_occurrence_ids: list[int | None] + # The occurrence the group keeps: the first one held by a detection in capture order, or None to create one. + keeper_id: int | None + # The cost of each link this run made inside the group, in capture order of the earlier detection. + link_costs: list[float] + + @property + def absorbed_ids(self) -> list[int]: + """Occurrences other than the keeper that held detections of the group, in id order.""" + return sorted({pk for pk in self.previous_occurrence_ids if pk is not None and pk != self.keeper_id}) + + +def merge_groups( + detection_order: Sequence[int], + occurrence_of: Mapping[int, int | None], + links: Mapping[int, int], + new_link_costs: Mapping[int, float], +) -> list[MergeGroup]: + """The groups whose grouping a run changes, in capture order of their first detection. + + ``detection_order`` lists the detections in capture order; ``occurrence_of`` gives each one's + occurrence; ``links`` maps a detection to the next one, for earlier links and this run's; + ``new_link_costs`` gives the cost of this run's links, keyed by the earlier detection. A group + already held by one occurrence, with every detection in it, is left out because nothing changes. + """ + parent = {pk: pk for pk in detection_order} + + def find(pk: int) -> int: + while parent[pk] != pk: + parent[pk] = parent[parent[pk]] + pk = parent[pk] + return pk + + def union(a: int, b: int) -> None: + root_a, root_b = find(a), find(b) + if root_a != root_b: + parent[root_b] = root_a + + for source, target in links.items(): + if source in parent and target in parent: + union(source, target) + first_with_occurrence: dict[int, int] = {} + for pk in detection_order: + occurrence_id = occurrence_of.get(pk) + if occurrence_id is None: + continue + if occurrence_id in first_with_occurrence: + union(first_with_occurrence[occurrence_id], pk) + else: + first_with_occurrence[occurrence_id] = pk + + members: dict[int, list[int]] = {} + for pk in detection_order: + members.setdefault(find(pk), []).append(pk) + + groups = [] + for detection_ids in members.values(): + previous = [occurrence_of.get(pk) for pk in detection_ids] + held_by = {pk for pk in previous if pk is not None} + if len(held_by) == 1 and None not in previous: + continue + keeper_id = next((pk for pk in previous if pk is not None), None) + costs = [round(new_link_costs[pk], 4) for pk in detection_ids if pk in new_link_costs] + groups.append(MergeGroup(detection_ids, previous, keeper_id, costs)) + return groups diff --git a/ami/ml/post_processing/tracking/config.py b/ami/ml/post_processing/tracking/config.py index 030376fcb..c30ddf02d 100644 --- a/ami/ml/post_processing/tracking/config.py +++ b/ami/ml/post_processing/tracking/config.py @@ -38,6 +38,16 @@ class TrackingConfig(pydantic.BaseModel): ), ) event_ids: list[int] = pydantic.Field([], title="Sessions") + detection_algorithm_id: int | None = reference( + "algorithm", + None, + title="Detector", + description=( + "Compare only the detections from this detection algorithm (its id). Leave blank to use the only detector " + "in each session; a session with detections from more than one detector is then skipped, because " + "two detectors find the same insect twice." + ), + ) cost_threshold: float = pydantic.Field( 1.0, @@ -119,8 +129,18 @@ class TrackingConfig(pydantic.BaseModel): True, title="Only track sessions that have not been tracked", description=( - "Skip a session when any of its occurrences already holds more than one detection. Turned off, a " - "run can add links and merges to such a session, but it never undoes earlier ones." + "Skip a session when any of its detections is already linked to a next one. Turned off, a run " + "adds links between detections that have none and merges the occurrences they join; it never " + "removes or replaces a link, and never takes a detection out of its occurrence." + ), + ) + + preview_only: bool = pydantic.Field( + False, + title="Preview only", + description=( + "Work out the links and merges and report how many there would be on the job, without changing " + "anything." ), ) diff --git a/ami/ml/post_processing/tracking/task.py b/ami/ml/post_processing/tracking/task.py index 8a468edb5..6d5058fcc 100644 --- a/ami/ml/post_processing/tracking/task.py +++ b/ami/ml/post_processing/tracking/task.py @@ -1,7 +1,7 @@ -"""The tracking post-processing task: links detections across consecutive captures and merges each chain. +"""The tracking post-processing task: links detections in consecutive captures and merges the occurrences they join. -The matching rules live in ``matching.py`` and the settings in ``config.py``, both free of Django; -this module reads and writes the database around them. +The matching rules live in ``matching.py``, the settings in ``config.py`` and the merge plan in +``chains.py``, all free of Django; this module reads and writes the database around them. """ import collections @@ -9,9 +9,10 @@ import logging import time import typing -from collections.abc import Iterable, Iterator, Sequence +from collections.abc import Iterator, Sequence -from django.db import transaction +from cachalot.api import cachalot_disabled +from django.db import connection, transaction from django.db.models import Count, Exists, OuterRef from ami.main.models import ( @@ -30,6 +31,7 @@ from ami.ml.post_processing.base import BasePostProcessingTask from ami.ml.results.schemas import TrackingResultData +from .chains import MergeGroup, merge_groups from .config import TrackingConfig from .matching import captures_too_far_apart, image_diagonal, select_links from .sessions import lock_sessions @@ -42,243 +44,72 @@ # Progress is saved after this many matched transitions, or after this many seconds, whichever comes first. PROGRESS_EVERY_TRANSITIONS = 25 PROGRESS_EVERY_SECONDS = 5.0 +# Rows per statement when occurrence determinations are written in bulk. +WRITE_BATCH_SIZE = 1000 +Link = tuple[int, int, float] -@dataclasses.dataclass -class LinkedOccurrence: - """An occurrence the run built from a chain of two or more detections, or by merging occurrences.""" - keeper: Occurrence - detections: list[Detection] - determination_before_id: int | None - merged_ids: list[int] - link_costs: list[float] +class SkipSession(Exception): + """Raised with the reason a session is not tracked; the job counts sessions by reason.""" -def event_is_fresh(event: Event) -> tuple[bool, str]: - """Has this event's detections already been grouped into chains? +def session_has_links(event: Event) -> bool: + """Whether a detection in this session already links to a next one, which marks a session as tracked.""" + return Detection.objects.filter(source_image__event=event, next_detection__isnull=False).exists() - The guard keeps tracking away from events that were already consolidated: merging - those again can delete an occurrence that carries identifications. An occurrence - spanning more than one detection is the signal for that. - A detection with no occurrence at all is not that signal, since the chain walk creates an - occurrence for a chain that has none. Occurrences are found through their detections' - captures, not ``Occurrence.event``, so one that reaches into this session from another - session also counts. - """ - multi_detection_occurrences = ( - Occurrence.objects.filter(pk__in=Detection.objects.filter(source_image__event=event).values("occurrence_id")) - .annotate(_n=Count("detections")) - .filter(_n__gt=1) - .count() - ) - if multi_detection_occurrences: - return False, f"{multi_detection_occurrences} occurrence(s) already span >1 detection" - return True, "" +def session_has_identifications(event: Event) -> bool: + """Whether someone has identified an occurrence that holds a detection in this session.""" + return Identification.objects.filter(occurrence__detections__source_image__event=event).exists() def processed_captures(event: Event) -> list[SourceImage]: - """Captures of a session that have at least one detection row, oldest first. + """Captures of a session that have at least one detection row and a timestamp, oldest first. A null-bbox marker row counts: it records that a capture was processed and found nothing. - Captures nobody processed are left out, so they cannot break the adjacency of their neighbours. + Captures nobody processed are left out, so they cannot break the adjacency of their neighbours, + and so are captures with no timestamp, whose place in the sequence is unknown. """ return list( - SourceImage.objects.filter(event=event) + SourceImage.objects.filter(event=event, timestamp__isnull=False) .filter(Exists(Detection.objects.filter(source_image=OuterRef("pk")))) .order_by("timestamp", "pk") ) -def record_tracking_results( - linked: Sequence[LinkedOccurrence], algorithm: Algorithm, job: "Job | None" -) -> dict[int, AlgorithmResult]: - """Write one tracking result per linked occurrence, keyed by occurrence id. - - The figures are computed in memory from the chains' detections plus one query for their terminal - classifications. Only machine labels count: classifications by post-processing algorithms (size filter, - class masking) are left out, and a classification with no algorithm counts as a source label. - Call after the determinations have settled. - """ - if not linked: - return {} - detection_ids = [d.pk for item in linked for d in item.detections] - labels: dict[int, list[int | None]] = collections.defaultdict(list) - for detection_id, taxon_id in ( - Classification.objects.filter(detection_id__in=detection_ids, terminal=True) - .exclude(algorithm__task_type=AlgorithmTaskType.POST_PROCESSING.value) - .values_list("detection_id", "taxon_id") - ): - labels[detection_id].append(taxon_id) - - results = [] - for item in linked: - figures = occurrence_figures( - boxes=[d.bbox for d in item.detections], - sizes=[(d.source_image.width, d.source_image.height) for d in item.detections], - labels=[label for d in item.detections for label in labels.get(d.pk, [])], - determination_id=item.keeper.determination_id, - ) - results.append( - AlgorithmResult( - occurrence=item.keeper, - algorithm=algorithm, - job=job, - kind=TrackingResultData.kind, - data={ - **dataclasses.asdict(figures), - "link_costs": item.link_costs, - "determination_before_id": item.determination_before_id, - "determination_after_id": item.keeper.determination_id, - "merged_occurrence_ids": item.merged_ids, - }, - ) - ) - return {result.occurrence_id: result for result in AlgorithmResult.objects.record_many(results)} - - -def assign_occurrences_from_detection_chains( - source_images: Sequence[SourceImage], - logger: logging.Logger, - record_as: Algorithm | None = None, - job: "Job | None" = None, - link_costs: dict[int, float] | None = None, -) -> dict[str, int]: - """Fold each chain of linked detections into one occurrence, keeping the first existing one. - - A chain never leaves the given captures, which belong to one session, so a link that crosses a - session boundary starts a new chain on each side. Identifications move onto the keeper before the - occurrences that held them are deleted, because deleting an occurrence deletes its identifications. - Results of the absorbed occurrences move onto the keeper first. With ``record_as`` set, every keeper - built from two or more detections or from a merge gets a tracking result attributed to ``job``, which - records the determination before and after, so no classification is written for it. ``link_costs`` maps a - detection to the cost of its link to the next one, for the links this run made. - """ - image_ids = [image.pk for image in source_images] - detections = list( - Detection.objects.valid().filter(source_image_id__in=image_ids).select_related("source_image", "occurrence") - ) - by_id = {det.pk: det for det in detections} - # Walk each capture in time order, so a chain starts at its earliest detection. - position = {image_id: i for i, image_id in enumerate(image_ids)} - detections.sort(key=lambda d: (position[d.source_image_id], d.pk)) - has_previous = {det.next_detection_id for det in detections if det.next_detection_id in by_id} - - visited: set[int] = set() - created = merged = identifications_moved = 0 - linked: dict[int, LinkedOccurrence] = {} - deleted: set[int] = set() - existing = Occurrence.objects.filter(detections__source_image_id__in=image_ids).distinct().count() - - for det in detections: - if det.pk in visited or det.pk in has_previous: - continue - chain: list[Detection] = [] - current: Detection | None = det - while current is not None and current.pk not in visited: - chain.append(current) - visited.add(current.pk) - current = by_id.get(current.next_detection_id) if current.next_detection_id else None - - old_occ_ids = {d.occurrence_id for d in chain if d.occurrence_id} - # A chain already held by exactly one occurrence needs no change. - if len(old_occ_ids) == 1 and all(d.occurrence_id is not None for d in chain): - continue - - keeper: Occurrence | None = next((d.occurrence for d in chain if d.occurrence_id), None) - previous_determination_id = keeper.determination_id if keeper is not None else None - if keeper is None: - first_image = chain[0].source_image - keeper = Occurrence.objects.create( - event=first_image.event, deployment=first_image.deployment, project=first_image.project - ) - created += 1 - - moving = [d for d in chain if d.occurrence_id != keeper.pk] - if moving: - Detection.objects.filter(pk__in=[d.pk for d in moving]).update(occurrence=keeper) - for d in moving: - d.occurrence = keeper - - # Only occurrences the chain emptied are merged away. One that still holds detections outside the - # chain (an earlier run's link this run did not repeat, or another session) keeps them and its records. - doomed = set( - Occurrence.objects.filter(pk__in=old_occ_ids - {keeper.pk}) - .exclude(Exists(Detection.objects.filter(occurrence_id=OuterRef("pk")))) - .values_list("pk", flat=True) - ) - if doomed: - identifications_moved += Identification.objects.filter(occurrence_id__in=doomed).update(occurrence=keeper) - # Deleting an occurrence deletes its results, so they move onto the keeper first. - AlgorithmResult.objects.move_to_occurrence(keeper, doomed) - Occurrence.objects.filter(pk__in=doomed).delete() - deleted |= doomed - merged += len(doomed) - - # Only the determination is written, so a column this run did not change keeps its stored value. - if update_occurrence_determination(keeper, save=False): - keeper.save(update_determination=False, update_fields=["determination", "determination_score"]) - if record_as is not None and (len(chain) > 1 or doomed): - costs = [round(link_costs[d.pk], 4) for d in chain[:-1] if d.pk in link_costs] - linked[keeper.pk] = LinkedOccurrence(keeper, chain, previous_determination_id, sorted(doomed), costs) - - results: dict[int, AlgorithmResult] = {} - if record_as is not None: - # A keeper that a later chain absorbed no longer exists. - final = [item for pk, item in linked.items() if pk not in deleted] - results = record_tracking_results(final, record_as, job) - - new_count = Occurrence.objects.filter(detections__source_image_id__in=image_ids).distinct().count() - logger.info( - f"Created {created} occurrences and merged {merged} across {len(image_ids)} captures " - f"(occurrences before: {existing}, after: {new_count}). Moved {identifications_moved} identification(s), " - f"recorded {len(results)} result(s)." - ) - return { - "occurrences_before": existing, - "occurrences_after": new_count, - "occurrences_created": created, - "occurrences_merged": merged, - "identifications_moved": identifications_moved, - "results_recorded": len(results), - } - - -def save_links(links: Iterable[tuple[Detection, Detection, float]], logger: logging.Logger) -> None: - """Store each link as ``next_detection``, first detaching any other detection that points at the target.""" - links = list(links) - if not links: - return - Detection.objects.filter(next_detection_id__in=[nxt.pk for _, nxt, _ in links]).update(next_detection=None) - for det, nxt, cost in links: - det.next_detection = nxt - logger.debug(f"Linked detection {det.pk} -> {nxt.pk} (cost {cost:.4f})") - Detection.objects.bulk_update([det for det, _, _ in links], ["next_detection"]) - - @dataclasses.dataclass class SessionPlan: - """The links proposed for one session, computed without writing anything. + """What a run would change in one session, worked out without writing anything. ``snapshot`` records each detection's ``next_detection`` and occurrence as they were read, so the write phase can tell whether the session changed while the links were being matched. """ + event: Event source_images: list[SourceImage] + detection_algorithm_id: int | None + detections: dict[int, Detection] snapshot: dict[int, tuple[int | None, int | None]] - proposals: list[list[tuple[Detection, Detection, float]] | None] + links: list[Link] + groups: list[MergeGroup] transitions_too_far_apart: int + transitions_without_dimensions: int + + def occurrence_ids(self) -> set[int]: + return {occurrence_id for _, occurrence_id in self.snapshot.values() if occurrence_id is not None} def iter_transition_links( source_images: Sequence[SourceImage], - detections_by_capture: dict[int, dict[int, Detection]], + detections_by_capture: dict[int, list[Detection]], config: TrackingConfig, logger: logging.Logger, -) -> Iterator[list[tuple[Detection, Detection, float]] | None]: +) -> Iterator[list[Link] | None]: """Yield the proposed links for each pair of consecutive captures, in order, saving nothing. + A detection that already links to a next one is not a candidate to link from, and one that an + earlier link already points at is not a candidate to link to, so a run never replaces a link. Yields None for a transition that is not compared (the earlier capture has no dimensions) and an empty list for one over the interval limit. """ @@ -291,15 +122,21 @@ def iter_transition_links( logger.warning(f"Capture {cur.pk} has no dimensions; not comparing it with the next capture.") yield None continue - current = detections_by_capture.get(cur.pk, {}) - following = detections_by_capture.get(nxt.pk, {}) links = select_links( - [(pk, det.bbox) for pk, det in current.items()], - [(pk, det.bbox) for pk, det in following.items()], + [(d.pk, d.bbox) for d in detections_by_capture.get(cur.pk, []) if d.next_detection_id is None], + [(d.pk, d.bbox) for d in detections_by_capture.get(nxt.pk, []) if not d.has_previous], image_diagonal(cur.width, cur.height), config, ) - yield [(current[from_id], following[to_id], cost) for from_id, to_id, cost in links] + yield list(links) + + +def session_detections(source_images: Sequence[SourceImage], detection_algorithm_id: int | None) -> list[Detection]: + """The detections tracking compares on these captures: those with a box, from one detector when one is given.""" + qs = Detection.objects.valid().filter(source_image_id__in=[image.pk for image in source_images]) + if detection_algorithm_id is not None: + qs = qs.filter(detection_algorithm_id=detection_algorithm_id) + return list(qs.annotate(has_previous=Exists(Detection.objects.filter(next_detection_id=OuterRef("pk"))))) def plan_session_links( @@ -307,74 +144,342 @@ def plan_session_links( logger: logging.Logger, config: TrackingConfig, progress_cb: typing.Callable[[float], None] | None = None, -) -> SessionPlan | None: - """Match the detections of one session's processed captures, writing nothing and holding no lock. +) -> SessionPlan: + """Match one session's processed captures and work out the merges, writing nothing and holding no lock. - Returns None when the session has fewer than two processed captures. ``progress_cb`` receives the - share of transitions matched after each one. + Raises ``SkipSession`` when the session has fewer than two processed captures, or when no detector + was chosen and its detections come from more than one. Two detectors find the same insect twice, + and their boxes would be linked into parallel chains. ``progress_cb`` receives the share of + transitions matched after each one. """ source_images = processed_captures(event) if len(source_images) < 2: - logger.warning(f"Session {event.pk}: fewer than two processed captures ({len(source_images)}).") - return None - detections_by_capture: dict[int, dict[int, Detection]] = collections.defaultdict(dict) - for det in Detection.objects.valid().filter(source_image_id__in=[image.pk for image in source_images]): - detections_by_capture[det.source_image_id][det.pk] = det - snapshot = { - det.pk: (det.next_detection_id, det.occurrence_id) - for detections in detections_by_capture.values() - for det in detections.values() - } + raise SkipSession("it has fewer than two processed captures") + detections = session_detections(source_images, config.detection_algorithm_id) + detectors = {d.detection_algorithm_id for d in detections} + if config.detection_algorithm_id is None and len(detectors) > 1: + logger.warning(f"Session {event.pk} has detections from more than one detector: {sorted(detectors, key=str)}.") + raise SkipSession("its detections come from more than one detector; choose one in the settings") + + position = {image.pk: i for i, image in enumerate(source_images)} + detections.sort(key=lambda d: (position[d.source_image_id], d.pk)) + detections_by_capture: dict[int, list[Detection]] = collections.defaultdict(list) + for det in detections: + detections_by_capture[det.source_image_id].append(det) + transitions = len(source_images) - 1 - proposals = [] + links: list[Link] = [] + without_dimensions = 0 for i, proposed in enumerate(iter_transition_links(source_images, detections_by_capture, config, logger)): - proposals.append(proposed) + if proposed is None: + without_dimensions += 1 + else: + links.extend(proposed) if progress_cb: progress_cb((i + 1) / transitions) too_far = sum( captures_too_far_apart(source_images[i].timestamp, source_images[i + 1].timestamp, config) for i in range(transitions) ) - return SessionPlan(source_images, snapshot, proposals, too_far) + + snapshot = {d.pk: (d.next_detection_id, d.occurrence_id) for d in detections} + all_links = {pk: next_id for pk, (next_id, _) in snapshot.items() if next_id is not None} + all_links.update({source: target for source, target, _ in links}) + groups = merge_groups( + [d.pk for d in detections], + {pk: occurrence_id for pk, (_, occurrence_id) in snapshot.items()}, + all_links, + {source: cost for source, _, cost in links}, + ) + return SessionPlan( + event=event, + source_images=source_images, + detection_algorithm_id=config.detection_algorithm_id, + detections={d.pk: d for d in detections}, + snapshot=snapshot, + links=links, + groups=groups, + transitions_too_far_apart=too_far, + transitions_without_dimensions=without_dimensions, + ) def plan_is_current(plan: SessionPlan) -> bool: - """Whether the session's detections still have the links and occurrences the plan was matched against.""" + """Whether the session still has the captures, links and occurrences the plan was worked out from.""" + if [image.pk for image in processed_captures(plan.event)] != [image.pk for image in plan.source_images]: + return False current = { - pk: (next_id, occurrence_id) - for pk, next_id, occurrence_id in Detection.objects.valid() - .filter(source_image_id__in=[image.pk for image in plan.source_images]) - .values_list("pk", "next_detection_id", "occurrence_id") + d.pk: (d.next_detection_id, d.occurrence_id) + for d in session_detections(plan.source_images, plan.detection_algorithm_id) } return current == plan.snapshot +def lock_occurrences(occurrence_ids: set[int]) -> None: + """Hold a row lock on each occurrence until the transaction ends, in id order so two writers cannot deadlock. + + An identification saved on a locked occurrence waits until the run commits, so the guard against + identified sessions sees it, and it is never written on an occurrence the run is about to delete. + """ + if occurrence_ids: + list( + Occurrence.objects.select_for_update() + .filter(pk__in=sorted(occurrence_ids)) + .order_by("pk") + .values_list("pk") + ) + + +def emptied_occurrences(plan: SessionPlan) -> dict[int, int]: + """Each absorbed occurrence the merge leaves empty, mapped to the occurrence that keeps its detections. + + An absorbed occurrence that also holds detections the plan does not move (in another session, or + from a detector this run does not compare) keeps them and its records, so it is not merged away. + The keeper of a group that has none yet is not known before writing, so such groups are left out; + they have no absorbed occurrences anyway. + """ + absorbed = {occurrence_id: group.keeper_id for group in plan.groups for occurrence_id in group.absorbed_ids} + if not absorbed: + return {} + in_plan = collections.Counter( + occurrence_id for _, occurrence_id in plan.snapshot.values() if occurrence_id in absorbed + ) + in_database = dict( + Detection.objects.filter(occurrence_id__in=list(absorbed)) + .values("occurrence_id") + .annotate(n=Count("pk")) + .values_list("occurrence_id", "n") + ) + return { + pk: keeper for pk, keeper in absorbed.items() if keeper is not None and in_database.get(pk, 0) == in_plan[pk] + } + + +def preview_counts(plan: SessionPlan, emptied: dict[int, int] | None = None) -> dict[str, int]: + """The counts a run would report for this session, without changing anything. + + ``emptied`` is ``emptied_occurrences(plan)`` when the caller already has it. + """ + emptied = emptied_occurrences(plan) if emptied is None else emptied + before = len(plan.occurrence_ids()) + created = sum(1 for group in plan.groups if group.keeper_id is None) + return { + "links_created": len(plan.links), + "occurrences_before": before, + "occurrences_after": before + created - len(emptied), + "occurrences_created": created, + "occurrences_merged": len(emptied), + "transitions_too_far_apart": plan.transitions_too_far_apart, + "transitions_without_dimensions": plan.transitions_without_dimensions, + } + + +def _set_detection_column(column: str, values: dict[int, int]) -> None: + """Set one id column on many detections in a single statement, ``{detection id: value}``. + + ``bulk_update`` would build a CASE with a branch per row, which takes seconds of Python for a busy + session. Django-cachalot invalidates the table for raw writes too. + """ + if not values: + return + table = connection.ops.quote_name(Detection._meta.db_table) + with connection.cursor() as cursor: + cursor.execute( + f"UPDATE {table} AS d SET {connection.ops.quote_name(column)} = v.value " + "FROM unnest(%s::bigint[], %s::bigint[]) AS v(id, value) WHERE d.id = v.id", + [list(values), list(values.values())], + ) + + +def _withdraw_duplicate_identifications(occurrence_ids: set[int]) -> list[int]: + """Leave each user one active identification per occurrence, the newest, as saving an identification does. + + Identifications moved by a merge skip ``Identification.save``, so a user who identified two of the + merged occurrences would otherwise hold two active identifications on the one that is kept. + Returns the ids withdrawn. + """ + seen: set[tuple[int, int]] = set() + withdraw: list[int] = [] + for pk, occurrence_id, user_id in ( + Identification.objects.filter(occurrence_id__in=occurrence_ids, withdrawn=False, user__isnull=False) + .order_by("occurrence_id", "user_id", "-created_at", "-pk") + .values_list("pk", "occurrence_id", "user_id") + ): + if (occurrence_id, user_id) in seen: + withdraw.append(pk) + else: + seen.add((occurrence_id, user_id)) + if withdraw: + Identification.objects.filter(pk__in=withdraw).update(withdrawn=True) + return withdraw + + +def record_tracking_results( + groups: Sequence[MergeGroup], + keepers: dict[int, Occurrence], + plan: SessionPlan, + algorithm: Algorithm, + job: "Job | None", + merged: dict[int, list[int]], + moved_identifications: dict[int, list[tuple[int, int]]], + withdrawn: dict[int, list[int]], + determination_before: dict[int, int | None], +) -> int: + """Write one tracking result per group of two or more detections, or that absorbed an occurrence. + + The figures are computed in memory from the plan's detections plus one query for their terminal + classifications. Only machine labels count: classifications by post-processing algorithms (size filter, + class masking) are left out, and a classification with no algorithm counts as a source label. + Each result also records the occurrence every detection was in before the run, and the + identifications moved or withdrawn, so a reset can put the earlier grouping back. Returns the + number of results written. + """ + recorded = [(g, keepers[g.keeper_id]) for g in groups if len(g.detection_ids) > 1 or merged.get(g.keeper_id)] + if not recorded: + return 0 + detection_ids = [pk for group, _ in recorded for pk in group.detection_ids] + labels: dict[int, list[int | None]] = collections.defaultdict(list) + for detection_id, taxon_id in ( + Classification.objects.filter(detection_id__in=detection_ids, terminal=True) + .exclude(algorithm__task_type=AlgorithmTaskType.POST_PROCESSING.value) + .values_list("detection_id", "taxon_id") + ): + labels[detection_id].append(taxon_id) + + images = {image.pk: image for image in plan.source_images} + results = [] + for group, keeper in recorded: + detections = [plan.detections[pk] for pk in group.detection_ids] + figures = occurrence_figures( + boxes=[d.bbox for d in detections], + sizes=[(images[d.source_image_id].width, images[d.source_image_id].height) for d in detections], + labels=[label for d in detections for label in labels.get(d.pk, [])], + determination_id=keeper.determination_id, + ) + results.append( + AlgorithmResult( + occurrence=keeper, + algorithm=algorithm, + job=job, + kind=TrackingResultData.kind, + data={ + **dataclasses.asdict(figures), + "link_costs": group.link_costs, + "determination_before_id": determination_before.get(keeper.pk), + "determination_after_id": keeper.determination_id, + "merged_occurrence_ids": merged.get(keeper.pk, []), + "detection_ids": group.detection_ids, + "previous_occurrence_ids": group.previous_occurrence_ids, + "moved_identifications": moved_identifications.get(keeper.pk, []), + "withdrawn_identification_ids": withdrawn.get(keeper.pk, []), + }, + ) + ) + return len(AlgorithmResult.objects.record_many(results)) + + def write_session_plan( plan: SessionPlan, logger: logging.Logger, record_as: Algorithm | None = None, job: "Job | None" = None, ) -> dict[str, int]: - """Save a plan's links and fold the chains into occurrences. + """Save a plan's links and merges, writing each table in a few bulk statements. - Call inside the session's transaction, after checking ``plan_is_current``. It does no progress - writes, since saving the job inside the transaction would keep its row locked. + Call inside the session's transaction, holding the session and occurrence locks, after checking + ``plan_is_current``. The statements do not grow with the number of detections; the determination + recompute costs a few queries per changed occurrence. It does no progress writes, since saving the + job inside the transaction would keep its row locked. """ - links = skipped_for_dimensions = 0 - costs: dict[int, float] = {} - for proposed in plan.proposals: - if proposed is None: - skipped_for_dimensions += 1 - continue - save_links(proposed, logger) - links += len(proposed) - costs.update({det.pk: cost for det, _, cost in proposed}) - counters = assign_occurrences_from_detection_chains( - plan.source_images, logger, record_as=record_as, job=job, link_costs=costs + emptied = emptied_occurrences(plan) + counters = preview_counts(plan, emptied) + + _set_detection_column("next_detection_id", {source: target for source, target, _ in plan.links}) + + # Groups with no occurrence get a new one in the session. + to_create = [group for group in plan.groups if group.keeper_id is None] + created = Occurrence.objects.bulk_create( + [ + Occurrence(event=plan.event, deployment_id=plan.event.deployment_id, project_id=plan.event.project_id) + for _ in to_create + ] + ) + for group, occurrence in zip(to_create, created): + group.keeper_id = occurrence.pk + + _set_detection_column( + "occurrence_id", + { + pk: group.keeper_id + for group in plan.groups + for pk, previous in zip(group.detection_ids, group.previous_occurrence_ids) + if previous != group.keeper_id + }, + ) + + keeper_ids = {group.keeper_id for group in plan.groups} + keepers = Occurrence.objects.select_related("determination").in_bulk(list(keeper_ids)) + determination_before = {pk: keeper.determination_id for pk, keeper in keepers.items()} + + merged: dict[int, list[int]] = collections.defaultdict(list) + for pk, keeper_id in sorted(emptied.items()): + merged[keeper_id].append(pk) + + # Identifications and results of the emptied occurrences move onto their keepers before the delete, + # which would otherwise cascade to them. + moved_identifications: dict[int, list[tuple[int, int]]] = collections.defaultdict(list) + for pk, occurrence_id in Identification.objects.filter(occurrence_id__in=list(emptied)).values_list( + "pk", "occurrence_id" + ): + moved_identifications[emptied[occurrence_id]].append((pk, occurrence_id)) + for keeper_id, moved in moved_identifications.items(): + Identification.objects.filter(pk__in=[pk for pk, _ in moved]).update(occurrence_id=keeper_id) + withdrawn: dict[int, list[int]] = collections.defaultdict(list) + for pk, occurrence_id in Identification.objects.filter( + pk__in=_withdraw_duplicate_identifications(set(moved_identifications)) + ).values_list("pk", "occurrence_id"): + withdrawn[occurrence_id].append(pk) + + with_results = set( + AlgorithmResult.objects.filter(occurrence_id__in=list(emptied)) + .values_list("occurrence_id", flat=True) + .distinct() + ) + for keeper_id, absorbed in merged.items(): + if with_results.intersection(absorbed): + AlgorithmResult.objects.move_to_occurrence(keepers[keeper_id], absorbed) + if emptied: + Occurrence.objects.filter(pk__in=list(emptied)).delete() + + # The transaction is about to change these tables, so caching its reads only costs a cache key per query. + # Writes still invalidate the cache. Entered by hand because the context manager does not restore on error. + uncached = cachalot_disabled() + uncached.__enter__() + try: + changed = [ + keeper + for keeper in keepers.values() + if update_occurrence_determination(keeper, current_determination=keeper.determination, save=False) + ] + finally: + uncached.__exit__(None, None, None) + Occurrence.objects.bulk_update(changed, ["determination", "determination_score"], batch_size=WRITE_BATCH_SIZE) + + counters["identifications_moved"] = sum(len(moved) for moved in moved_identifications.values()) + counters["identifications_withdrawn"] = sum(len(ids) for ids in withdrawn.values()) + counters["results_recorded"] = ( + record_tracking_results( + plan.groups, keepers, plan, record_as, job, merged, moved_identifications, withdrawn, determination_before + ) + if record_as is not None + else 0 + ) + logger.info( + f"Session {plan.event.pk}: {counters['links_created']} links, created {counters['occurrences_created']} " + f"occurrences and merged {counters['occurrences_merged']} (occurrences before: " + f"{counters['occurrences_before']}, after: {counters['occurrences_after']}). Moved " + f"{counters['identifications_moved']} identification(s), recorded {counters['results_recorded']} result(s)." ) - counters["links_created"] = links - counters["transitions_too_far_apart"] = plan.transitions_too_far_apart - counters["transitions_without_dimensions"] = skipped_for_dimensions return counters @@ -386,10 +491,10 @@ def nothing_tracked_summary(skip_reasons: collections.Counter[str]) -> str: class TrackingTask(BasePostProcessingTask): - """Link detections across consecutive processed captures and fold each chain into one occurrence. + """Link detections across consecutive processed captures and merge the occurrences each chain joins. Sets each detection's ``next_detection`` link from bounding-box overlap, size and distance, - then merges every chain into a single occurrence per session. + then merges every group of linked detections into a single occurrence per session. """ key = "tracking" @@ -431,44 +536,36 @@ def _resolve_events(self) -> list[Event]: self.logger.warning(f"Tracking requested {sorted(missing)} but those sessions were not found.") return events - def _skip_reason(self, event: Event) -> str | None: - """Why this session must not be tracked, or None. Called under the session lock.""" - if self.config.require_fresh_event: - fresh, detail = event_is_fresh(event) - if not fresh: - self.logger.info(f"Skipping session {event.pk}: already tracked or edited ({detail}).") - return "it was already tracked or edited" - if ( - self.config.skip_if_human_identifications - and Occurrence.objects.filter(event=event, identifications__isnull=False).exists() - ): - self.logger.info(f"Skipping session {event.pk}: has human identifications.") - return "it has human identifications" - return None + def _check_guards(self, event: Event) -> None: + """Raise ``SkipSession`` when a guard in the settings keeps this session from being tracked.""" + if self.config.require_fresh_event and session_has_links(event): + raise SkipSession("it was already tracked") + if self.config.skip_if_human_identifications and session_has_identifications(event): + raise SkipSession("it has human identifications") def _track_session(self, event: Event, progress_cb: typing.Callable[[float], None] | None = None): """Track one session in two phases, returning ``(counters, None)`` or ``(None, skip reason)``. Matching reads only and runs outside any transaction, so the job can keep saving progress - while it works. Writing is one short transaction that locks the session, repeats the guards, - and refuses to write when the session changed since it was matched. + while it works. Writing is one short transaction that locks the session and its occurrences, + repeats the guards, and refuses to write when the session changed since it was matched. In + preview mode the plan's counts are returned and nothing is written. """ - if self._skip_reason(event) is not None: - plan = None # the locked check below reports the reason - else: + try: + self._check_guards(event) plan = plan_session_links(event, self.logger, self.config, progress_cb) - with transaction.atomic(): - lock_sessions([event.pk]) - reason = self._skip_reason(event) - if reason is not None: - return None, reason - if plan is None: - return None, "it has fewer than two processed captures" - if not plan_is_current(plan): - self.logger.warning(f"Skipping session {event.pk}: it changed while its links were being matched.") - return None, "it changed while it was being tracked" - counters = write_session_plan(plan, self.logger, record_as=self.algorithm, job=self.job) - return counters, None + if self.config.preview_only: + return preview_counts(plan), None + with transaction.atomic(): + lock_sessions([event.pk]) + lock_occurrences(plan.occurrence_ids()) + self._check_guards(event) + if not plan_is_current(plan): + raise SkipSession("it changed while it was being tracked") + return write_session_plan(plan, self.logger, record_as=self.algorithm, job=self.job), None + except SkipSession as skip: + self.logger.info(f"Skipping session {event.pk}: {skip}.") + return None, str(skip) def _throttled_progress(self, index: int, total: int) -> typing.Callable[[float], None]: """A callback that saves progress for session ``index`` of ``total`` every few transitions or seconds.""" @@ -495,9 +592,11 @@ def run(self) -> None: captures in the set, because a chain needs the captures between its detections. A session that fails is rolled back, logged and counted, and the run goes on with the next one. After the counts of the sessions that were tracked are refreshed and the metrics reported, the run raises if any - session failed, so the job is marked failed. + session failed, so the job is marked failed. In preview mode the same counts are reported and + nothing is written. """ - self.logger.info(f"Tracking starting with config: {self.config.dict()}") + preview = self.config.preview_only + self.logger.info(f"Tracking {'preview ' if preview else ''}starting with config: {self.config.dict()}") events = self._resolve_events() total = len(events) @@ -520,26 +619,30 @@ def run(self) -> None: if counters is not None: totals["tracked"] += 1 tracked_event_ids.append(event.pk) - for key in ("links_created", "occurrences_merged", "results_recorded", "transitions_too_far_apart"): - totals[key] += counters.get(key, 0) + totals.update(counters) elif reason is not None: totals["skipped"] += 1 skip_reasons[reason] += 1 # Saved between sessions, after the transaction has committed, so the bar moves per session. self.update_progress(idx / total) - # Merging occurrences changes the session and station counts, which no save refreshes. - # This already runs in a background job, so the station refresh stays inline. - update_calculated_fields_for_sessions_and_stations(tracked_event_ids, stations_async=False) + if not preview: + # Merging occurrences changes the session and station counts, which no save refreshes. + update_calculated_fields_for_sessions_and_stations(tracked_event_ids) + would = "that would be " if preview else "" metrics: dict[str, typing.Any] = { - "Sessions tracked": totals["tracked"], + "Sessions tracked" if not preview else "Sessions previewed": totals["tracked"], "Sessions skipped": totals["skipped"], "Sessions failed": len(failed_event_ids), - "Detection links created": totals["links_created"], - "Occurrences merged": totals["occurrences_merged"], - "Occurrences recorded": totals["results_recorded"], + f"Detection links {would}created": totals["links_created"], + f"Occurrences {would}merged away": totals["occurrences_merged"], + "Occurrences before": totals["occurrences_before"], + "Occurrences after" if not preview else "Occurrences after the run": totals["occurrences_after"], } + if not preview: + metrics["Occurrences recorded"] = totals["results_recorded"] + metrics["Identifications moved"] = totals["identifications_moved"] if self.config.max_capture_interval_seconds is not None: metrics["Capture pairs too far apart to compare"] = totals["transitions_too_far_apart"] # The job still succeeds when every session is skipped, so this line is written on every run: @@ -549,6 +652,12 @@ def run(self) -> None: f"Tracking failed for {len(failed_event_ids)} of {total} session(s) " f"(ids {failed_event_ids}); {totals['tracked']} were tracked." ) + elif totals["tracked"] and preview: + metrics["Result"] = ( + f"Preview only, nothing was changed. Tracking {totals['tracked']} session(s) would create " + f"{totals['links_created']} links and merge away {totals['occurrences_merged']} occurrences " + f"({totals['occurrences_before']} before, {totals['occurrences_after']} after)." + ) elif totals["tracked"]: metrics["Result"] = f"Tracked {totals['tracked']} session(s)." elif skip_reasons: diff --git a/ami/ml/results/schemas.py b/ami/ml/results/schemas.py index ed86f71c0..fccab5945 100644 --- a/ami/ml/results/schemas.py +++ b/ami/ml/results/schemas.py @@ -90,6 +90,13 @@ class TrackingResultData(DeterminationSnapshot): link_costs: list[float] = [] # Occurrences the run folded into this one. They are deleted, so these are plain ids, not references. merged_occurrence_ids: list[int] = [] + # The grouping before the run, so a reset can restore it: the occurrence's detections in capture order, + # the occurrence each one was in, each identification moved here as (identification, earlier occurrence), + # and the identifications withdrawn because their user had another active one on a merged occurrence. + detection_ids: list[int] = [] + previous_occurrence_ids: list[int | None] = [] + moved_identifications: list[tuple[int, int]] = [] + withdrawn_identification_ids: list[int] = [] ALGORITHM_RESULT_DATA_MODELS: tuple[type[AlgorithmResultData], ...] = ( diff --git a/ui/src/data-services/models/occurrence-history.ts b/ui/src/data-services/models/occurrence-history.ts index 1da8236ee..e05ed0f29 100644 --- a/ui/src/data-services/models/occurrence-history.ts +++ b/ui/src/data-services/models/occurrence-history.ts @@ -97,6 +97,12 @@ export interface TrackingResultData extends ServerDeterminationSnapshot { path_length: number /** The largest box area over the smallest. */ size_change: number + /** The grouping before the run: the detections in capture order and the occurrence each was in. */ + detection_ids?: number[] + previous_occurrence_ids?: (number | null)[] + /** Identifications moved here, as [identification id, earlier occurrence id]. */ + moved_identifications?: [number, number][] + withdrawn_identification_ids?: number[] } export interface ServerIdentificationDetails { From 049ce4b24b1bf1a341b170026875c880f98ea1fd Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Tue, 6 Oct 2026 18:42:24 -0700 Subject: [PATCH 26/28] fix(tracking): skip the regroup split where nothing was tracked, and refresh stations inline Every regroup searched the sessions it touched for occurrences spanning a new session boundary. On the largest stations of a production copy that meant literal id lists of 600,000 to 800,000 occurrences and about 1.5 to 2 s per regroup, although no occurrence spanned sessions. Tracking is what merges detections of several captures into one occurrence, so the search now runs only when a detection of those sessions has a tracking link, which one indexed query answers. An occurrence grouped some other way, without links, is no longer split; a test pins that. The station count refresh after tracking ran through a background task that nothing used in the background, behind a flag only tracking set. The helper now refreshes the stations inline and the unused task is removed. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01QT59KePky4u4nbCsTAggtc --- ami/main/models.py | 29 ++++++++++++----------------- ami/main/tasks.py | 9 --------- ami/main/tests.py | 13 +++++++++++-- 3 files changed, 23 insertions(+), 28 deletions(-) diff --git a/ami/main/models.py b/ami/main/models.py index c777afac8..309b3ffc6 100644 --- a/ami/main/models.py +++ b/ami/main/models.py @@ -1441,26 +1441,18 @@ def update_calculated_fields_for_events( return to_update -def update_calculated_fields_for_sessions_and_stations( - event_ids: typing.Iterable[int | None], stations_async: bool = True -) -> None: +def update_calculated_fields_for_sessions_and_stations(event_ids: typing.Iterable[int | None]) -> None: """Refresh the cached counts of these sessions and of the stations they belong to. - Call once after occurrences are created, merged or split, which neither the - occurrence nor the detection saves do. The station refresh scans the whole station, so by - default it runs in a background task after the transaction commits. + Call once after occurrences are created, merged or split, which neither the occurrence nor the + detection saves do. The station refresh scans each whole station, so call it from a background job. """ - from ami.main.tasks import refresh_deployment_cached_counts - pks = sorted({pk for pk in event_ids if pk is not None}) if not pks: return update_calculated_fields_for_events(pks=pks) - deployment_ids = list(Deployment.objects.filter(events__pk__in=pks).values_list("pk", flat=True).distinct()) - if stations_async: - transaction.on_commit(lambda: refresh_deployment_cached_counts.delay(deployment_ids)) - else: - refresh_deployment_cached_counts(deployment_ids) + for deployment in Deployment.objects.filter(events__pk__in=pks).distinct(): + deployment.update_calculated_fields(save=True) def audit_event_lengths(deployment: Deployment): @@ -1769,13 +1761,16 @@ def _split_occurrences_at_session_boundaries(job: "Job | None", event_pks: set[i """Split every occurrence whose detections now span several sessions, among the sessions a regroup touched. An occurrence is expected to belong to one session, so a regroup that draws a session - boundary through it leaves one piece per session. Only an occurrence with a detection on a capture of - ``event_pks`` can have been cut, so the search starts from those captures. The ids are looked up - in steps with literal id lists: left as a subquery, Postgres scans the whole detection table - (measured on a copy of production data). Returns how many occurrences were split. + boundary through it leaves one piece per session. Tracking is what merges detections of several + captures into one occurrence, so when no detection of these sessions has a tracking link the search is + skipped after one indexed query; an occurrence grouped some other way, with no links, is then not split. + The search itself reads every occurrence of the touched sessions. Returns how many occurrences were split. """ from ami.ml.post_processing.tracking.sessions import lock_sessions, split_at_session_boundaries + if not Detection.objects.filter(source_image__event_id__in=event_pks, next_detection__isnull=False).exists(): + return 0 + def find_spanning_ids() -> list[int]: capture_ids = list(SourceImage.objects.filter(event_id__in=event_pks).values_list("pk", flat=True)) touched_occurrence_ids = list( diff --git a/ami/main/tasks.py b/ami/main/tasks.py index 19ea81321..16f927a3f 100644 --- a/ami/main/tasks.py +++ b/ami/main/tasks.py @@ -23,12 +23,3 @@ def refresh_project_cached_counts(project_id: int) -> None: logger.info(f"Refreshing cached counts for project {project.pk} ({project.name})") project.update_related_calculated_fields() - - -@celery_app.task(ignore_result=True) -def refresh_deployment_cached_counts(deployment_ids: list[int]) -> None: - """Refresh the cached counts of these stations after occurrences were created, merged or split.""" - from ami.main.models import Deployment - - for deployment in Deployment.objects.filter(pk__in=deployment_ids): - deployment.update_calculated_fields(save=True) diff --git a/ami/main/tests.py b/ami/main/tests.py index b086ac67a..560b2b8ba 100644 --- a/ami/main/tests.py +++ b/ami/main/tests.py @@ -8572,7 +8572,7 @@ def _group(self, gap_hours: int) -> list[Event]: capture.refresh_from_db() return list(Event.objects.filter(deployment=self.deployment).order_by("start")) - def _make_occurrence(self, captures: list[SourceImage]) -> tuple[Occurrence, list[Detection]]: + def _make_occurrence(self, captures: list[SourceImage], linked: bool = True) -> tuple[Occurrence, list[Detection]]: occurrence = Occurrence.objects.create( event=captures[0].event, deployment=self.deployment, project=self.project ) @@ -8583,12 +8583,21 @@ def _make_occurrence(self, captures: list[SourceImage]) -> tuple[Occurrence, lis ) detection.classifications.create(taxon=self.taxon, score=0.9, timestamp=capture.timestamp) detections.append(detection) - for earlier, later in zip(detections, detections[1:]): + for earlier, later in zip(detections, detections[1:] if linked else []): earlier.next_detection = later earlier.save(update_fields=["next_detection"]) occurrence.save() return occurrence, detections + def test_sessions_without_tracking_links_are_not_searched(self): + """Only tracking merges detections across captures, so a regroup with no links skips the search.""" + self._group(gap_hours=6) + occurrence, detections = self._make_occurrence(self.captures, linked=False) + + self._group(gap_hours=2) + + self.assertEqual(self._detection_ids(occurrence), [d.pk for d in detections]) + def _split_one_occurrence(self) -> tuple[Occurrence, Occurrence, list[Detection], list[Event]]: self._group(gap_hours=6) occurrence, detections = self._make_occurrence(self.captures) From d4bb07b302b30d796a75f53e5b7a735c28ec1757 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Tue, 6 Oct 2026 18:42:28 -0700 Subject: [PATCH 27/28] fix(migrations): give up after 10 s rather than queue behind a long query on detections Adding the next_detection column, attaching its unique constraint and adding its foreign key each take a brief strong lock on the detection table. With no lock timeout, a long query already reading the table, such as an export, would make the statement wait, and every other query on the table would wait behind it. The migrations now set a 10 s lock timeout around those statements, so they fail and can be rerun instead. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01QT59KePky4u4nbCsTAggtc --- ami/main/migrations/0100_detection_next_detection.py | 3 +++ .../0101_detection_next_detection_constraints.py | 10 +++++++++- 2 files changed, 12 insertions(+), 1 deletion(-) diff --git a/ami/main/migrations/0100_detection_next_detection.py b/ami/main/migrations/0100_detection_next_detection.py index 825935ee0..960753ed7 100644 --- a/ami/main/migrations/0100_detection_next_detection.py +++ b/ami/main/migrations/0100_detection_next_detection.py @@ -17,6 +17,9 @@ class Migration(migrations.Migration): ] operations = [ + # Adding the column takes a brief exclusive lock. Give up rather than queue behind a long query on + # the table, which would block every other query until it ends; rerun the migration if it gives up. + migrations.RunSQL(sql="SET LOCAL lock_timeout = '10s';", reverse_sql=migrations.RunSQL.noop), migrations.SeparateDatabaseAndState( state_operations=[ migrations.AddField( diff --git a/ami/main/migrations/0101_detection_next_detection_constraints.py b/ami/main/migrations/0101_detection_next_detection_constraints.py index b0673903b..6da28691f 100644 --- a/ami/main/migrations/0101_detection_next_detection_constraints.py +++ b/ami/main/migrations/0101_detection_next_detection_constraints.py @@ -11,7 +11,8 @@ class Migration(migrations.Migration): light lock. See 0093 for why the statement timeout is cleared and restored around the build. The constraint names match the ones Django generates, so later AlterField migrations find them. - If the index build is interrupted it leaves an invalid index of the same name; drop it before retrying. + If the index build is interrupted it leaves an invalid index of the same name, and if a lock times out the + index already exists; drop it before retrying. """ atomic = False @@ -32,6 +33,9 @@ class Migration(migrations.Migration): ), reverse_sql=migrations.RunSQL.noop, ), + # The two ALTER TABLE statements below take brief strong locks. Give up rather than queue behind a long + # query on the table, which would block every other query until it ends. + migrations.RunSQL(sql="SET lock_timeout = '10s';", reverse_sql=migrations.RunSQL.noop), migrations.RunSQL( sql=( 'ALTER TABLE "main_detection" ADD CONSTRAINT "main_detection_next_detection_id_key" ' @@ -57,6 +61,10 @@ class Migration(migrations.Migration): ), reverse_sql=migrations.RunSQL.noop, ), + migrations.RunSQL( + sql="RESET lock_timeout;", + reverse_sql=migrations.RunSQL.noop, + ), migrations.RunSQL( sql="RESET statement_timeout;", reverse_sql=migrations.RunSQL.noop, From 8b387eb41fd26fc7e3cf64bef601f3d2444b4ff7 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Tue, 6 Oct 2026 18:42:28 -0700 Subject: [PATCH 28/28] refactor(tracking): reuse the admin's error mapping and drop the card's tracking special case The Events admin action mapped settings errors onto form fields with its own copy of the post-processing admin's helper; it now calls that helper. The history card showed "detections affected" for every kind except tracking; it now shows the figure whenever the result has classifications to count, which is what the exception stood for. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01QT59KePky4u4nbCsTAggtc --- ami/ml/post_processing/admin/tracking_actions.py | 7 ++----- .../identification-card/algorithm-result.tsx | 2 +- 2 files changed, 3 insertions(+), 6 deletions(-) diff --git a/ami/ml/post_processing/admin/tracking_actions.py b/ami/ml/post_processing/admin/tracking_actions.py index 61966b17b..f1962ab95 100644 --- a/ami/ml/post_processing/admin/tracking_actions.py +++ b/ami/ml/post_processing/admin/tracking_actions.py @@ -13,7 +13,7 @@ from django.db import transaction from ami.jobs.models import Job -from ami.ml.post_processing.admin.actions import ConfigValidationErrors +from ami.ml.post_processing.admin.actions import ConfigValidationErrors, _schema_errors_to_form_fields from ami.ml.post_processing.base import BasePostProcessingTask @@ -43,10 +43,7 @@ def build_tracking_jobs_for_events( try: validated.append((project_id, task_cls.config_schema(**{**config, "event_ids": sorted(event_ids)}))) except pydantic.ValidationError as exc: - for err in exc.errors(): - loc = err.get("loc") or () - field = str(loc[0]) if loc and str(loc[0]) in form_field_names else None - errors.append((field, err.get("msg", "Invalid value"))) + errors.extend(_schema_errors_to_form_fields(exc, form_field_names)) if errors: raise ConfigValidationErrors(errors) diff --git a/ui/src/pages/occurrence-details/identification-card/algorithm-result.tsx b/ui/src/pages/occurrence-details/identification-card/algorithm-result.tsx index d40ea91fa..661271ee2 100644 --- a/ui/src/pages/occurrence-details/identification-card/algorithm-result.tsx +++ b/ui/src/pages/occurrence-details/identification-card/algorithm-result.tsx @@ -198,7 +198,7 @@ export const AlgorithmResult = ({ break } } - if (entry.kind !== 'tracking') { + if (entry.classifications.length) { stats.push({ label: translate(STRING.HISTORY_DETECTIONS_AFFECTED), value: new Set(entry.classifications.map((c) => c.detection_id)).size,