diff --git a/ami/main/admin.py b/ami/main/admin.py index 47b383a00..2c4aa2543 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 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/main/migrations/0100_detection_next_detection.py b/ami/main/migrations/0100_detection_next_detection.py new file mode 100644 index 000000000..960753ed7 --- /dev/null +++ b/ami/main/migrations/0100_detection_next_detection.py @@ -0,0 +1,45 @@ +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 = [ + # 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( + 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..6da28691f --- /dev/null +++ b/ami/main/migrations/0101_detection_next_detection_constraints.py @@ -0,0 +1,72 @@ +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, and if a lock times out the + index already exists; 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, + ), + # 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" ' + '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 lock_timeout;", + reverse_sql=migrations.RunSQL.noop, + ), + migrations.RunSQL( + sql="RESET statement_timeout;", + reverse_sql=migrations.RunSQL.noop, + ), + ] diff --git a/ami/main/models.py b/ami/main/models.py index a1330e88a..309b3ffc6 100644 --- a/ami/main/models.py +++ b/ami/main/models.py @@ -1441,6 +1441,20 @@ def update_calculated_fields_for_events( return to_update +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 each whole station, so call it from a background job. + """ + pks = sorted({pk for pk in event_ids if pk is not None}) + if not pks: + return + update_calculated_fields_for_events(pks=pks) + for deployment in Deployment.objects.filter(events__pk__in=pks).distinct(): + deployment.update_calculated_fields(save=True) + + def audit_event_lengths(deployment: Deployment): logger.info("Checking for unusual event durations") @@ -1648,6 +1662,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(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 # time (Detection.associate_new_occurrence and Pipeline.save_results both @@ -1728,6 +1744,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 +1757,62 @@ def _group_images_into_events_locked( return events +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. 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( + 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 + # 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 + + def deployment_events_need_update(deployment: Deployment) -> bool: """ Returns True if there are any SourceImages in the deployment @@ -3239,6 +3312,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 @@ -3896,13 +3978,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 diff --git a/ami/main/tests.py b/ami/main/tests.py index ef58ce8d9..560b2b8ba 100644 --- a/ami/main/tests.py +++ b/ami/main/tests.py @@ -8538,3 +8538,157 @@ 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. + """ + + @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) + cls.captures = [ + SourceImage.objects.create( + deployment=cls.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], linked: bool = True) -> 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:] 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) + 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) + + 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) + + 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, occurrence.pk) + self.assertFalse(AlgorithmResult.objects.filter(occurrence=piece).exists()) 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/__init__.py b/ami/ml/post_processing/__init__.py index c94be9ae9..2f45726b1 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 # noqa: F401 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)) 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..f1962ab95 --- /dev/null +++ b/ami/ml/post_processing/admin/tracking_actions.py @@ -0,0 +1,62 @@ +"""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, _schema_errors_to_form_fields +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: + errors.extend(_schema_errors_to_form_fields(exc, form_field_names)) + + 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..a1c39dc43 --- /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 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..0b9156f6e 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 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..b1d9f69e9 --- /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 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_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_matching.py b/ami/ml/post_processing/tests/test_tracking_matching.py new file mode 100644 index 000000000..b74605bb9 --- /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]) + + # 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)]) + + +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_stats.py b/ami/ml/post_processing/tests/test_tracking_stats.py new file mode 100644 index 000000000..f102a754f --- /dev/null +++ b/ami/ml/post_processing/tests/test_tracking_stats.py @@ -0,0 +1,71 @@ +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_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.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) + + +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 TestSizeChange(SimpleTestCase): + def test_is_largest_area_over_smallest(self): + 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_change([[0, 0, 0, 0], [0, 0, 10, 10]]), 100.0) + + def test_is_one_without_boxes(self): + self.assertEqual(stats.size_change([]), 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.label_agreement([1, 1, 2, 3], 1), 0.5) + + def test_agreement_is_none_without_labels(self): + self.assertIsNone(stats.label_agreement([], 1)) + + def test_missing_taxa_never_agree_with_a_missing_determination(self): + self.assertEqual(stats.label_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, 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 new file mode 100644 index 000000000..7833a4e70 --- /dev/null +++ b/ami/ml/post_processing/tests/test_tracking_task.py @@ -0,0 +1,641 @@ +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, + 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 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 + +logger = logging.getLogger(__name__) + +BOX = [100, 100, 200, 200] + + +class _TrackingCase(TestCase): + @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) + 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 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]) + event = captures[0].event + self.assertEqual(Occurrence.objects.filter(event=event).count(), 3) + + self.run_task(event) + + self.assertEqual(self.occurrence_sizes(event), [3]) + 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_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"]) + + # 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, require_fresh_event=False) + + 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): + # 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[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.run_task(event, require_fresh_event=False) + self.assertEqual(self.occurrence_sizes(event), [3]) + + 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 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]) + 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), [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]) + 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}) + # 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]) + 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_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, 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()) + + # 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) + + occurrence.refresh_from_db() + self.assertEqual(occurrence.determination, self.taxa[2]) + + +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(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), (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) + self.assertEqual(result.data["distinct_taxa"], 1) + 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() + 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_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_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) + capture.detections.get().occurrence.save() + job = self.make_job() + + task = self.run_task(captures[0].event, job=job) + + result = AlgorithmResult.objects.get(kind="tracking") + self.assertEqual(result.data["determination_after_id"], self.taxa[1].pk) + 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]) + 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, first.pk) + 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]) + 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).") + + +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.write_session_plan + + 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=failing).update(determination_score=0.123) + raise ValueError("boom") + return real(plan, *args, **kwargs) + + 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() + + 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 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, captures_count: int, insects: int, start: datetime.datetime) -> int: + from cachalot.api import cachalot_disabled + + 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: + 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), [captures_count] * insects) + return len(queries) + + 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/__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/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 new file mode 100644 index 000000000..c30ddf02d --- /dev/null +++ b/ami/ml/post_processing/tracking/config.py @@ -0,0 +1,155 @@ +"""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 + +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." +) + + +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 = reference( + "capture_set", + 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] = 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, + 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 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." + ), + ) + + @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/ml/post_processing/tracking/sessions.py b/ami/ml/post_processing/tracking/sessions.py new file mode 100644 index 000000000..fd906bd47 --- /dev/null +++ b/ami/ml/post_processing/tracking/sessions.py @@ -0,0 +1,113 @@ +"""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 +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 + +# Order of detections within an occurrence: capture time, then capture, then detection. +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. 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]) + 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) + 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/ml/post_processing/tracking/stats.py b/ami/ml/post_processing/tracking/stats.py new file mode 100644 index 000000000..2461cc23b --- /dev/null +++ b/ami/ml/post_processing/tracking/stats.py @@ -0,0 +1,101 @@ +"""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 + path_length: float + size_change: float + distinct_taxa: int + label_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 _raw_path_length(boxes: Sequence[BBox]) -> float: + centres = [((box[0] + box[2]) / 2, (box[1] + box[3]) / 2) for box in boxes] + 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. + + 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 + 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 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) + + +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.""" + diagonal = frame_diagonal(sizes, boxes) + return OccurrenceFigures( + detection_count=len(boxes), + motion=motion(boxes, diagonal), + path_length=path_length(boxes, diagonal), + size_change=size_change(boxes), + distinct_taxa=distinct_taxa(labels), + 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 new file mode 100644 index 000000000..6d5058fcc --- /dev/null +++ b/ami/ml/post_processing/tracking/task.py @@ -0,0 +1,672 @@ +"""The tracking post-processing task: links detections in consecutive captures and merges the occurrences they join. + +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 +import dataclasses +import logging +import time +import typing +from collections.abc import Iterator, Sequence + +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 ( + Classification, + Detection, + Event, + Identification, + Occurrence, + SourceImage, + SourceImageCollection, + update_calculated_fields_for_sessions_and_stations, + 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 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 +from .stats import occurrence_figures + +if typing.TYPE_CHECKING: + 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 +# Rows per statement when occurrence determinations are written in bulk. +WRITE_BATCH_SIZE = 1000 + +Link = tuple[int, int, float] + + +class SkipSession(Exception): + """Raised with the reason a session is not tracked; the job counts sessions by reason.""" + + +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() + + +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 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, + and so are captures with no timestamp, whose place in the sequence is unknown. + """ + return list( + SourceImage.objects.filter(event=event, timestamp__isnull=False) + .filter(Exists(Detection.objects.filter(source_image=OuterRef("pk")))) + .order_by("timestamp", "pk") + ) + + +@dataclasses.dataclass +class SessionPlan: + """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]] + 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, list[Detection]], + config: TrackingConfig, + logger: logging.Logger, +) -> 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. + """ + 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 [] + 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 + links = select_links( + [(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 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( + event: Event, + logger: logging.Logger, + config: TrackingConfig, + progress_cb: typing.Callable[[float], None] | None = None, +) -> SessionPlan: + """Match one session's processed captures and work out the merges, writing nothing and holding no lock. + + 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: + 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 + links: list[Link] = [] + without_dimensions = 0 + for i, proposed in enumerate(iter_transition_links(source_images, detections_by_capture, config, logger)): + 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) + ) + + 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 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 = { + 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 merges, writing each table in a few bulk statements. + + 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. + """ + 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)." + ) + 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 merge the occurrences each chain joins. + + Sets each detection's ``next_detection`` link from bounding-box overlap, size and distance, + then merges every group of linked detections into a single occurrence per session. + """ + + key = "tracking" + name = "Occurrence tracking" + config_schema = TrackingConfig + result_models = (TrackingResultData,) + + 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 _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 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. + """ + try: + self._check_guards(event) + plan = plan_session_links(event, self.logger, self.config, progress_cb) + 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.""" + 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, 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 + 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. In preview mode the same counts are reported and + nothing is written. + """ + 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) + self.logger.info(f"Tracking: {total} session(s) in scope") + + 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})") + try: + 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) + counters, reason = None, None + if counters is not None: + totals["tracked"] += 1 + tracked_event_ids.append(event.pk) + 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) + + 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" if not preview else "Sessions previewed": totals["tracked"], + "Sessions skipped": totals["skipped"], + "Sessions failed": len(failed_event_ids), + 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: + # a retry keeps text params, and a stale line would contradict the counts. + 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"] 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: + 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)}") + if failed_event_ids: + raise RuntimeError(metrics["Result"]) diff --git a/ami/ml/results/schemas.py b/ami/ml/results/schemas.py index f4be31b83..fccab5945 100644 --- a/ami/ml/results/schemas.py +++ b/ami/ml/results/schemas.py @@ -68,7 +68,42 @@ 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" + value_field: ClassVar[str | None] = "motion" + + detection_count: int + # Mean distance per step between consecutive detection centres, as a fraction of the image diagonal; + # 0 for one detection. + motion: float + # 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. + 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] = [] + # 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], ...] = ( + ClassMaskingResultData, + SizeFilterResultData, + TrackingResultData, +) ALGORITHM_RESULT_DATA_SCHEMAS: dict[str, type[AlgorithmResultData]] = { model.kind: model for model in ALGORITHM_RESULT_DATA_MODELS 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 diff --git a/ui/src/data-services/models/occurrence-history.test.ts b/ui/src/data-services/models/occurrence-history.test.ts index be8ba5440..a746d7012 100644 --- a/ui/src/data-services/models/occurrence-history.test.ts +++ b/ui/src/data-services/models/occurrence-history.test.ts @@ -111,6 +111,32 @@ 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: {}, + label_agreement: null, + link_costs: [0.12, 0.2], + merged_occurrence_ids: [11, 12], + motion: 0.0071, + path_length: 0.0141, + size_change: 1.15, + }, + determination_after: NOCTUA, + determination_before: NOCTUA, + id: 6, + kind: 'tracking', + score: null, + type: 'algorithm_result', + value: 0.0071, +} + const ownIdentification: HumanIdentification = { comment: '', createdAt: '2026-04-29T22:00:00', @@ -150,6 +176,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 +247,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..e05ed0f29 100644 --- a/ui/src/data-services/models/occurrence-history.ts +++ b/ui/src/data-services/models/occurrence-history.ts @@ -81,6 +81,30 @@ export interface SizeFilterResultData extends ServerDeterminationSnapshot { relative_size: number } +export interface TrackingResultData extends ServerDeterminationSnapshot { + detection_count: number + /** 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 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_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 { agreed_with_identification_id: number | null agreed_with_prediction_id: number | null @@ -131,9 +155,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 +197,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..661271ee2 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,56 @@ 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_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_change * 100) / 100}`, + }), + }, + { + label: translate(STRING.HISTORY_TRACKING_TAXA), + value: data.distinct_taxa, + }, + { + label: translate(STRING.HISTORY_TRACKING_LABEL_AGREEMENT), + value: + data.label_agreement !== null + ? formatPercent(data.label_agreement) + : translate(STRING.VALUE_NOT_AVAILABLE), + }, + { + label: translate(STRING.HISTORY_TRACKING_MERGED), + value: data.merged_occurrence_ids.length, + } + ) + break + } + } + if (entry.classifications.length) { + 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..88d2e1d69 100644 --- a/ui/src/utils/language.ts +++ b/ui/src/utils/language.ts @@ -337,6 +337,16 @@ export enum STRING { HISTORY_RECORD_ID, HISTORY_SIZE_FILTER, HISTORY_SUPERSEDED_BY, + HISTORY_TRACKING, + 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, ID_APPLIED, INFO, INTERMEDIATE_CLASSIFICATION, @@ -794,6 +804,18 @@ 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_LABEL_AGREEMENT]: 'Label agreement', + [STRING.HISTORY_TRACKING_DETECTIONS]: 'Detections', + [STRING.HISTORY_TRACKING_MERGED]: 'Occurrences merged in', + [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', + [STRING.HISTORY_TRACKING_TAXA]: 'Taxa', [STRING.ID_APPLIED]: 'ID applied', [STRING.INFO]: 'Info', [STRING.INTERMEDIATE_CLASSIFICATION]: 'Intermediate classification',