From 0762b4b3fb2fa3b73de85e801ca87cfa897f4c50 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Thu, 1 Oct 2026 19:01:49 -0700 Subject: [PATCH 1/5] feat: move confirmed occurrences, identifications and vectors between Antenna instances Four management commands and the modules behind them. export_validated_occurrences writes the occurrences people confirmed or identified in a project to a versioned JSON bundle keyed by natural keys only (capture path, timestamp, station, box, detector; taxon by GBIF key or name and rank; user by email). import_validated_occurrences replays a bundle onto another project: it finds each detection again (exact box, else IoU >= 0.7), regroups confirmed occurrences with the track-edit operations the review interface uses, records each confirmation under its original reviewer and time, and re-creates the identifications with their user and timestamp. Partially found occurrences are reported and never confirmed; a second run changes nothing. export_embeddings and import_embeddings do the same for detection feature vectors (vectors.npy plus an index of detection keys), writing DetectionEmbedding rows where that table exists. The confirmation write is isolated in confirm_grouping_as_of so it can become a ValidationReview row in a few lines once the model-outputs schema lands. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01C7Xf6VPbwWtTumhjjF15g8 --- .../management/commands/export_embeddings.py | 55 ++ .../commands/export_validated_occurrences.py | 42 ++ .../management/commands/import_embeddings.py | 74 ++ .../commands/import_validated_occurrences.py | 95 +++ ami/main/models_future/detection_matching.py | 274 +++++++ ami/main/models_future/embedding_transfer.py | 288 ++++++++ .../models_future/validated_occurrences.py | 683 ++++++++++++++++++ ami/main/test_validated_occurrences.py | 326 +++++++++ 8 files changed, 1837 insertions(+) create mode 100644 ami/main/management/commands/export_embeddings.py create mode 100644 ami/main/management/commands/export_validated_occurrences.py create mode 100644 ami/main/management/commands/import_embeddings.py create mode 100644 ami/main/management/commands/import_validated_occurrences.py create mode 100644 ami/main/models_future/detection_matching.py create mode 100644 ami/main/models_future/embedding_transfer.py create mode 100644 ami/main/models_future/validated_occurrences.py create mode 100644 ami/main/test_validated_occurrences.py diff --git a/ami/main/management/commands/export_embeddings.py b/ami/main/management/commands/export_embeddings.py new file mode 100644 index 000000000..848c01a92 --- /dev/null +++ b/ami/main/management/commands/export_embeddings.py @@ -0,0 +1,55 @@ +""" +Write one algorithm's detection feature vectors for a project to a directory. + +The vectors go to ``vectors.npy`` and each row's detection is named in ``index.csv`` by +its natural key (capture path, timestamp, station, box, detector), so +``import_embeddings`` can put them back on a database where the ids differ. +""" + +import pathlib + +from django.core.management.base import BaseCommand, CommandError + +from ami.main.models import Project +from ami.main.models_future.embedding_transfer import DEFAULT_VECTOR_KEY, export_embeddings, resolve_algorithm +from ami.ml.models import Algorithm + + +class Command(BaseCommand): + help = "Export detection feature vectors from one algorithm as vectors.npy plus an index of detection keys." + + def add_arguments(self, parser): + parser.add_argument("--project", type=int, required=True, help="Project ID to export from.") + parser.add_argument( + "--algorithm", required=True, help="Key, name or ID of the algorithm whose vectors to export." + ) + parser.add_argument("--output", type=pathlib.Path, required=True, help="Directory to write into.") + parser.add_argument( + "--vector-key", + default=DEFAULT_VECTOR_KEY, + help=f"Which of the algorithm's vectors to export when it stores several (default {DEFAULT_VECTOR_KEY}).", + ) + parser.add_argument( + "--dtype", default="float32", choices=["float16", "float32"], help="Storage precision (default float32)." + ) + + def handle(self, *args, **options): + try: + project = Project.objects.get(pk=options["project"]) + except Project.DoesNotExist as err: + raise CommandError(f"Project {options['project']} does not exist") from err + try: + algorithm = resolve_algorithm(options["algorithm"]) + except Algorithm.DoesNotExist as err: + raise CommandError(str(err)) from err + + try: + manifest = export_embeddings( + project, algorithm, options["output"], vector_key=options["vector_key"], dtype=options["dtype"] + ) + except ValueError as err: + raise CommandError(str(err)) from err + self.stdout.write( + f"Wrote {manifest.count} vectors of {manifest.dimensions} dimensions ({manifest.dtype}) from " + f"{manifest.source_store} for algorithm {algorithm.key!r} in project #{project.pk} to {options['output']}" + ) diff --git a/ami/main/management/commands/export_validated_occurrences.py b/ami/main/management/commands/export_validated_occurrences.py new file mode 100644 index 000000000..eb783d8c2 --- /dev/null +++ b/ami/main/management/commands/export_validated_occurrences.py @@ -0,0 +1,42 @@ +""" +Write the occurrences people confirmed or identified in one project to a portable bundle. + +The bundle names detections, taxa and users by natural keys only (capture path, timestamp, +station, bounding box, detector; GBIF key or name and rank; email), so it can be replayed +onto a database where every id is different with ``import_validated_occurrences``. +""" + +import json +import pathlib + +from django.core.management.base import BaseCommand, CommandError + +from ami.main.models import Project +from ami.main.models_future.validated_occurrences import build_bundle + + +class Command(BaseCommand): + help = "Export a project's confirmed occurrences and identifications as a bundle keyed by natural keys." + + def add_arguments(self, parser): + parser.add_argument("--project", type=int, required=True, help="Project ID to export from.") + parser.add_argument("--output", type=pathlib.Path, required=True, help="Path of the JSON bundle to write.") + + def handle(self, *args, **options): + try: + project = Project.objects.get(pk=options["project"]) + except Project.DoesNotExist as err: + raise CommandError(f"Project {options['project']} does not exist") from err + + bundle = build_bundle(project) + output: pathlib.Path = options["output"] + output.parent.mkdir(parents=True, exist_ok=True) + output.write_text(json.dumps(bundle.as_dict(), indent=1, ensure_ascii=False)) + + confirmed = sum(1 for record in bundle.occurrences if record.confirmation) + identifications = sum(len(record.identifications) for record in bundle.occurrences) + detections = sum(len(record.detections) for record in bundle.occurrences) + self.stdout.write( + f"Wrote {len(bundle.occurrences)} occurrences ({confirmed} confirmed, {identifications} identifications, " + f"{detections} detections) from project #{project.pk} to {output}" + ) diff --git a/ami/main/management/commands/import_embeddings.py b/ami/main/management/commands/import_embeddings.py new file mode 100644 index 000000000..47877b9c4 --- /dev/null +++ b/ami/main/management/commands/import_embeddings.py @@ -0,0 +1,74 @@ +""" +Put detection feature vectors from an ``export_embeddings`` directory onto a project. + +Each row's detection is found again by its natural key (capture path and box, with an +overlap fallback) and the vector is stored under the same algorithm. By default this is a +dry run that reports how many detections were found; ``--execute`` writes the rows. A +detection that already has a vector from that algorithm is skipped unless ``--replace``. +""" + +import pathlib + +from django.core.management.base import BaseCommand, CommandError + +from ami.main.models import Project +from ami.main.models_future.detection_matching import DEFAULT_IOU_THRESHOLD +from ami.main.models_future.embedding_transfer import import_embeddings, resolve_algorithm +from ami.ml.models import Algorithm + + +class Command(BaseCommand): + help = "Import detection feature vectors from an export directory onto a project (dry run by default)." + + def add_arguments(self, parser): + parser.add_argument("--project", type=int, required=True, help="Project ID to import into.") + parser.add_argument( + "--input", type=pathlib.Path, required=True, help="Directory written by export_embeddings." + ) + parser.add_argument("--execute", action="store_true", default=False, help="Write the rows.") + parser.add_argument( + "--algorithm", + default=None, + help="Key, name or ID of the algorithm to store the vectors under (default: the one in the manifest).", + ) + parser.add_argument("--iou-threshold", type=float, default=DEFAULT_IOU_THRESHOLD) + parser.add_argument( + "--replace", action="store_true", default=False, help="Overwrite vectors the detections already have." + ) + + def handle(self, *args, **options): + try: + project = Project.objects.get(pk=options["project"]) + except Project.DoesNotExist as err: + raise CommandError(f"Project {options['project']} does not exist") from err + algorithm = None + if options["algorithm"]: + try: + algorithm = resolve_algorithm(options["algorithm"]) + except Algorithm.DoesNotExist as err: + raise CommandError(str(err)) from err + + try: + report = import_embeddings( + project, + options["input"], + execute=options["execute"], + iou_threshold=options["iou_threshold"], + replace=options["replace"], + algorithm=algorithm, + ) + except (OSError, ValueError, RuntimeError, Algorithm.DoesNotExist) as err: + raise CommandError(str(err)) from err + + summary = report.summary() + mode = "Applied" if options["execute"] else "Dry run" + self.stdout.write( + f"{mode} on project #{project.pk}: {summary['vectors_total']} vectors, detections {summary['detections']}" + ) + if options["execute"]: + self.stdout.write( + f" written: {summary['written']} skipped (already stored): {summary['skipped_existing']}" + f" replaced: {summary['replaced']}" + ) + else: + self.stdout.write("Nothing was changed. Re-run with --execute to write the vectors.") diff --git a/ami/main/management/commands/import_validated_occurrences.py b/ami/main/management/commands/import_validated_occurrences.py new file mode 100644 index 000000000..5224f066a --- /dev/null +++ b/ami/main/management/commands/import_validated_occurrences.py @@ -0,0 +1,95 @@ +""" +Replay a bundle of confirmed occurrences and identifications onto a project. + +By default this is a dry run: it finds each detection again and reports how many matched +exactly, how many only by overlap, and how many are missing, without changing a row. With +``--execute`` it regroups the matched detections with the same operations the review +interface uses, records each confirmation under its original reviewer and time, and +re-creates the identifications. Occurrences with a missing detection are reported as +partial and left unconfirmed. Running it twice changes nothing the second time. +""" + +import json +import pathlib + +from django.core.management.base import BaseCommand, CommandError + +from ami.main.models import Project +from ami.main.models_future.detection_matching import DEFAULT_IOU_THRESHOLD +from ami.main.models_future.validated_occurrences import Bundle, ImportOptions, import_bundle + + +class Command(BaseCommand): + help = "Replay confirmed occurrences and identifications from a bundle onto a project (dry run by default)." + + def add_arguments(self, parser): + parser.add_argument("--project", type=int, required=True, help="Project ID to import into.") + parser.add_argument("--input", type=pathlib.Path, required=True, help="Path of the JSON bundle to read.") + parser.add_argument( + "--execute", + action="store_true", + default=False, + help="Write the changes. Without this flag the command only reports what it would do.", + ) + parser.add_argument( + "--iou-threshold", + type=float, + default=DEFAULT_IOU_THRESHOLD, + help=( + "Lowest intersection over union accepted when no detection has the exact box " + f"(default {DEFAULT_IOU_THRESHOLD})." + ), + ) + parser.add_argument( + "--create-missing-detections", + action="store_true", + default=False, + help="Recreate a box whose capture exists but that no detector found, as a reviewer-drawn detection.", + ) + parser.add_argument( + "--report", type=pathlib.Path, default=None, help="Write the per-occurrence match report as JSON here." + ) + + def handle(self, *args, **options): + try: + project = Project.objects.get(pk=options["project"]) + except Project.DoesNotExist as err: + raise CommandError(f"Project {options['project']} does not exist") from err + try: + bundle = Bundle.from_dict(json.loads(options["input"].read_text())) + except (OSError, ValueError, KeyError) as err: + raise CommandError(f"Could not read bundle {options['input']}: {err}") from err + + import_options = ImportOptions( + execute=options["execute"], + iou_threshold=options["iou_threshold"], + create_missing_detections=options["create_missing_detections"], + ) + report = import_bundle(project, bundle, import_options) + summary = report.summary() + + mode = "Applied" if import_options.execute else "Dry run" + self.stdout.write(f"{mode} on project #{project.pk} ({project.name}), bundle from {bundle.project_name!r}:") + self.stdout.write(f" occurrences: {summary['occurrences_total']} {dict(summary['occurrences'])}") + self.stdout.write(f" detections: {summary['detections_total']} {dict(summary['detections'])}") + self.stdout.write( + f" confirmed: {summary['confirmed']} identifications applied: {summary['identifications_applied']}" + f" skipped: {summary['identifications_skipped']} detections created: {summary['detections_created']}" + ) + if summary["determination_mismatches"]: + self.stdout.write(f" determinations differing from the bundle: {summary['determination_mismatches']}") + if summary["missing_users"]: + self.stdout.write(self.style.WARNING(f" users not found: {', '.join(summary['missing_users'])}")) + if summary["missing_taxa"]: + names = ", ".join(f"{t['name']} ({t['rank']})" for t in summary["missing_taxa"]) + self.stdout.write(self.style.WARNING(f" taxa not found: {names}")) + for outcome in report.outcomes: + if outcome.outcome == "error": + self.stdout.write(self.style.ERROR(f" record {outcome.ref}: {outcome.error}")) + + if options["report"]: + options["report"].parent.mkdir(parents=True, exist_ok=True) + options["report"].write_text(json.dumps(report.as_dict(), indent=1, default=str)) + self.stdout.write(f" report written to {options['report']}") + if not import_options.execute: + self.stdout.write("Nothing was changed. Re-run with --execute to apply.") diff --git a/ami/main/models_future/detection_matching.py b/ami/main/models_future/detection_matching.py new file mode 100644 index 000000000..a70f3181f --- /dev/null +++ b/ami/main/models_future/detection_matching.py @@ -0,0 +1,274 @@ +"""Find a detection again on another Antenna instance, without its primary key. + +Primary keys do not survive a move between databases: a fresh import assigns new ids to +every capture, detection and occurrence. What does survive is the capture's path under its +data source, its timestamp and station, the bounding box, and the name of the detector +that drew it. A ``DetectionKey`` carries exactly those, and ``match_detections`` turns a +list of keys back into detection ids on the target database. + +Matching is exact when the detector rerun gives the same box, and falls back to the +candidate with the highest intersection over union when it does not (a newer model +version, or a rounding difference). Every match records which of the two it was and the +IoU, so a caller can refuse to trust a fuzzy match where exactness matters. +""" + +from __future__ import annotations + +import dataclasses +import datetime +import logging +from collections import defaultdict +from collections.abc import Iterable + +from ami.main.models import Detection, Project, SourceImage + +logger = logging.getLogger(__name__) + +DEFAULT_IOU_THRESHOLD = 0.7 + +# Two boxes within this many pixels on every edge are the same box: JSON round trips and +# float formatting can shift a coordinate by a fraction of a pixel. +EXACT_BOX_TOLERANCE_PX = 0.5 + +# How a key was, or was not, found again on the target database. +MATCH_EXACT = "exact" +MATCH_IOU = "iou" +MATCH_NO_CANDIDATE = "no_candidate" # the capture exists but no box on it is close enough +MATCH_NO_CAPTURE = "no_capture" # the capture itself is not in the project + + +@dataclasses.dataclass(frozen=True) +class DetectionKey: + """The natural key of a detection: where its capture is and what box it holds.""" + + capture_path: str + bbox: tuple[float, float, float, float] + capture_timestamp: str | None = None + deployment: str | None = None + detector: str | None = None + + @classmethod + def for_detection(cls, detection: Detection) -> DetectionKey: + """Build the key of a loaded detection. ``source_image`` and its ``deployment`` should be + selected already, and ``detection_algorithm`` too, or this costs three queries per call.""" + capture = detection.source_image + return cls( + capture_path=capture.path, + bbox=normalise_bbox(detection.bbox), + capture_timestamp=capture.timestamp.isoformat() if capture.timestamp else None, + deployment=capture.deployment.name if capture.deployment_id else None, + detector=detection.detection_algorithm.name if detection.detection_algorithm_id else None, + ) + + def as_dict(self) -> dict: + data = dataclasses.asdict(self) + data["bbox"] = list(self.bbox) + return data + + @classmethod + def from_dict(cls, data: dict) -> DetectionKey: + return cls( + capture_path=data["capture_path"], + bbox=normalise_bbox(data["bbox"]), + capture_timestamp=data.get("capture_timestamp"), + deployment=data.get("deployment"), + detector=data.get("detector"), + ) + + +@dataclasses.dataclass +class DetectionMatch: + """What ``match_detections`` found for one key.""" + + key: DetectionKey + status: str + detection_id: int | None = None + capture_id: int | None = None + occurrence_id: int | None = None + event_id: int | None = None + iou: float | None = None + + @property + def found(self) -> bool: + return self.detection_id is not None + + def as_dict(self) -> dict: + return { + "key": self.key.as_dict(), + "status": self.status, + "detection_id": self.detection_id, + "capture_id": self.capture_id, + "occurrence_id": self.occurrence_id, + "iou": self.iou, + } + + +def normalise_bbox(bbox: Iterable[float] | None) -> tuple[float, float, float, float]: + if bbox is None: + raise ValueError("A detection without a bounding box has no natural key.") + x1, y1, x2, y2 = (float(value) for value in bbox) + return (x1, y1, x2, y2) + + +def bbox_iou(a: Iterable[float], b: Iterable[float]) -> float: + """Intersection over union of two ``[x1, y1, x2, y2]`` boxes; 0 when either is empty.""" + ax1, ay1, ax2, ay2 = normalise_bbox(a) + bx1, by1, bx2, by2 = normalise_bbox(b) + inter_w = min(ax2, bx2) - max(ax1, bx1) + inter_h = min(ay2, by2) - max(ay1, by1) + if inter_w <= 0 or inter_h <= 0: + return 0.0 + intersection = inter_w * inter_h + union = (ax2 - ax1) * (ay2 - ay1) + (bx2 - bx1) * (by2 - by1) - intersection + return intersection / union if union > 0 else 0.0 + + +def boxes_are_the_same(a: Iterable[float], b: Iterable[float], tolerance: float = EXACT_BOX_TOLERANCE_PX) -> bool: + return all(abs(p - q) <= tolerance for p, q in zip(normalise_bbox(a), normalise_bbox(b))) + + +def _parse_timestamp(value: str | None) -> datetime.datetime | None: + if not value: + return None + # Capture timestamps are naive local time (USE_TZ is off); drop any offset an exporter added. + return datetime.datetime.fromisoformat(value).replace(tzinfo=None) + + +@dataclasses.dataclass +class _Capture: + id: int + path: str + timestamp: datetime.datetime | None + deployment: str | None + event_id: int | None + + +@dataclasses.dataclass +class _Candidate: + id: int + bbox: list[float] + detector: str | None + occurrence_id: int | None + + +def _captures_for_keys(project: Project, keys: Iterable[DetectionKey]) -> dict[DetectionKey, _Capture | None]: + """Resolve each key's capture: by path within the project, else by station and timestamp.""" + keys = list(keys) + by_path: dict[str, list[_Capture]] = defaultdict(list) + rows = ( + SourceImage.objects.filter(project=project, path__in={key.capture_path for key in keys}) + .order_by("pk") + .values_list("pk", "path", "timestamp", "deployment__name", "event_id") + ) + for pk, path, timestamp, deployment, event_id in rows: + by_path[path].append(_Capture(pk, path, timestamp, deployment, event_id)) + + resolved: dict[DetectionKey, _Capture | None] = {} + unresolved: list[DetectionKey] = [] + for key in keys: + candidates = by_path.get(key.capture_path, []) + if len(candidates) > 1: + # The same path under two data sources: the station and the timestamp decide. + narrowed = [c for c in candidates if key.deployment is None or c.deployment == key.deployment] + wanted = _parse_timestamp(key.capture_timestamp) + if wanted is not None and len(narrowed) > 1: + narrowed = [c for c in narrowed if c.timestamp == wanted] + candidates = narrowed or candidates + if candidates: + resolved[key] = candidates[0] + else: + unresolved.append(key) + + if unresolved: + # A re-import under a different prefix changes the path but not the station or the time. + wanted = {(key.deployment, _parse_timestamp(key.capture_timestamp)) for key in unresolved} + wanted.discard((None, None)) + rows = ( + SourceImage.objects.filter( + project=project, + deployment__name__in={deployment for deployment, _ in wanted if deployment}, + timestamp__in={timestamp for _, timestamp in wanted if timestamp}, + ) + .order_by("pk") + .values_list("pk", "path", "timestamp", "deployment__name", "event_id") + ) + by_station_time = { + (deployment, timestamp): _Capture(pk, path, timestamp, deployment, event_id) + for pk, path, timestamp, deployment, event_id in rows + } + for key in unresolved: + resolved[key] = by_station_time.get((key.deployment, _parse_timestamp(key.capture_timestamp))) + return resolved + + +def _candidates_by_capture(capture_ids: Iterable[int]) -> dict[int, list[_Candidate]]: + candidates: dict[int, list[_Candidate]] = defaultdict(list) + rows = ( + Detection.objects.valid() + .filter(source_image_id__in=list(capture_ids), bbox__isnull=False) + .order_by("pk") + .values_list("pk", "source_image_id", "bbox", "detection_algorithm__name", "occurrence_id") + ) + for pk, capture_id, bbox, detector, occurrence_id in rows: + candidates[capture_id].append(_Candidate(pk, bbox, detector, occurrence_id)) + return candidates + + +def _pick(key: DetectionKey, candidates: list[_Candidate], threshold: float) -> tuple[_Candidate | None, str, float]: + """The best unclaimed candidate for ``key``: the same box, else the closest box past the threshold. + + Candidates from the key's own detector are preferred when any exist, so a box the + classifier's own detector drew wins over a coincidentally similar box from another one. + """ + same_detector = [c for c in candidates if key.detector and c.detector == key.detector] + pool = same_detector or candidates + for candidate in pool: + if boxes_are_the_same(candidate.bbox, key.bbox): + return candidate, MATCH_EXACT, 1.0 + best, best_iou = None, 0.0 + for candidate in pool: + iou = bbox_iou(candidate.bbox, key.bbox) + if iou > best_iou: + best, best_iou = candidate, iou + if best is not None and best_iou >= threshold: + return best, MATCH_IOU, best_iou + return None, MATCH_NO_CANDIDATE, best_iou + + +def match_detections( + project: Project, keys: Iterable[DetectionKey], iou_threshold: float = DEFAULT_IOU_THRESHOLD +) -> list[DetectionMatch]: + """Find each key's detection in ``project``, in the order given, with two queries per call. + + A detection on the target is claimed by at most one key, so two keys that both + resemble one box do not both land on it. Keys are matched in the order given, and a + key whose exact box was already claimed falls through to its next best candidate. + """ + keys = list(keys) + captures = _captures_for_keys(project, keys) + candidates = _candidates_by_capture({capture.id for capture in captures.values() if capture}) + claimed: set[int] = set() + matches: list[DetectionMatch] = [] + for key in keys: + capture = captures.get(key) + if capture is None: + matches.append(DetectionMatch(key=key, status=MATCH_NO_CAPTURE)) + continue + available = [c for c in candidates.get(capture.id, []) if c.id not in claimed] + picked, status, iou = _pick(key, available, iou_threshold) + if picked is None: + matches.append(DetectionMatch(key=key, status=status, capture_id=capture.id, event_id=capture.event_id)) + continue + claimed.add(picked.id) + matches.append( + DetectionMatch( + key=key, + status=status, + detection_id=picked.id, + capture_id=capture.id, + occurrence_id=picked.occurrence_id, + event_id=capture.event_id, + iou=round(iou, 4), + ) + ) + return matches diff --git a/ami/main/models_future/embedding_transfer.py b/ami/main/models_future/embedding_transfer.py new file mode 100644 index 000000000..fb312e016 --- /dev/null +++ b/ami/main/models_future/embedding_transfer.py @@ -0,0 +1,288 @@ +"""Move detection feature vectors between Antenna databases without their ids. + +A vector is worth GPU time to recompute, so when a project is re-imported the vectors +should follow it. Each vector is written next to the natural key of its detection +(``detection_matching.DetectionKey``) and the key and name of the algorithm that produced +it; on the way back in, the detections are found again by capture path and box and the +rows are written under the same algorithm. + +Layout of an export directory: ``vectors.npy`` (one row per vector), ``index.csv`` (one line +per row, the detection's natural key) and ``manifest.json`` (format version, algorithm, +vector key, dimensions, count, dtype). Vectors are read from ``DetectionEmbedding`` rows +when that table exists on this branch, otherwise from the classification vectors +(``Classification.features_2048``) the same algorithm stored, which is all older data has. +Writing always targets ``DetectionEmbedding`` and refuses when the table does not exist. +""" + +from __future__ import annotations + +import csv +import dataclasses +import datetime +import json +import logging +import pathlib +from collections.abc import Iterable, Iterator +from typing import Any + +import numpy as np +from django.apps import apps +from django.db.models import Model + +from ami.main.models import Classification, Project +from ami.main.models_future.detection_matching import ( + DEFAULT_IOU_THRESHOLD, + DetectionKey, + DetectionMatch, + match_detections, +) +from ami.ml.models import Algorithm + +logger = logging.getLogger(__name__) + +EMBEDDINGS_FORMAT = "antenna-detection-embeddings" +EMBEDDINGS_VERSION = 1 +DEFAULT_VECTOR_KEY = "embedding" +VECTORS_FILE = "vectors.npy" +INDEX_FILE = "index.csv" +MANIFEST_FILE = "manifest.json" +INDEX_COLUMNS = ["row", "capture_path", "capture_timestamp", "deployment", "detector", "x1", "y1", "x2", "y2"] +MATCH_CHUNK = 5000 +WRITE_BATCH = 1000 + + +def embedding_model() -> type[Model] | None: + """The ``DetectionEmbedding`` model, or None on a branch that does not have it yet.""" + try: + return apps.get_model("main", "DetectionEmbedding") + except LookupError: + return None + + +def _model_field_names(model: type[Model]) -> set[str]: + return {field.name for field in model._meta.get_fields()} + + +def resolve_algorithm(reference: str) -> Algorithm: + """An algorithm by key, then by name, then by id.""" + algorithm = Algorithm.objects.filter(key=reference).first() or Algorithm.objects.filter(name=reference).first() + if algorithm is None and reference.isdigit(): + algorithm = Algorithm.objects.filter(pk=int(reference)).first() + if algorithm is None: + raise Algorithm.DoesNotExist(f"No algorithm with key, name or id {reference!r}.") + return algorithm + + +@dataclasses.dataclass +class EmbeddingManifest: + algorithm_key: str + algorithm_name: str + vector_key: str + dimensions: int + count: int + dtype: str + source_store: str + project_name: str | None = None + exported_at: str = dataclasses.field(default_factory=lambda: datetime.datetime.now().isoformat()) + version: int = EMBEDDINGS_VERSION + + def as_dict(self) -> dict: + return {"format": EMBEDDINGS_FORMAT, **dataclasses.asdict(self)} + + @classmethod + def read(cls, directory: pathlib.Path) -> EmbeddingManifest: + data = json.loads((directory / MANIFEST_FILE).read_text()) + if data.get("format") != EMBEDDINGS_FORMAT: + raise ValueError(f"{directory} is not a {EMBEDDINGS_FORMAT} export.") + if int(data.get("version", 0)) > EMBEDDINGS_VERSION: + raise ValueError(f"Export version {data['version']} is newer than this code understands.") + data.pop("format") + return cls(**data) + + +def _vector_rows( + project: Project, algorithm: Algorithm, vector_key: str +) -> tuple[str, Iterator[tuple[DetectionKey, Any]]]: + """(store name, iterator of (detection key, vector)) for one algorithm's vectors in a project.""" + model = embedding_model() + if model is not None: + rows = model.objects.filter(detection__source_image__project=project, algorithm=algorithm) + if "key" in _model_field_names(model): + rows = rows.filter(key=vector_key) + rows = rows.select_related("detection__source_image__deployment", "detection__detection_algorithm").order_by( + "pk" + ) + return "detection_embedding", ( + (DetectionKey.for_detection(row.detection), row.vector) for row in rows.iterator() + ) + + # Older data: the classifier's backbone vector lives on the classification. Newest per detection. + rows = ( + Classification.objects.filter( + detection__source_image__project=project, algorithm=algorithm, features_2048__isnull=False + ) + .select_related("detection__source_image__deployment", "detection__detection_algorithm") + .order_by("detection_id", "-timestamp", "-pk") + .distinct("detection_id") + ) + return "classification_features", ( + (DetectionKey.for_detection(row.detection), row.features_2048) for row in rows.iterator() + ) + + +def export_embeddings( + project: Project, + algorithm: Algorithm, + directory: pathlib.Path, + vector_key: str = DEFAULT_VECTOR_KEY, + dtype: str = "float32", +) -> EmbeddingManifest: + """Write one algorithm's vectors for ``project`` to ``directory``; returns the manifest written.""" + directory.mkdir(parents=True, exist_ok=True) + store, rows = _vector_rows(project, algorithm, vector_key) + vectors: list[np.ndarray] = [] + with (directory / INDEX_FILE).open("w", newline="") as index_file: + writer = csv.writer(index_file) + writer.writerow(INDEX_COLUMNS) + for row, (key, vector) in enumerate(rows): + # pgvector returns lists on some versions and arrays on others; np.asarray takes both. + vectors.append(np.asarray(list(vector), dtype=dtype)) + writer.writerow( + [ + row, + key.capture_path, + key.capture_timestamp or "", + key.deployment or "", + key.detector or "", + *key.bbox, + ] + ) + if vectors: + lengths = {len(v) for v in vectors} + if len(lengths) > 1: + raise ValueError( + f"Algorithm {algorithm.key} stored vectors of several lengths {sorted(lengths)}; export one at a time." + ) + matrix = np.stack(vectors) + else: + matrix = np.zeros((0, 0), dtype=dtype) + np.save(directory / VECTORS_FILE, matrix) + manifest = EmbeddingManifest( + algorithm_key=algorithm.key, + algorithm_name=algorithm.name, + vector_key=vector_key, + dimensions=int(matrix.shape[1]) if matrix.size else 0, + count=int(matrix.shape[0]), + dtype=dtype, + source_store=store, + project_name=project.name, + ) + (directory / MANIFEST_FILE).write_text(json.dumps(manifest.as_dict(), indent=1)) + return manifest + + +def read_index(directory: pathlib.Path) -> list[DetectionKey]: + keys = [] + with (directory / INDEX_FILE).open(newline="") as index_file: + for line in csv.DictReader(index_file): + keys.append( + DetectionKey( + capture_path=line["capture_path"], + bbox=(float(line["x1"]), float(line["y1"]), float(line["x2"]), float(line["y2"])), + capture_timestamp=line["capture_timestamp"] or None, + deployment=line["deployment"] or None, + detector=line["detector"] or None, + ) + ) + return keys + + +def _chunks(items: list, size: int) -> Iterable[list]: + for start in range(0, len(items), size): + yield items[start : start + size] + + +@dataclasses.dataclass +class EmbeddingImportReport: + matches: list[DetectionMatch] + written: int = 0 + skipped_existing: int = 0 + replaced: int = 0 + execute: bool = False + + def summary(self) -> dict: + from collections import Counter + + counts = Counter(match.status for match in self.matches) + return { + "mode": "execute" if self.execute else "dry-run", + "vectors_total": len(self.matches), + "detections": dict(counts), + "written": self.written, + "skipped_existing": self.skipped_existing, + "replaced": self.replaced, + } + + +def import_embeddings( + project: Project, + directory: pathlib.Path, + execute: bool = False, + iou_threshold: float = DEFAULT_IOU_THRESHOLD, + replace: bool = False, + algorithm: Algorithm | None = None, +) -> EmbeddingImportReport: + """Write the vectors in ``directory`` as ``DetectionEmbedding`` rows on ``project``'s detections. + + Detections are found by natural key. A detection that already has a vector from the + same algorithm (and key) is skipped unless ``replace`` is set. A dry run only matches. + """ + manifest = EmbeddingManifest.read(directory) + algorithm = algorithm or resolve_algorithm(manifest.algorithm_key) + keys = read_index(directory) + if len(keys) != manifest.count: + raise ValueError(f"{INDEX_FILE} has {len(keys)} rows but the manifest says {manifest.count}.") + matches: list[DetectionMatch] = [] + for chunk in _chunks(keys, MATCH_CHUNK): + matches.extend(match_detections(project, chunk, iou_threshold)) + report = EmbeddingImportReport(matches=matches, execute=execute) + if not execute: + return report + + model = embedding_model() + if model is None: + raise RuntimeError("This branch has no DetectionEmbedding table; vectors cannot be imported here.") + fields = _model_field_names(model) + matrix = np.load(directory / VECTORS_FILE, mmap_mode="r") + if matrix.shape[0] != manifest.count: + raise ValueError(f"{VECTORS_FILE} has {matrix.shape[0]} rows but the manifest says {manifest.count}.") + + found = [(row, match.detection_id) for row, match in enumerate(matches) if match.found] + for chunk in _chunks(found, WRITE_BATCH): + detection_ids = [detection_id for _, detection_id in chunk] + existing = model.objects.filter(detection_id__in=detection_ids, algorithm=algorithm) + if "key" in fields: + existing = existing.filter(key=manifest.vector_key) + existing_ids = set(existing.values_list("detection_id", flat=True)) + if replace and existing_ids: + existing.delete() + report.replaced += len(existing_ids) + rows = [] + for row, detection_id in chunk: + if detection_id in existing_ids and not replace: + report.skipped_existing += 1 + continue + values: dict[str, Any] = { + "detection_id": detection_id, + "algorithm": algorithm, + "vector": matrix[row].tolist(), + } + # Fields the settled schema adds; absent on the draft table. + if "key" in fields: + values["key"] = manifest.vector_key + if "project" in fields: + values["project"] = project + rows.append(model(**values)) + model.objects.bulk_create(rows, batch_size=WRITE_BATCH) + report.written += len(rows) + return report diff --git a/ami/main/models_future/validated_occurrences.py b/ami/main/models_future/validated_occurrences.py new file mode 100644 index 000000000..0195a45d2 --- /dev/null +++ b/ami/main/models_future/validated_occurrences.py @@ -0,0 +1,683 @@ +"""Move what people decided about occurrences between Antenna databases. + +Two kinds of human judgement attach to an occurrence: a reviewer confirmed that its +detections are one complete individual (``grouping_verified_at``), and people identified +it as a taxon (``Identification``). Both are lost when a project is re-imported or +re-processed, because every id changes. This module writes them to a portable bundle keyed +by natural keys only (``detection_matching.DetectionKey``, taxon name and GBIF key, user +email), and replays a bundle onto another database. + +Replaying a confirmed occurrence means finding its detections again, moving them into one +occurrence with the same track-edit operations the review interface uses (``tracks``), and +recording the confirmation under the original reviewer and time. An occurrence whose +detections cannot all be found is reported as partial and is never confirmed: a reviewer +confirmed a complete set, not whatever subset survived. + +Identifications are re-created under the original user and timestamp on the occurrence +that now holds the identified detections. Users and taxa are never created here; a missing +one is reported and its rows are skipped. + +The bundle is versioned so a later schema (reviews as rows of their own) can read today's +files. Nothing in it names a database id except as an opaque ``ref`` used to resolve +"agreed with" links inside the same bundle. +""" + +from __future__ import annotations + +import dataclasses +import datetime +import logging +from collections import Counter +from collections.abc import Iterable + +from django.db import transaction +from django.db.models import Prefetch + +from ami.main.models import ( + Classification, + Detection, + Identification, + Occurrence, + Project, + SourceImage, + Taxon, + User, + update_occurrence_determination, +) +from ami.main.models_future.detection_matching import ( + DEFAULT_IOU_THRESHOLD, + MATCH_NO_CANDIDATE, + DetectionKey, + DetectionMatch, + match_detections, +) +from ami.main.models_future.tracks import TrackEditError, add_detections, detach_detection, verify_grouping + +logger = logging.getLogger(__name__) + +BUNDLE_FORMAT = "antenna-validated-occurrences" +BUNDLE_VERSION = 1 + + +# -- Bundle schema ------------------------------------------------------------------------ + + +@dataclasses.dataclass(frozen=True) +class TaxonKey: + name: str + rank: str + gbif_taxon_key: int | None = None + + @classmethod + def for_taxon(cls, taxon: Taxon | None) -> TaxonKey | None: + if taxon is None: + return None + return cls(name=taxon.name, rank=taxon.rank, gbif_taxon_key=taxon.gbif_taxon_key) + + @classmethod + def from_dict(cls, data: dict | None) -> TaxonKey | None: + if not data: + return None + return cls(name=data["name"], rank=data["rank"], gbif_taxon_key=data.get("gbif_taxon_key")) + + +@dataclasses.dataclass +class GroupingConfirmation: + user_email: str | None + verified_at: str + + @classmethod + def from_dict(cls, data: dict | None) -> GroupingConfirmation | None: + if not data: + return None + return cls(user_email=data.get("user_email"), verified_at=data["verified_at"]) + + +@dataclasses.dataclass +class IdentificationRecord: + ref: str + created_at: str + user_email: str | None + taxon: TaxonKey | None + withdrawn: bool = False + comment: str = "" + agreed_with_identification_ref: str | None = None + # The prediction agreed with, by taxon and algorithm name: its row will not exist on the target. + agreed_with_prediction_taxon: TaxonKey | None = None + agreed_with_prediction_algorithm: str | None = None + + def as_dict(self) -> dict: + return dataclasses.asdict(self) + + @classmethod + def from_dict(cls, data: dict) -> IdentificationRecord: + return cls( + ref=str(data["ref"]), + created_at=data["created_at"], + user_email=data.get("user_email"), + taxon=TaxonKey.from_dict(data.get("taxon")), + withdrawn=bool(data.get("withdrawn", False)), + comment=data.get("comment") or "", + agreed_with_identification_ref=data.get("agreed_with_identification_ref"), + agreed_with_prediction_taxon=TaxonKey.from_dict(data.get("agreed_with_prediction_taxon")), + agreed_with_prediction_algorithm=data.get("agreed_with_prediction_algorithm"), + ) + + +@dataclasses.dataclass +class OccurrenceRecord: + """One occurrence worth carrying: confirmed as a grouping, identified, or both.""" + + ref: str + detections: list[DetectionKey] + confirmation: GroupingConfirmation | None = None + determination: TaxonKey | None = None + identifications: list[IdentificationRecord] = dataclasses.field(default_factory=list) + + def as_dict(self) -> dict: + return { + "ref": self.ref, + "detections": [key.as_dict() for key in self.detections], + "confirmation": dataclasses.asdict(self.confirmation) if self.confirmation else None, + "determination": dataclasses.asdict(self.determination) if self.determination else None, + "identifications": [identification.as_dict() for identification in self.identifications], + } + + @classmethod + def from_dict(cls, data: dict) -> OccurrenceRecord: + return cls( + ref=str(data["ref"]), + detections=[DetectionKey.from_dict(key) for key in data["detections"]], + confirmation=GroupingConfirmation.from_dict(data.get("confirmation")), + determination=TaxonKey.from_dict(data.get("determination")), + identifications=[IdentificationRecord.from_dict(row) for row in data.get("identifications", [])], + ) + + +@dataclasses.dataclass +class Bundle: + occurrences: list[OccurrenceRecord] + project_name: str | None = None + exported_at: str = dataclasses.field(default_factory=lambda: datetime.datetime.now().isoformat()) + source: str | None = None + version: int = BUNDLE_VERSION + + def as_dict(self) -> dict: + return { + "format": BUNDLE_FORMAT, + "version": self.version, + "exported_at": self.exported_at, + "project_name": self.project_name, + "source": self.source, + "occurrences": [record.as_dict() for record in self.occurrences], + } + + @classmethod + def from_dict(cls, data: dict) -> Bundle: + if data.get("format") != BUNDLE_FORMAT: + raise ValueError(f"Not a {BUNDLE_FORMAT} bundle (format={data.get('format')!r}).") + if int(data.get("version", 0)) > BUNDLE_VERSION: + raise ValueError(f"Bundle version {data['version']} is newer than this code understands.") + return cls( + occurrences=[OccurrenceRecord.from_dict(row) for row in data["occurrences"]], + project_name=data.get("project_name"), + exported_at=data.get("exported_at") or "", + source=data.get("source"), + version=int(data.get("version", BUNDLE_VERSION)), + ) + + +# -- Export ------------------------------------------------------------------------------- + + +def _iso(value: datetime.datetime | None) -> str | None: + return value.isoformat() if value else None + + +def build_bundle(project: Project) -> Bundle: + """Every occurrence in ``project`` that a person confirmed or identified, with its detection keys.""" + detections = Detection.objects.valid().select_related("source_image__deployment", "detection_algorithm") + identifications = Identification.objects.select_related( + "user", "taxon", "agreed_with_prediction__taxon", "agreed_with_prediction__algorithm" + ).order_by("created_at", "pk") + occurrences = ( + Occurrence.objects.filter(project=project) + .filter(grouping_verified_at__isnull=False) + .union(Occurrence.objects.filter(project=project, identifications__isnull=False)) + .values_list("pk", flat=True) + ) + occurrences = ( + Occurrence.objects.filter(pk__in=list(occurrences)) + .select_related("determination", "grouping_verified_by") + .prefetch_related( + Prefetch("detections", queryset=detections.order_by("source_image__timestamp", "source_image_id", "pk")), + Prefetch("identifications", queryset=identifications), + ) + .order_by("pk") + ) + + records = [] + for occurrence in occurrences: + keys = [DetectionKey.for_detection(d) for d in occurrence.detections.all() if d.bbox] + if not keys: + logger.warning(f"Occurrence {occurrence.pk} has no detection with a box and was left out of the bundle.") + continue + confirmation = None + if occurrence.grouping_verified_at is not None: + reviewer = occurrence.grouping_verified_by + confirmation = GroupingConfirmation( + user_email=reviewer.email if reviewer else None, + verified_at=occurrence.grouping_verified_at.isoformat(), + ) + records.append( + OccurrenceRecord( + ref=str(occurrence.pk), + detections=keys, + confirmation=confirmation, + determination=TaxonKey.for_taxon(occurrence.determination), + identifications=[ + IdentificationRecord( + ref=str(identification.pk), + created_at=identification.created_at.isoformat(), + user_email=identification.user.email if identification.user else None, + taxon=TaxonKey.for_taxon(identification.taxon), + withdrawn=identification.withdrawn, + comment=identification.comment or "", + agreed_with_identification_ref=( + str(identification.agreed_with_identification_id) + if identification.agreed_with_identification_id + else None + ), + agreed_with_prediction_taxon=( + TaxonKey.for_taxon(identification.agreed_with_prediction.taxon) + if identification.agreed_with_prediction_id + else None + ), + agreed_with_prediction_algorithm=( + identification.agreed_with_prediction.algorithm.name + if identification.agreed_with_prediction_id + and identification.agreed_with_prediction.algorithm_id + else None + ), + ) + for identification in occurrence.identifications.all() + ], + ) + ) + return Bundle(occurrences=records, project_name=project.name, source=f"project:{project.pk}") + + +# -- Resolving users and taxa ------------------------------------------------------------- + + +class _Lookups: + """Users by email and taxa by GBIF key or name and rank, each fetched once per import.""" + + def __init__(self, bundle: Bundle): + emails: set[str] = set() + taxa: set[TaxonKey] = set() + for record in bundle.occurrences: + if record.confirmation and record.confirmation.user_email: + emails.add(record.confirmation.user_email.lower()) + for identification in record.identifications: + if identification.user_email: + emails.add(identification.user_email.lower()) + if identification.taxon: + taxa.add(identification.taxon) + if identification.agreed_with_prediction_taxon: + taxa.add(identification.agreed_with_prediction_taxon) + if record.determination: + taxa.add(record.determination) + + self.users = {user.email.lower(): user for user in User.objects.filter(email__in=emails)} if emails else {} + self._by_gbif: dict[int, Taxon] = {} + self._by_name: dict[tuple[str, str], Taxon] = {} + gbif_keys = {key.gbif_taxon_key for key in taxa if key.gbif_taxon_key} + names = {key.name for key in taxa} + if taxa: + for taxon in Taxon.objects.filter(gbif_taxon_key__in=gbif_keys): + self._by_gbif.setdefault(taxon.gbif_taxon_key, taxon) + for taxon in Taxon.objects.filter(name__in=names): + self._by_name.setdefault((taxon.name.lower(), taxon.rank), taxon) + self.missing_users: set[str] = set() + self.missing_taxa: set[TaxonKey] = set() + + def user(self, email: str | None) -> User | None: + if not email: + return None + user = self.users.get(email.lower()) + if user is None: + self.missing_users.add(email) + return user + + def taxon(self, key: TaxonKey | None) -> Taxon | None: + if key is None: + return None + taxon = self._by_gbif.get(key.gbif_taxon_key) if key.gbif_taxon_key else None + taxon = taxon or self._by_name.get((key.name.lower(), key.rank)) + if taxon is None: + self.missing_taxa.add(key) + return taxon + + +# -- Import ------------------------------------------------------------------------------- + +# Outcome of one occurrence record. ``applied`` and ``unchanged`` are the two good ends. +OUTCOME_APPLIED = "applied" # detections regrouped and/or confirmation written +OUTCOME_UNCHANGED = "unchanged" # the target already matched the record (idempotent re-run) +OUTCOME_PARTIAL = "partial" # some detections were not found; grouping left alone, not confirmed +OUTCOME_ERROR = "error" # a track edit refused (sessions differ, two boxes on one capture, ...) + + +@dataclasses.dataclass +class ImportOptions: + execute: bool = False + iou_threshold: float = DEFAULT_IOU_THRESHOLD + # Recreate a box the reviewer drew by hand when its capture exists but no detector found it. + create_missing_detections: bool = False + + +@dataclasses.dataclass +class OccurrenceOutcome: + ref: str + outcome: str + matches: list[DetectionMatch] + target_occurrence_id: int | None = None + # How the matched detections were spread over occurrences before the import touched them. + occurrences_before: int = 0 + detached: int = 0 + added: int = 0 + created: int = 0 + confirmed: bool = False + identifications_applied: int = 0 + identifications_skipped: int = 0 + determination_matches: bool | None = None + error: str | None = None + + @property + def match_counts(self) -> Counter: + return Counter(match.status for match in self.matches) + + def as_dict(self) -> dict: + data = dataclasses.asdict(self) + data["matches"] = [match.as_dict() for match in self.matches] + return data + + +@dataclasses.dataclass +class ImportReport: + options: ImportOptions + outcomes: list[OccurrenceOutcome] = dataclasses.field(default_factory=list) + missing_users: list[str] = dataclasses.field(default_factory=list) + missing_taxa: list[dict] = dataclasses.field(default_factory=list) + + def detection_counts(self) -> Counter: + counts: Counter = Counter() + for outcome in self.outcomes: + counts.update(outcome.match_counts) + return counts + + def outcome_counts(self) -> Counter: + return Counter(outcome.outcome for outcome in self.outcomes) + + def summary(self) -> dict: + detections = self.detection_counts() + return { + "mode": "execute" if self.options.execute else "dry-run", + "iou_threshold": self.options.iou_threshold, + "occurrences": dict(self.outcome_counts()), + "occurrences_total": len(self.outcomes), + "detections": dict(detections), + "detections_total": sum(detections.values()), + "confirmed": sum(1 for o in self.outcomes if o.confirmed), + "detections_created": sum(o.created for o in self.outcomes), + "identifications_applied": sum(o.identifications_applied for o in self.outcomes), + "identifications_skipped": sum(o.identifications_skipped for o in self.outcomes), + "determination_mismatches": sum(1 for o in self.outcomes if o.determination_matches is False), + "missing_users": self.missing_users, + "missing_taxa": self.missing_taxa, + } + + def as_dict(self) -> dict: + return {"summary": self.summary(), "occurrences": [outcome.as_dict() for outcome in self.outcomes]} + + +def _parse_datetime(value: str) -> datetime.datetime: + # Stored times are naive local time (USE_TZ is off); drop any offset an exporter added. + return datetime.datetime.fromisoformat(value).replace(tzinfo=None) + + +def confirm_grouping_as_of(occurrence: Occurrence, user: User, verified_at: datetime.datetime) -> None: + """Record a grouping confirmation made by ``user`` at ``verified_at`` on another database. + + The one place the import writes a confirmation. ``verify_grouping`` stamps the current + time, so the exported time is written over it afterwards to keep the provenance. When + confirmations become review rows of their own, this function changes and nothing else does. + """ + verify_grouping(occurrence, user) + occurrence.grouping_verified_at = verified_at + Occurrence.objects.filter(pk=occurrence.pk).update(grouping_verified_at=verified_at) + + +def _is_confirmed_as(occurrence: Occurrence, user: User | None, verified_at: datetime.datetime | None) -> bool: + return ( + occurrence.grouping_verified_at == verified_at + and user is not None + and occurrence.grouping_verified_by_id == user.pk + ) + + +def _holder_counts(detection_ids: Iterable[int]) -> Counter: + """How many of ``detection_ids`` each occurrence holds right now; ``None`` for unattached ones.""" + return Counter( + Detection.objects.filter(pk__in=list(detection_ids)).order_by().values_list("occurrence_id", flat=True) + ) + + +def _majority_holder(holders: Counter) -> int | None: + attached = {pk: n for pk, n in holders.items() if pk is not None} + return max(attached, key=lambda pk: (attached[pk], -pk)) if attached else None + + +def _target_occurrence(project: Project, matches: list[DetectionMatch], holders: Counter) -> Occurrence: + """The occurrence to rebuild the record on: the one already holding most of its detections. + + Ties go to the lowest id. When no matched detection is attached to any occurrence (a + detector-only project), a new occurrence is created in the session of the first capture. + """ + best = _majority_holder(holders) + if best is not None: + return Occurrence.objects.get(pk=best) + first = next(match for match in matches if match.found) + capture = SourceImage.objects.select_related("deployment").get(pk=first.capture_id) + return Occurrence.objects.create(project=project, deployment=capture.deployment, event_id=capture.event_id) + + +def _create_missing_detections(record: OccurrenceRecord, matches: list[DetectionMatch]) -> int: + """Recreate boxes a reviewer added by hand: their capture exists but no detector drew them.""" + created = 0 + for match in matches: + if match.status != MATCH_NO_CANDIDATE or match.capture_id is None: + continue + capture = SourceImage.objects.get(pk=match.capture_id) + detection = Detection.objects.create( + source_image=capture, bbox=list(match.key.bbox), timestamp=capture.timestamp + ) + match.detection_id = detection.pk + match.status = "created" + match.iou = 1.0 + created += 1 + return created + + +def _regroup(target: Occurrence, matched_ids: set[int], outcome: OccurrenceOutcome) -> None: + """Make ``target`` hold exactly ``matched_ids`` using the track-edit operations. + + Detections the target holds that are not in the record leave first, one by one, each + into an occurrence of its own; then the record's detections still elsewhere move in. + Every step keeps the chain, statistics and session counts the way the review interface does. + """ + current = set(Detection.objects.valid().filter(occurrence=target).values_list("pk", flat=True)) + for pk in sorted(current - matched_ids): + detach_detection(target, Detection.objects.get(pk=pk)) + outcome.detached += 1 + incoming = matched_ids - current + if incoming: + add_detections(target, Detection.objects.filter(pk__in=incoming).select_related("source_image")) + outcome.added += len(incoming) + + +def _apply_identifications( + target: Occurrence, + record: OccurrenceRecord, + lookups: _Lookups, + outcome: OccurrenceOutcome, + created_by_ref: dict[str, Identification], +) -> None: + """Re-create the record's identifications on ``target`` under their original user and time. + + Oldest first, so ``Identification.save`` withdraws earlier ones by the same user in the + order it did originally; the exported ``withdrawn`` flag is then written as the final + word. An identification that already exists with the same user, taxon and time is skipped. + """ + for row in sorted(record.identifications, key=lambda r: r.created_at): + user = lookups.user(row.user_email) + taxon = lookups.taxon(row.taxon) + if (row.user_email and user is None) or (row.taxon and taxon is None): + outcome.identifications_skipped += 1 + continue + created_at = _parse_datetime(row.created_at) + existing = Identification.objects.filter( + occurrence=target, user=user, taxon=taxon, created_at=created_at + ).first() + if existing is not None: + created_by_ref[row.ref] = existing + outcome.identifications_skipped += 1 + continue + agreed_prediction = None + agreed_taxon = lookups.taxon(row.agreed_with_prediction_taxon) if row.agreed_with_prediction_taxon else None + if agreed_taxon is not None: + predictions = Classification.objects.filter(detection__occurrence=target, taxon=agreed_taxon) + if row.agreed_with_prediction_algorithm: + predictions = predictions.filter(algorithm__name=row.agreed_with_prediction_algorithm) + agreed_prediction = predictions.order_by("-score").first() + identification = Identification( + occurrence=target, + user=user, + taxon=taxon, + withdrawn=row.withdrawn, + comment=row.comment, + agreed_with_prediction=agreed_prediction, + ) + identification.save() + # auto_now_add ignores a value given at creation; the original time is the record of who said what when. + Identification.objects.filter(pk=identification.pk).update( + created_at=created_at, updated_at=created_at, withdrawn=row.withdrawn + ) + created_by_ref[row.ref] = identification + outcome.identifications_applied += 1 + + +def _link_agreements(bundle: Bundle, created_by_ref: dict[str, Identification]) -> None: + for record in bundle.occurrences: + for row in record.identifications: + if row.agreed_with_identification_ref and row.ref in created_by_ref: + agreed = created_by_ref.get(row.agreed_with_identification_ref) + if agreed is not None: + Identification.objects.filter(pk=created_by_ref[row.ref].pk).update( + agreed_with_identification=agreed + ) + + +def _pending_identifications(target_pk: int | None, record: OccurrenceRecord, lookups: _Lookups) -> int: + """How many of the record's identifications are not yet on the target occurrence.""" + pending = 0 + for row in record.identifications: + user = lookups.user(row.user_email) + taxon = lookups.taxon(row.taxon) + if (row.user_email and user is None) or (row.taxon and taxon is None): + continue + exists = ( + target_pk is not None + and Identification.objects.filter( + occurrence_id=target_pk, user=user, taxon=taxon, created_at=_parse_datetime(row.created_at) + ).exists() + ) + pending += 0 if exists else 1 + return pending + + +def _apply_record( + project: Project, + record: OccurrenceRecord, + lookups: _Lookups, + options: ImportOptions, + created_by_ref: dict[str, Identification], +) -> OccurrenceOutcome: + """Replay one record: regroup and confirm when it carries a confirmation, then re-attach its identifications. + + Only a confirmed record has its grouping rebuilt. An identified-but-unconfirmed + occurrence keeps whatever grouping the target has; its identifications land on the + occurrence holding most of its detections, and the report says over how many + occurrences they were spread. + """ + matches = match_detections(project, record.detections, options.iou_threshold) + outcome = OccurrenceOutcome(ref=record.ref, outcome=OUTCOME_PARTIAL, matches=matches) + + if options.execute and options.create_missing_detections: + outcome.created = _create_missing_detections(record, matches) + complete = all(match.found for match in matches) + if not options.execute and options.create_missing_detections: + # A dry run counts the boxes an execute run would recreate as found. + complete = all(m.found or (m.status == MATCH_NO_CANDIDATE and m.capture_id) for m in matches) + matched_ids = {match.detection_id for match in matches if match.found} + if not matched_ids: + return outcome + holders = _holder_counts(matched_ids) + outcome.occurrences_before = len(holders) + reviewer = lookups.user(record.confirmation.user_email) if record.confirmation else None + verified_at = _parse_datetime(record.confirmation.verified_at) if record.confirmation else None + rebuild = complete and record.confirmation is not None + + if not options.execute: + # Report what an execute run would do without touching a row. + target_pk = _majority_holder(holders) + outcome.target_occurrence_id = target_pk + pending = _pending_identifications(target_pk, record, lookups) + if rebuild: + grouped = ( + target_pk is not None + and len(holders) == 1 + and Detection.objects.valid().filter(occurrence_id=target_pk).count() == len(matched_ids) + ) + confirmed = grouped and _is_confirmed_as(Occurrence.objects.get(pk=target_pk), reviewer, verified_at) + outcome.confirmed = bool(reviewer) and not confirmed + changed = outcome.confirmed or not grouped + else: + changed = False + if complete: + outcome.outcome = OUTCOME_APPLIED if (changed or pending) else OUTCOME_UNCHANGED + return outcome + + try: + with transaction.atomic(): + target = _target_occurrence(project, matches, holders) + outcome.target_occurrence_id = target.pk + changed = False + if rebuild: + _regroup(target, matched_ids, outcome) + changed = bool(outcome.detached or outcome.added or outcome.created or len(holders) != 1) + if reviewer is None: + logger.warning(f"Record {record.ref}: reviewer {record.confirmation.user_email} not found.") + elif changed or not _is_confirmed_as(target, reviewer, verified_at): + confirm_grouping_as_of(target, reviewer, verified_at) + outcome.confirmed = True + changed = True + # A partial record is never confirmed, but what people said about it still applies to + # the occurrence that holds most of its detections. + _apply_identifications(target, record, lookups, outcome, created_by_ref) + if outcome.identifications_applied: + target.refresh_from_db() + update_occurrence_determination(target) + changed = True + if complete: + outcome.outcome = OUTCOME_APPLIED if changed else OUTCOME_UNCHANGED + if record.determination: + target.refresh_from_db() + wanted = lookups.taxon(record.determination) + outcome.determination_matches = wanted is not None and target.determination_id == wanted.pk + except TrackEditError as err: + outcome.outcome = OUTCOME_ERROR + outcome.error = str(err) + return outcome + + +def import_bundle(project: Project, bundle: Bundle, options: ImportOptions | None = None) -> ImportReport: + """Replay ``bundle`` onto ``project``; a dry run (the default) only reports what would happen. + + Each record is applied in its own transaction, so a refused edit on one occurrence + leaves the others in place. Re-running on an already replayed project changes nothing + and reports every record as unchanged. + """ + options = options or ImportOptions() + lookups = _Lookups(bundle) + report = ImportReport(options=options) + created_by_ref: dict[str, Identification] = {} + for record in bundle.occurrences: + report.outcomes.append(_apply_record(project, record, lookups, options, created_by_ref)) + if options.execute: + _link_agreements(bundle, created_by_ref) + report.missing_users = sorted(lookups.missing_users) + report.missing_taxa = [dataclasses.asdict(key) for key in sorted(lookups.missing_taxa, key=lambda k: k.name)] + return report + + +def occurrence_detection_keys(occurrence: Occurrence) -> list[DetectionKey]: + """The keys of an occurrence's detections in capture order; for verifying a replay.""" + detections = ( + Detection.objects.valid() + .filter(occurrence=occurrence) + .select_related("source_image__deployment", "detection_algorithm") + .order_by("source_image__timestamp", "source_image_id", "pk") + ) + return [DetectionKey.for_detection(d) for d in detections if d.bbox] diff --git a/ami/main/test_validated_occurrences.py b/ami/main/test_validated_occurrences.py new file mode 100644 index 000000000..d0000306e --- /dev/null +++ b/ami/main/test_validated_occurrences.py @@ -0,0 +1,326 @@ +"""Replaying confirmed occurrences and identifications onto another project. + +Every test exports from one project and imports into a second one that holds the same +captures and boxes but, as a pipeline leaves them, one occurrence per detection and no +confirmations. What a replay must preserve: the reviewer and time of each confirmation, +the user and time of each identification, and the rule that a partially found occurrence +is never confirmed. +""" + +import dataclasses +import datetime +import pathlib +import tempfile + +import numpy as np +from django.test import TestCase + +from ami.main.models import ( + Detection, + Identification, + Occurrence, + SourceImage, + Taxon, + TaxonRank, + User, + group_images_into_events, +) +from ami.main.models_future.detection_matching import MATCH_EXACT, MATCH_IOU, MATCH_NO_CANDIDATE, bbox_iou +from ami.main.models_future.embedding_transfer import ( + EmbeddingManifest, + embedding_model, + export_embeddings, + import_embeddings, + read_index, +) +from ami.main.models_future.tracks import verify_grouping +from ami.main.models_future.validated_occurrences import ( + OUTCOME_APPLIED, + OUTCOME_PARTIAL, + OUTCOME_UNCHANGED, + Bundle, + ImportOptions, + build_bundle, + import_bundle, + occurrence_detection_keys, +) +from ami.ml.models import Algorithm +from ami.tests.fixtures.main import create_taxa, setup_test_project +from ami.tests.fixtures.tracking import pgvector_is_available + +CONFIRMED_AT = datetime.datetime(2024, 6, 10, 9, 30, 0) +IDENTIFIED_AT = datetime.datetime(2024, 6, 9, 8, 0, 0) +WITHDRAWN_AT = datetime.datetime(2024, 6, 8, 8, 0, 0) +TRACK_BOX = [10.0, 10.0, 40.0, 40.0] +OTHER_BOX = [100.0, 100.0, 140.0, 140.0] + + +class ReplayTestCase(TestCase): + def setUp(self) -> None: + self.source_project, self.source_deployment = setup_test_project(reuse=False) + self.target_project, self.target_deployment = setup_test_project(reuse=False) + create_taxa(self.source_project) + self.taxon, self.other_taxon = list( + Taxon.objects.filter(rank=TaxonRank.SPECIES.name, projects=self.source_project).order_by("pk")[:2] + ) + self.reviewer = User.objects.create_user(email="reviewer@example.org", name="Reviewer") # type: ignore + self.identifier = User.objects.create_user(email="identifier@example.org") # type: ignore[attr-defined] + self.source_captures = self._make_captures(self.source_deployment) + self.target_captures = self._make_captures(self.target_deployment) + self.track, self.identified = self._make_source_occurrences() + + def _make_captures(self, deployment) -> list[SourceImage]: + start = datetime.datetime(2024, 6, 1, 22, 0) + captures = [ + SourceImage.objects.create( + deployment=deployment, + project=deployment.project, + timestamp=start + datetime.timedelta(minutes=i), + path=f"replay/capture-{i}.jpg", + width=640, + height=480, + ) + for i in range(6) + ] + group_images_into_events(deployment) + for capture in captures: + capture.refresh_from_db() + return captures + + def _make_source_occurrences(self) -> tuple[Occurrence, Occurrence]: + """A confirmed four-frame track with two identifications, and an identified two-frame occurrence.""" + track = self._occurrence_with_boxes(self.source_project, self.source_captures[:4], TRACK_BOX) + Identification.objects.create(occurrence=track, user=self.identifier, taxon=self.other_taxon) + Identification.objects.filter(occurrence=track).update(created_at=WITHDRAWN_AT, updated_at=WITHDRAWN_AT) + Identification.objects.create(occurrence=track, user=self.identifier, taxon=self.taxon, comment="sure") + Identification.objects.filter(occurrence=track, taxon=self.taxon).update( + created_at=IDENTIFIED_AT, updated_at=IDENTIFIED_AT + ) + verify_grouping(track, self.reviewer) + Occurrence.objects.filter(pk=track.pk).update(grouping_verified_at=CONFIRMED_AT) + + identified = self._occurrence_with_boxes(self.source_project, self.source_captures[4:], OTHER_BOX) + Identification.objects.create(occurrence=identified, user=self.reviewer, taxon=self.taxon) + Identification.objects.filter(occurrence=identified).update(created_at=IDENTIFIED_AT, updated_at=IDENTIFIED_AT) + return Occurrence.objects.get(pk=track.pk), Occurrence.objects.get(pk=identified.pk) + + def _occurrence_with_boxes(self, project, captures, bbox) -> Occurrence: + occurrence = Occurrence.objects.create( + event=captures[0].event, deployment=captures[0].deployment, project=project + ) + for capture in captures: + Detection.objects.create( + source_image=capture, timestamp=capture.timestamp, bbox=bbox, occurrence=occurrence + ) + occurrence.save() + return occurrence + + def _populate_target(self, shift: float = 0.0, skip_capture_index: int | None = None) -> None: + """One occurrence per detection, the way a detector run leaves a project.""" + for index, capture in enumerate(self.target_captures): + if index == skip_capture_index: + continue + bbox = [coordinate + shift for coordinate in (TRACK_BOX if index < 4 else OTHER_BOX)] + self._occurrence_with_boxes(self.target_project, [capture], bbox) + + def _replay(self, **options) -> tuple[Bundle, object]: + bundle = Bundle.from_dict(build_bundle(self.source_project).as_dict()) + return bundle, import_bundle(self.target_project, bundle, ImportOptions(**options)) + + def _target_track(self) -> Occurrence: + return Occurrence.objects.get(project=self.target_project, grouping_verified_at__isnull=False) + + +class TestReplayConfirmedOccurrences(ReplayTestCase): + def test_the_bundle_names_detections_by_capture_and_box_only(self): + bundle = build_bundle(self.source_project) + + self.assertEqual({record.ref for record in bundle.occurrences}, {str(self.track.pk), str(self.identified.pk)}) + track = next(record for record in bundle.occurrences if record.ref == str(self.track.pk)) + self.assertEqual([key.capture_path for key in track.detections], [c.path for c in self.source_captures[:4]]) + self.assertEqual(track.detections[0].bbox, tuple(TRACK_BOX)) + self.assertEqual(track.confirmation.user_email, self.reviewer.email) + self.assertEqual(track.confirmation.verified_at, CONFIRMED_AT.isoformat()) + self.assertEqual(len(track.identifications), 2) + + def test_a_round_trip_rebuilds_the_track_under_the_original_reviewer_and_time(self): + self._populate_target() + + _, report = self._replay(execute=True) + + self.assertEqual(report.summary()["detections"], {MATCH_EXACT: 6}) + self.assertEqual(report.outcome_counts(), {OUTCOME_APPLIED: 2}) + rebuilt = self._target_track() + self.assertEqual( + [key.capture_path for key in occurrence_detection_keys(rebuilt)], + [c.path for c in self.target_captures[:4]], + ) + self.assertEqual(rebuilt.grouping_verified_at, CONFIRMED_AT, "The confirmation keeps the exported time") + self.assertEqual(rebuilt.grouping_verified_by, self.reviewer) + self.assertEqual(Occurrence.objects.filter(project=self.target_project).count(), 3, "4 singles became 1 track") + + def test_identifications_keep_their_user_time_and_withdrawn_state(self): + self._populate_target() + + self._replay(execute=True) + + rebuilt = self._target_track() + current = Identification.objects.get(occurrence=rebuilt, withdrawn=False) + self.assertEqual( + (current.user, current.taxon, current.created_at, current.comment), + (self.identifier, self.taxon, IDENTIFIED_AT, "sure"), + ) + withdrawn = Identification.objects.get(occurrence=rebuilt, withdrawn=True) + self.assertEqual((withdrawn.taxon, withdrawn.created_at), (self.other_taxon, WITHDRAWN_AT)) + self.assertEqual(rebuilt.determination, self.taxon) + other = Identification.objects.get(user=self.reviewer, occurrence__project=self.target_project) + self.assertEqual(other.occurrence.detections.first().source_image.path, self.target_captures[4].path) + self.assertIsNone(other.occurrence.grouping_verified_at, "An identified-only occurrence is not confirmed") + + def test_boxes_that_moved_a_little_match_by_overlap(self): + self._populate_target(shift=2.0) + self.assertGreater(bbox_iou(TRACK_BOX, [c + 2.0 for c in TRACK_BOX]), 0.7) + + _, report = self._replay(execute=True) + + self.assertEqual(report.summary()["detections"], {MATCH_IOU: 6}) + self.assertEqual(self._target_track().grouping_verified_at, CONFIRMED_AT) + + def test_a_stricter_overlap_threshold_leaves_moved_boxes_unmatched(self): + self._populate_target(shift=2.0) + + _, report = self._replay(execute=True, iou_threshold=0.9) + + self.assertEqual(report.summary()["detections"], {MATCH_NO_CANDIDATE: 6}) + self.assertEqual(report.outcome_counts(), {OUTCOME_PARTIAL: 2}) + self.assertFalse( + Occurrence.objects.filter(project=self.target_project, grouping_verified_at__isnull=False).exists() + ) + + def test_an_occurrence_with_a_missing_detection_is_reported_partial_and_not_confirmed(self): + self._populate_target(skip_capture_index=2) + + _, report = self._replay(execute=True) + + track_outcome = next(o for o in report.outcomes if o.ref == str(self.track.pk)) + self.assertEqual(track_outcome.outcome, OUTCOME_PARTIAL) + self.assertEqual(track_outcome.match_counts, {MATCH_EXACT: 3, MATCH_NO_CANDIDATE: 1}) + self.assertFalse(track_outcome.confirmed) + self.assertFalse( + Occurrence.objects.filter(project=self.target_project, grouping_verified_at__isnull=False).exists() + ) + self.assertEqual( + Occurrence.objects.filter(project=self.target_project).count(), 5, "The grouping was left alone" + ) + self.assertEqual( + track_outcome.identifications_applied, 2, "What people said still lands on the nearest occurrence" + ) + + def test_a_hand_drawn_box_can_be_recreated_on_its_capture(self): + self._populate_target(skip_capture_index=2) + + _, report = self._replay(execute=True, create_missing_detections=True) + + track_outcome = next(o for o in report.outcomes if o.ref == str(self.track.pk)) + self.assertEqual((track_outcome.outcome, track_outcome.created), (OUTCOME_APPLIED, 1)) + rebuilt = self._target_track() + self.assertEqual(rebuilt.detections.count(), 4) + recreated = rebuilt.detections.get(source_image=self.target_captures[2]) + self.assertEqual((recreated.bbox, recreated.timestamp), (TRACK_BOX, self.target_captures[2].timestamp)) + + def test_rerunning_the_import_changes_nothing(self): + self._populate_target() + self._replay(execute=True) + before = ( + Occurrence.objects.filter(project=self.target_project).count(), + Identification.objects.filter(occurrence__project=self.target_project).count(), + self._target_track().grouping_verified_at, + ) + + _, report = self._replay(execute=True) + + self.assertEqual(report.outcome_counts(), {OUTCOME_UNCHANGED: 2}) + self.assertEqual(report.summary()["identifications_skipped"], 3) + self.assertEqual(report.summary()["identifications_applied"], 0) + after = ( + Occurrence.objects.filter(project=self.target_project).count(), + Identification.objects.filter(occurrence__project=self.target_project).count(), + self._target_track().grouping_verified_at, + ) + self.assertEqual(before, after) + + def test_a_dry_run_reports_the_plan_and_writes_nothing(self): + self._populate_target() + + _, report = self._replay(execute=False) + + self.assertEqual(report.outcome_counts(), {OUTCOME_APPLIED: 2}) + self.assertEqual(report.summary()["detections"], {MATCH_EXACT: 6}) + self.assertTrue(next(o for o in report.outcomes if o.ref == str(self.track.pk)).confirmed, "would confirm") + self.assertEqual(Occurrence.objects.filter(project=self.target_project).count(), 6) + self.assertFalse(Identification.objects.filter(occurrence__project=self.target_project).exists()) + + def test_a_missing_reviewer_or_taxon_is_reported_and_skipped(self): + self._populate_target() + bundle = build_bundle(self.source_project) + track = next(record for record in bundle.occurrences if record.ref == str(self.track.pk)) + track.confirmation.user_email = "nobody@example.org" + first = track.identifications[0] + track.identifications[0] = dataclasses.replace( + first, taxon=dataclasses.replace(first.taxon, name="Not a real taxon", gbif_taxon_key=None) + ) + + report = import_bundle(self.target_project, bundle, ImportOptions(execute=True)) + + self.assertEqual(report.missing_users, ["nobody@example.org"]) + self.assertEqual([t["name"] for t in report.missing_taxa], ["Not a real taxon"]) + track_outcome = next(o for o in report.outcomes if o.ref == str(self.track.pk)) + self.assertFalse(track_outcome.confirmed) + self.assertEqual((track_outcome.identifications_applied, track_outcome.identifications_skipped), (1, 1)) + + +class TestEmbeddingTransfer(ReplayTestCase): + def setUp(self) -> None: + if not pgvector_is_available(): + self.skipTest("pgvector is not installed in this database") + super().setUp() + self.algorithm = Algorithm.objects.create(key="replay-test-backbone", name="Replay Test Backbone") + rng = np.random.default_rng(7) + self.vectors = {} + for detection in Detection.objects.filter(source_image__project=self.source_project): + vector = rng.standard_normal(2048).astype(np.float32) + detection.classifications.create( + taxon=self.taxon, + score=0.5, + algorithm=self.algorithm, + timestamp=detection.timestamp, + features_2048=vector.tolist(), + ) + self.vectors[detection.source_image.path] = vector + + def test_vectors_round_trip_by_detection_key(self): + self._populate_target() + with tempfile.TemporaryDirectory() as directory: + directory = pathlib.Path(directory) + manifest = export_embeddings(self.source_project, self.algorithm, directory) + + self.assertEqual( + (manifest.count, manifest.dimensions, manifest.algorithm_key), (6, 2048, "replay-test-backbone") + ) + self.assertEqual(EmbeddingManifest.read(directory).count, 6) + keys = read_index(directory) + matrix = np.load(directory / "vectors.npy") + for row, key in enumerate(keys): + np.testing.assert_array_equal(matrix[row], self.vectors[key.capture_path]) + + report = import_embeddings(self.target_project, directory, execute=False) + self.assertEqual(report.summary()["detections"], {MATCH_EXACT: 6}) + + if embedding_model() is None: + with self.assertRaises(RuntimeError): + import_embeddings(self.target_project, directory, execute=True) + return + report = import_embeddings(self.target_project, directory, execute=True) + self.assertEqual(report.written, 6) + again = import_embeddings(self.target_project, directory, execute=True) + self.assertEqual((again.written, again.skipped_existing), (0, 6)) From dd13e08da7009e20afddd18ebfbdbfbf6118ab2a Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Thu, 1 Oct 2026 19:06:05 -0700 Subject: [PATCH 2/5] refactor: let verify_grouping take the time of a confirmation made elsewhere A replayed confirmation must keep its original reviewer time. Instead of stamping now and overwriting the field afterwards, verify_grouping accepts an optional timestamp, so the import goes through the same function the review interface uses and whatever that function records (today the cached fields, later a review row as well) carries the original time. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01C7Xf6VPbwWtTumhjjF15g8 --- ami/main/models_future/tracks.py | 7 ++++--- ami/main/models_future/validated_occurrences.py | 10 ++++------ ami/main/test_validated_occurrences.py | 13 +++++++++++-- 3 files changed, 19 insertions(+), 11 deletions(-) diff --git a/ami/main/models_future/tracks.py b/ami/main/models_future/tracks.py index 176fe08b5..bed815ca5 100644 --- a/ami/main/models_future/tracks.py +++ b/ami/main/models_future/tracks.py @@ -614,14 +614,15 @@ def add_detections(target: Occurrence, detections: Iterable[Detection]) -> Occur @transaction.atomic -def verify_grouping(occurrence: Occurrence, user: User) -> Occurrence: +def verify_grouping(occurrence: Occurrence, user: User, timestamp: datetime.datetime | None = None) -> Occurrence: """Record that a person confirmed this occurrence holds the right detections. This is the label the tracking methods are scored against, so it is deliberately - an explicit act — no operation in this module sets it as a side effect. + an explicit act — no operation in this module sets it as a side effect. ``timestamp`` + is for replaying a confirmation made elsewhere; a live confirmation is stamped now. """ _lock_for_edit(occurrence) - occurrence.grouping_verified_at = timezone.now() + occurrence.grouping_verified_at = timestamp or timezone.now() occurrence.grouping_verified_by = user occurrence.save(update_fields=["grouping_verified_at", "grouping_verified_by"]) return occurrence diff --git a/ami/main/models_future/validated_occurrences.py b/ami/main/models_future/validated_occurrences.py index 0195a45d2..153e5f6f3 100644 --- a/ami/main/models_future/validated_occurrences.py +++ b/ami/main/models_future/validated_occurrences.py @@ -410,13 +410,11 @@ def _parse_datetime(value: str) -> datetime.datetime: def confirm_grouping_as_of(occurrence: Occurrence, user: User, verified_at: datetime.datetime) -> None: """Record a grouping confirmation made by ``user`` at ``verified_at`` on another database. - The one place the import writes a confirmation. ``verify_grouping`` stamps the current - time, so the exported time is written over it afterwards to keep the provenance. When - confirmations become review rows of their own, this function changes and nothing else does. + The one place the import writes a confirmation, through the same function the review + interface uses, so whatever ``verify_grouping`` records (today the cached fields, later a + review row as well) is recorded here too, under the original time. """ - verify_grouping(occurrence, user) - occurrence.grouping_verified_at = verified_at - Occurrence.objects.filter(pk=occurrence.pk).update(grouping_verified_at=verified_at) + verify_grouping(occurrence, user, timestamp=verified_at) def _is_confirmed_as(occurrence: Occurrence, user: User | None, verified_at: datetime.datetime | None) -> bool: diff --git a/ami/main/test_validated_occurrences.py b/ami/main/test_validated_occurrences.py index d0000306e..33f333a86 100644 --- a/ami/main/test_validated_occurrences.py +++ b/ami/main/test_validated_occurrences.py @@ -96,8 +96,7 @@ def _make_source_occurrences(self) -> tuple[Occurrence, Occurrence]: Identification.objects.filter(occurrence=track, taxon=self.taxon).update( created_at=IDENTIFIED_AT, updated_at=IDENTIFIED_AT ) - verify_grouping(track, self.reviewer) - Occurrence.objects.filter(pk=track.pk).update(grouping_verified_at=CONFIRMED_AT) + verify_grouping(track, self.reviewer, timestamp=CONFIRMED_AT) identified = self._occurrence_with_boxes(self.source_project, self.source_captures[4:], OTHER_BOX) Identification.objects.create(occurrence=identified, user=self.reviewer, taxon=self.taxon) @@ -159,6 +158,16 @@ def test_a_round_trip_rebuilds_the_track_under_the_original_reviewer_and_time(se self.assertEqual(rebuilt.grouping_verified_by, self.reviewer) self.assertEqual(Occurrence.objects.filter(project=self.target_project).count(), 3, "4 singles became 1 track") + def test_verify_grouping_accepts_the_time_of_a_confirmation_made_elsewhere(self): + occurrence = self._occurrence_with_boxes(self.target_project, self.target_captures[:1], TRACK_BOX) + + verify_grouping(occurrence, self.reviewer, timestamp=CONFIRMED_AT) + + occurrence.refresh_from_db() + self.assertEqual( + (occurrence.grouping_verified_at, occurrence.grouping_verified_by), (CONFIRMED_AT, self.reviewer) + ) + def test_identifications_keep_their_user_time_and_withdrawn_state(self): self._populate_target() From b5eaa77d6f0dab9230966b0ea1596447acf7a79a Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Thu, 1 Oct 2026 21:51:57 -0700 Subject: [PATCH 3/5] fix: export vectors from both stores and re-confirm when the review row is missing Two problems found on a trial merge with the model-outputs branch. The vector export read only DetectionEmbedding when that table existed and so wrote nothing for data whose vectors live on the classifications; it now carries both stores, one matrix file per source since their lengths differ, and the import writes each back to its own home (embedding rows, or features_2048 on the matching classification, reported as skipped when the target has no such classification). The "already confirmed" check compared only the cached grouping_verified_at/by, so an occurrence confirmed before reviews existed never got its ValidationReview row; where that model exists the check now also requires a standing grouping review by that person at that time, and re-confirms through verify_grouping when it is missing. On a branch without reviews the check stays cache-only. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01C7Xf6VPbwWtTumhjjF15g8 --- .../management/commands/import_embeddings.py | 4 +- ami/main/models_future/embedding_transfer.py | 333 ++++++++++++------ .../models_future/validated_occurrences.py | 43 ++- ami/main/test_validated_occurrences.py | 160 +++++++-- 4 files changed, 402 insertions(+), 138 deletions(-) diff --git a/ami/main/management/commands/import_embeddings.py b/ami/main/management/commands/import_embeddings.py index 47877b9c4..2561ae807 100644 --- a/ami/main/management/commands/import_embeddings.py +++ b/ami/main/management/commands/import_embeddings.py @@ -63,12 +63,14 @@ def handle(self, *args, **options): summary = report.summary() mode = "Applied" if options["execute"] else "Dry run" self.stdout.write( - f"{mode} on project #{project.pk}: {summary['vectors_total']} vectors, detections {summary['detections']}" + f"{mode} on project #{project.pk}: {summary['vectors_total']} vectors {summary['sources']}, " + f"detections {summary['detections']}" ) if options["execute"]: self.stdout.write( f" written: {summary['written']} skipped (already stored): {summary['skipped_existing']}" f" replaced: {summary['replaced']}" + f" classifier features without a classification on the target: {summary['skipped_no_classification']}" ) else: self.stdout.write("Nothing was changed. Re-run with --execute to write the vectors.") diff --git a/ami/main/models_future/embedding_transfer.py b/ami/main/models_future/embedding_transfer.py index fb312e016..57c4fe3f9 100644 --- a/ami/main/models_future/embedding_transfer.py +++ b/ami/main/models_future/embedding_transfer.py @@ -6,12 +6,18 @@ it; on the way back in, the detections are found again by capture path and box and the rows are written under the same algorithm. -Layout of an export directory: ``vectors.npy`` (one row per vector), ``index.csv`` (one line -per row, the detection's natural key) and ``manifest.json`` (format version, algorithm, -vector key, dimensions, count, dtype). Vectors are read from ``DetectionEmbedding`` rows -when that table exists on this branch, otherwise from the classification vectors -(``Classification.features_2048``) the same algorithm stored, which is all older data has. -Writing always targets ``DetectionEmbedding`` and refuses when the table does not exist. +Vectors live in two stores and an export carries both: ``DetectionEmbedding`` rows (an +extractor's vector for a detection, under a vector key) where that table exists, and the +backbone features a classifier stored on its own ``Classification`` rows +(``features_2048``), which is all older data has. The two have different lengths, so each +source gets its own matrix file; ``index.csv`` says which source every row came from. + +Layout of an export directory: ``vectors..npy`` per source, ``index.csv`` (one line +per vector: source, row in that source's matrix, the detection's natural key) and +``manifest.json`` (format version, algorithm, vector key, per-source count and dimensions). +On import, embeddings become ``DetectionEmbedding`` rows and classifier features go back +onto the matching classification when the target has one; a feature vector whose detection +has no classification from that algorithm is reported as skipped. """ from __future__ import annotations @@ -22,6 +28,7 @@ import json import logging import pathlib +from collections import Counter from collections.abc import Iterable, Iterator from typing import Any @@ -43,13 +50,28 @@ EMBEDDINGS_FORMAT = "antenna-detection-embeddings" EMBEDDINGS_VERSION = 1 DEFAULT_VECTOR_KEY = "embedding" -VECTORS_FILE = "vectors.npy" INDEX_FILE = "index.csv" MANIFEST_FILE = "manifest.json" -INDEX_COLUMNS = ["row", "capture_path", "capture_timestamp", "deployment", "detector", "x1", "y1", "x2", "y2"] +INDEX_COLUMNS = [ + "source", + "row", + "capture_path", + "capture_timestamp", + "deployment", + "detector", + "x1", + "y1", + "x2", + "y2", +] MATCH_CHUNK = 5000 WRITE_BATCH = 1000 +# Where a vector was read from, and so where it is written back to. +SOURCE_EMBEDDING = "embedding" # a DetectionEmbedding row +SOURCE_CLASSIFIER_FEATURES = "classifier_features" # Classification.features_2048 +SOURCES = (SOURCE_EMBEDDING, SOURCE_CLASSIFIER_FEATURES) + def embedding_model() -> type[Model] | None: """The ``DetectionEmbedding`` model, or None on a branch that does not have it yet.""" @@ -63,6 +85,10 @@ def _model_field_names(model: type[Model]) -> set[str]: return {field.name for field in model._meta.get_fields()} +def vectors_file(source: str) -> str: + return f"vectors.{source}.npy" + + def resolve_algorithm(reference: str) -> Algorithm: """An algorithm by key, then by name, then by id.""" algorithm = Algorithm.objects.filter(key=reference).first() or Algorithm.objects.filter(name=reference).first() @@ -78,10 +104,10 @@ class EmbeddingManifest: algorithm_key: str algorithm_name: str vector_key: str - dimensions: int + # Per source: {"file": ..., "count": ..., "dimensions": ...} + sources: dict[str, dict] count: int dtype: str - source_store: str project_name: str | None = None exported_at: str = dataclasses.field(default_factory=lambda: datetime.datetime.now().isoformat()) version: int = EMBEDDINGS_VERSION @@ -100,23 +126,19 @@ def read(cls, directory: pathlib.Path) -> EmbeddingManifest: return cls(**data) -def _vector_rows( - project: Project, algorithm: Algorithm, vector_key: str -) -> tuple[str, Iterator[tuple[DetectionKey, Any]]]: - """(store name, iterator of (detection key, vector)) for one algorithm's vectors in a project.""" +def _embedding_rows(project: Project, algorithm: Algorithm, vector_key: str) -> Iterator[tuple[DetectionKey, Any]]: model = embedding_model() - if model is not None: - rows = model.objects.filter(detection__source_image__project=project, algorithm=algorithm) - if "key" in _model_field_names(model): - rows = rows.filter(key=vector_key) - rows = rows.select_related("detection__source_image__deployment", "detection__detection_algorithm").order_by( - "pk" - ) - return "detection_embedding", ( - (DetectionKey.for_detection(row.detection), row.vector) for row in rows.iterator() - ) + if model is None: + return iter(()) + rows = model.objects.filter(detection__source_image__project=project, algorithm=algorithm) + if "key" in _model_field_names(model): + rows = rows.filter(key=vector_key) + rows = rows.select_related("detection__source_image__deployment", "detection__detection_algorithm").order_by("pk") + return ((DetectionKey.for_detection(row.detection), row.vector) for row in rows.iterator()) - # Older data: the classifier's backbone vector lives on the classification. Newest per detection. + +def _classifier_feature_rows(project: Project, algorithm: Algorithm) -> Iterator[tuple[DetectionKey, Any]]: + """The newest feature vector the algorithm stored on each detection's classifications.""" rows = ( Classification.objects.filter( detection__source_image__project=project, algorithm=algorithm, features_2048__isnull=False @@ -125,9 +147,7 @@ def _vector_rows( .order_by("detection_id", "-timestamp", "-pk") .distinct("detection_id") ) - return "classification_features", ( - (DetectionKey.for_detection(row.detection), row.features_2048) for row in rows.iterator() - ) + return ((DetectionKey.for_detection(row.detection), row.features_2048) for row in rows.iterator()) def export_embeddings( @@ -137,64 +157,87 @@ def export_embeddings( vector_key: str = DEFAULT_VECTOR_KEY, dtype: str = "float32", ) -> EmbeddingManifest: - """Write one algorithm's vectors for ``project`` to ``directory``; returns the manifest written.""" + """Write one algorithm's vectors for ``project`` to ``directory``, from both stores. + + A detection with a vector in both stores appears twice in the index, once per source; + within a source each detection appears once. Returns the manifest written. + """ directory.mkdir(parents=True, exist_ok=True) - store, rows = _vector_rows(project, algorithm, vector_key) - vectors: list[np.ndarray] = [] + readers = { + SOURCE_EMBEDDING: _embedding_rows(project, algorithm, vector_key), + SOURCE_CLASSIFIER_FEATURES: _classifier_feature_rows(project, algorithm), + } + sources: dict[str, dict] = {} with (directory / INDEX_FILE).open("w", newline="") as index_file: writer = csv.writer(index_file) writer.writerow(INDEX_COLUMNS) - for row, (key, vector) in enumerate(rows): - # pgvector returns lists on some versions and arrays on others; np.asarray takes both. - vectors.append(np.asarray(list(vector), dtype=dtype)) - writer.writerow( - [ - row, - key.capture_path, - key.capture_timestamp or "", - key.deployment or "", - key.detector or "", - *key.bbox, - ] - ) - if vectors: - lengths = {len(v) for v in vectors} - if len(lengths) > 1: - raise ValueError( - f"Algorithm {algorithm.key} stored vectors of several lengths {sorted(lengths)}; export one at a time." - ) - matrix = np.stack(vectors) - else: - matrix = np.zeros((0, 0), dtype=dtype) - np.save(directory / VECTORS_FILE, matrix) + for source, rows in readers.items(): + vectors: list[np.ndarray] = [] + for row, (key, vector) in enumerate(rows): + # pgvector returns lists on some versions and arrays on others; np.asarray takes both. + vectors.append(np.asarray(list(vector), dtype=dtype)) + writer.writerow( + [ + source, + row, + key.capture_path, + key.capture_timestamp or "", + key.deployment or "", + key.detector or "", + *key.bbox, + ] + ) + if not vectors: + continue + lengths = {len(v) for v in vectors} + if len(lengths) > 1: + raise ValueError( + f"Algorithm {algorithm.key} stored {source} vectors of several lengths {sorted(lengths)}; " + "export one at a time." + ) + matrix = np.stack(vectors) + np.save(directory / vectors_file(source), matrix) + sources[source] = {"file": vectors_file(source), "count": len(vectors), "dimensions": int(matrix.shape[1])} manifest = EmbeddingManifest( algorithm_key=algorithm.key, algorithm_name=algorithm.name, vector_key=vector_key, - dimensions=int(matrix.shape[1]) if matrix.size else 0, - count=int(matrix.shape[0]), + sources=sources, + count=sum(entry["count"] for entry in sources.values()), dtype=dtype, - source_store=store, project_name=project.name, ) (directory / MANIFEST_FILE).write_text(json.dumps(manifest.as_dict(), indent=1)) return manifest -def read_index(directory: pathlib.Path) -> list[DetectionKey]: - keys = [] +@dataclasses.dataclass(frozen=True) +class IndexEntry: + source: str + row: int + key: DetectionKey + + +def read_index(directory: pathlib.Path) -> list[IndexEntry]: + entries = [] with (directory / INDEX_FILE).open(newline="") as index_file: for line in csv.DictReader(index_file): - keys.append( - DetectionKey( - capture_path=line["capture_path"], - bbox=(float(line["x1"]), float(line["y1"]), float(line["x2"]), float(line["y2"])), - capture_timestamp=line["capture_timestamp"] or None, - deployment=line["deployment"] or None, - detector=line["detector"] or None, + if line["source"] not in SOURCES: + raise ValueError(f"Unknown vector source {line['source']!r} in {INDEX_FILE}.") + entries.append( + IndexEntry( + source=line["source"], + row=int(line["row"]), + key=DetectionKey( + capture_path=line["capture_path"], + bbox=(float(line["x1"]), float(line["y1"]), float(line["x2"]), float(line["y2"])), + capture_timestamp=line["capture_timestamp"] or None, + deployment=line["deployment"] or None, + detector=line["detector"] or None, + ), ) ) - return keys + return entries def _chunks(items: list, size: int) -> Iterable[list]: @@ -204,73 +247,64 @@ def _chunks(items: list, size: int) -> Iterable[list]: @dataclasses.dataclass class EmbeddingImportReport: + entries: list[IndexEntry] matches: list[DetectionMatch] - written: int = 0 - skipped_existing: int = 0 - replaced: int = 0 execute: bool = False + written: Counter = dataclasses.field(default_factory=Counter) # per source + skipped_existing: Counter = dataclasses.field(default_factory=Counter) # per source + replaced: Counter = dataclasses.field(default_factory=Counter) # per source + # Classifier features whose detection has no classification from the algorithm on the target. + skipped_no_classification: int = 0 def summary(self) -> dict: - from collections import Counter - - counts = Counter(match.status for match in self.matches) return { "mode": "execute" if self.execute else "dry-run", - "vectors_total": len(self.matches), - "detections": dict(counts), - "written": self.written, - "skipped_existing": self.skipped_existing, - "replaced": self.replaced, + "vectors_total": len(self.entries), + "sources": dict(Counter(entry.source for entry in self.entries)), + "detections": dict(Counter(match.status for match in self.matches)), + "written": dict(self.written), + "skipped_existing": dict(self.skipped_existing), + "replaced": dict(self.replaced), + "skipped_no_classification": self.skipped_no_classification, } -def import_embeddings( - project: Project, - directory: pathlib.Path, - execute: bool = False, - iou_threshold: float = DEFAULT_IOU_THRESHOLD, - replace: bool = False, - algorithm: Algorithm | None = None, -) -> EmbeddingImportReport: - """Write the vectors in ``directory`` as ``DetectionEmbedding`` rows on ``project``'s detections. +def _load_matrix(directory: pathlib.Path, manifest: EmbeddingManifest, source: str) -> np.ndarray: + entry = manifest.sources.get(source) + if entry is None: + raise ValueError(f"The manifest lists no {source} vectors but {INDEX_FILE} has rows of that source.") + matrix = np.load(directory / entry["file"], mmap_mode="r") + if matrix.shape[0] != entry["count"]: + raise ValueError(f"{entry['file']} has {matrix.shape[0]} rows but the manifest says {entry['count']}.") + return matrix - Detections are found by natural key. A detection that already has a vector from the - same algorithm (and key) is skipped unless ``replace`` is set. A dry run only matches. - """ - manifest = EmbeddingManifest.read(directory) - algorithm = algorithm or resolve_algorithm(manifest.algorithm_key) - keys = read_index(directory) - if len(keys) != manifest.count: - raise ValueError(f"{INDEX_FILE} has {len(keys)} rows but the manifest says {manifest.count}.") - matches: list[DetectionMatch] = [] - for chunk in _chunks(keys, MATCH_CHUNK): - matches.extend(match_detections(project, chunk, iou_threshold)) - report = EmbeddingImportReport(matches=matches, execute=execute) - if not execute: - return report +def _write_embeddings( + project: Project, + algorithm: Algorithm, + vector_key: str, + matrix: np.ndarray, + found: list[tuple[int, int]], + replace: bool, + report: EmbeddingImportReport, +) -> None: model = embedding_model() if model is None: - raise RuntimeError("This branch has no DetectionEmbedding table; vectors cannot be imported here.") + raise RuntimeError("This branch has no DetectionEmbedding table; embedding vectors cannot be imported here.") fields = _model_field_names(model) - matrix = np.load(directory / VECTORS_FILE, mmap_mode="r") - if matrix.shape[0] != manifest.count: - raise ValueError(f"{VECTORS_FILE} has {matrix.shape[0]} rows but the manifest says {manifest.count}.") - - found = [(row, match.detection_id) for row, match in enumerate(matches) if match.found] for chunk in _chunks(found, WRITE_BATCH): detection_ids = [detection_id for _, detection_id in chunk] existing = model.objects.filter(detection_id__in=detection_ids, algorithm=algorithm) if "key" in fields: - existing = existing.filter(key=manifest.vector_key) + existing = existing.filter(key=vector_key) existing_ids = set(existing.values_list("detection_id", flat=True)) if replace and existing_ids: existing.delete() - report.replaced += len(existing_ids) + report.replaced[SOURCE_EMBEDDING] += len(existing_ids) rows = [] for row, detection_id in chunk: if detection_id in existing_ids and not replace: - report.skipped_existing += 1 + report.skipped_existing[SOURCE_EMBEDDING] += 1 continue values: dict[str, Any] = { "detection_id": detection_id, @@ -279,10 +313,85 @@ def import_embeddings( } # Fields the settled schema adds; absent on the draft table. if "key" in fields: - values["key"] = manifest.vector_key + values["key"] = vector_key if "project" in fields: values["project"] = project rows.append(model(**values)) model.objects.bulk_create(rows, batch_size=WRITE_BATCH) - report.written += len(rows) + report.written[SOURCE_EMBEDDING] += len(rows) + + +def _write_classifier_features( + algorithm: Algorithm, + matrix: np.ndarray, + found: list[tuple[int, int]], + replace: bool, + report: EmbeddingImportReport, +) -> None: + """Put each feature vector back on the newest classification the algorithm made of its detection.""" + for chunk in _chunks(found, WRITE_BATCH): + detection_ids = [detection_id for _, detection_id in chunk] + newest: dict[int, Classification] = {} + for classification in ( + Classification.objects.filter(detection_id__in=detection_ids, algorithm=algorithm) + .order_by("detection_id", "-timestamp", "-pk") + .distinct("detection_id") + .only("pk", "detection_id", "features_2048") + ): + newest[classification.detection_id] = classification + to_update = [] + for row, detection_id in chunk: + classification = newest.get(detection_id) + if classification is None: + report.skipped_no_classification += 1 + continue + if classification.features_2048 is not None: + if not replace: + report.skipped_existing[SOURCE_CLASSIFIER_FEATURES] += 1 + continue + report.replaced[SOURCE_CLASSIFIER_FEATURES] += 1 + classification.features_2048 = matrix[row].tolist() + to_update.append(classification) + Classification.objects.bulk_update(to_update, ["features_2048"], batch_size=WRITE_BATCH) + report.written[SOURCE_CLASSIFIER_FEATURES] += len(to_update) + + +def import_embeddings( + project: Project, + directory: pathlib.Path, + execute: bool = False, + iou_threshold: float = DEFAULT_IOU_THRESHOLD, + replace: bool = False, + algorithm: Algorithm | None = None, +) -> EmbeddingImportReport: + """Put the vectors in ``directory`` back on ``project``'s detections, each in its own store. + + Detections are found by natural key. A vector the target already holds (an embedding + from the same algorithm and key, or a classification that already has features) is + skipped unless ``replace`` is set. A dry run only matches. + """ + manifest = EmbeddingManifest.read(directory) + algorithm = algorithm or resolve_algorithm(manifest.algorithm_key) + entries = read_index(directory) + if len(entries) != manifest.count: + raise ValueError(f"{INDEX_FILE} has {len(entries)} rows but the manifest says {manifest.count}.") + matches: list[DetectionMatch] = [] + for chunk in _chunks(entries, MATCH_CHUNK): + matches.extend(match_detections(project, [entry.key for entry in chunk], iou_threshold)) + report = EmbeddingImportReport(entries=entries, matches=matches, execute=execute) + if not execute: + return report + + found_by_source: dict[str, list[tuple[int, int]]] = {source: [] for source in SOURCES} + for entry, match in zip(entries, matches): + if match.found: + found_by_source[entry.source].append((entry.row, match.detection_id)) + if found_by_source[SOURCE_EMBEDDING]: + matrix = _load_matrix(directory, manifest, SOURCE_EMBEDDING) + _write_embeddings( + project, algorithm, manifest.vector_key, matrix, found_by_source[SOURCE_EMBEDDING], replace, report + ) + if found_by_source[SOURCE_CLASSIFIER_FEATURES]: + matrix = _load_matrix(directory, manifest, SOURCE_CLASSIFIER_FEATURES) + _write_classifier_features(algorithm, matrix, found_by_source[SOURCE_CLASSIFIER_FEATURES], replace, report) return report diff --git a/ami/main/models_future/validated_occurrences.py b/ami/main/models_future/validated_occurrences.py index 153e5f6f3..26e02c058 100644 --- a/ami/main/models_future/validated_occurrences.py +++ b/ami/main/models_future/validated_occurrences.py @@ -30,6 +30,7 @@ from collections import Counter from collections.abc import Iterable +from django.apps import apps from django.db import transaction from django.db.models import Prefetch @@ -417,12 +418,44 @@ def confirm_grouping_as_of(occurrence: Occurrence, user: User, verified_at: date verify_grouping(occurrence, user, timestamp=verified_at) +def review_model(): + """The ``ValidationReview`` model, or None on a branch that does not have it yet.""" + try: + return apps.get_model("main", "ValidationReview") + except LookupError: + return None + + +def _has_current_grouping_review(occurrence: Occurrence, user: User, verified_at: datetime.datetime) -> bool | None: + """Whether a standing grouping review by ``user`` at ``verified_at`` exists; None where reviews do not exist.""" + model = review_model() + if model is None: + return None + return model.objects.filter( + occurrence=occurrence, + aspect="grouping", + user=user, + timestamp=verified_at, + is_current=True, + withdrawn=False, + ).exists() + + def _is_confirmed_as(occurrence: Occurrence, user: User | None, verified_at: datetime.datetime | None) -> bool: - return ( - occurrence.grouping_verified_at == verified_at - and user is not None - and occurrence.grouping_verified_by_id == user.pk - ) + """Whether the occurrence already carries this confirmation in full. + + The cached ``grouping_verified_at/by`` must match, and where the schema keeps reviews as + rows of their own a standing grouping review by that person at that time must exist too. + A cache that matches without its review (an occurrence confirmed before reviews existed) + is not "confirmed as", so the import re-confirms it through ``verify_grouping`` and the + review gets written. + """ + if user is None or verified_at is None: + return False + if occurrence.grouping_verified_at != verified_at or occurrence.grouping_verified_by_id != user.pk: + return False + has_review = _has_current_grouping_review(occurrence, user, verified_at) + return has_review is None or has_review def _holder_counts(detection_ids: Iterable[int]) -> Counter: diff --git a/ami/main/test_validated_occurrences.py b/ami/main/test_validated_occurrences.py index 33f333a86..e7225651b 100644 --- a/ami/main/test_validated_occurrences.py +++ b/ami/main/test_validated_occurrences.py @@ -11,11 +11,13 @@ import datetime import pathlib import tempfile +from unittest import mock import numpy as np from django.test import TestCase from ami.main.models import ( + Classification, Detection, Identification, Occurrence, @@ -27,11 +29,14 @@ ) from ami.main.models_future.detection_matching import MATCH_EXACT, MATCH_IOU, MATCH_NO_CANDIDATE, bbox_iou from ami.main.models_future.embedding_transfer import ( + SOURCE_CLASSIFIER_FEATURES, + SOURCE_EMBEDDING, EmbeddingManifest, embedding_model, export_embeddings, import_embeddings, read_index, + vectors_file, ) from ami.main.models_future.tracks import verify_grouping from ami.main.models_future.validated_occurrences import ( @@ -43,6 +48,7 @@ build_bundle, import_bundle, occurrence_detection_keys, + review_model, ) from ami.ml.models import Algorithm from ami.tests.fixtures.main import create_taxa, setup_test_project @@ -289,14 +295,19 @@ def test_a_missing_reviewer_or_taxon_is_reported_and_skipped(self): class TestEmbeddingTransfer(ReplayTestCase): + """Vectors from both stores travel by detection key and land back in their own store.""" + def setUp(self) -> None: if not pgvector_is_available(): self.skipTest("pgvector is not installed in this database") super().setUp() self.algorithm = Algorithm.objects.create(key="replay-test-backbone", name="Replay Test Backbone") rng = np.random.default_rng(7) - self.vectors = {} - for detection in Detection.objects.filter(source_image__project=self.source_project): + self.features = {} + self.embeddings = {} + for index, detection in enumerate( + Detection.objects.filter(source_image__project=self.source_project).order_by("source_image__timestamp") + ): vector = rng.standard_normal(2048).astype(np.float32) detection.classifications.create( taxon=self.taxon, @@ -305,31 +316,140 @@ def setUp(self) -> None: timestamp=detection.timestamp, features_2048=vector.tolist(), ) - self.vectors[detection.source_image.path] = vector + self.features[detection.source_image.path] = vector + # Where the embeddings table exists, half the detections also carry an extractor vector. + if embedding_model() is not None and index < 3: + embedding = rng.standard_normal(16).astype(np.float32) + self._store_embedding(detection, embedding) + self.embeddings[detection.source_image.path] = embedding + + def _store_embedding(self, detection, vector) -> None: + model = embedding_model() + fields = {f.name for f in model._meta.get_fields()} + values = {"detection": detection, "algorithm": self.algorithm, "vector": vector.tolist()} + if "key" in fields: + values["key"] = "embedding" + if "project" in fields: + values["project"] = detection.source_image.project + model.objects.create(**values) + + def _target_classifications(self) -> None: + for detection in Detection.objects.filter(source_image__project=self.target_project): + detection.classifications.create( + taxon=self.taxon, score=0.4, algorithm=self.algorithm, timestamp=detection.timestamp + ) - def test_vectors_round_trip_by_detection_key(self): + def test_classifier_features_are_exported_and_put_back_on_the_classification(self): self._populate_target() with tempfile.TemporaryDirectory() as directory: directory = pathlib.Path(directory) manifest = export_embeddings(self.source_project, self.algorithm, directory) + self.assertEqual(manifest.sources[SOURCE_CLASSIFIER_FEATURES]["count"], 6) + self.assertEqual(manifest.sources[SOURCE_CLASSIFIER_FEATURES]["dimensions"], 2048) + self.assertEqual(EmbeddingManifest.read(directory).count, manifest.count) + matrix = np.load(directory / vectors_file(SOURCE_CLASSIFIER_FEATURES)) + for entry in read_index(directory): + if entry.source == SOURCE_CLASSIFIER_FEATURES: + np.testing.assert_array_equal(matrix[entry.row], self.features[entry.key.capture_path]) + + # No classification on the target yet: nowhere to put the features, so they are skipped, not lost. + report = import_embeddings(self.target_project, directory, execute=True) + self.assertEqual(report.summary()["detections"][MATCH_EXACT], manifest.count) + self.assertEqual(report.skipped_no_classification, 6) + self.assertEqual(report.written[SOURCE_CLASSIFIER_FEATURES], 0) + + self._target_classifications() + report = import_embeddings(self.target_project, directory, execute=True) + self.assertEqual(report.written[SOURCE_CLASSIFIER_FEATURES], 6) + stored = Classification.objects.get( + detection__source_image=self.target_captures[0], detection__source_image__project=self.target_project + ) + np.testing.assert_array_equal( + np.asarray(list(stored.features_2048), dtype=np.float32), self.features[self.target_captures[0].path] + ) + + again = import_embeddings(self.target_project, directory, execute=True) self.assertEqual( - (manifest.count, manifest.dimensions, manifest.algorithm_key), (6, 2048, "replay-test-backbone") + (again.written[SOURCE_CLASSIFIER_FEATURES], again.skipped_existing[SOURCE_CLASSIFIER_FEATURES]), (0, 6) ) - self.assertEqual(EmbeddingManifest.read(directory).count, 6) - keys = read_index(directory) - matrix = np.load(directory / "vectors.npy") - for row, key in enumerate(keys): - np.testing.assert_array_equal(matrix[row], self.vectors[key.capture_path]) - - report = import_embeddings(self.target_project, directory, execute=False) - self.assertEqual(report.summary()["detections"], {MATCH_EXACT: 6}) - - if embedding_model() is None: - with self.assertRaises(RuntimeError): - import_embeddings(self.target_project, directory, execute=True) - return + + def test_embeddings_are_exported_beside_the_features_and_rewritten_as_rows(self): + if embedding_model() is None: + self.skipTest("This branch has no DetectionEmbedding table") + self._populate_target() + with tempfile.TemporaryDirectory() as directory: + directory = pathlib.Path(directory) + manifest = export_embeddings(self.source_project, self.algorithm, directory) + + self.assertEqual(manifest.sources[SOURCE_EMBEDDING]["count"], 3) + self.assertEqual(manifest.sources[SOURCE_EMBEDDING]["dimensions"], 16) + self.assertEqual( + manifest.sources[SOURCE_CLASSIFIER_FEATURES]["count"], 6, "both stores, no double counting" + ) + self.assertEqual(manifest.count, 9) + report = import_embeddings(self.target_project, directory, execute=True) - self.assertEqual(report.written, 6) + self.assertEqual(report.written[SOURCE_EMBEDDING], 3) + rows = embedding_model().objects.filter( + detection__source_image__project=self.target_project, algorithm=self.algorithm + ) + self.assertEqual(rows.count(), 3) + row = rows.get(detection__source_image=self.target_captures[0]) + np.testing.assert_array_equal( + np.asarray(list(row.vector), dtype=np.float32), self.embeddings[self.target_captures[0].path] + ) + again = import_embeddings(self.target_project, directory, execute=True) - self.assertEqual((again.written, again.skipped_existing), (0, 6)) + self.assertEqual((again.written[SOURCE_EMBEDDING], again.skipped_existing[SOURCE_EMBEDDING]), (0, 3)) + + +class TestConfirmationNeedsItsReview(ReplayTestCase): + """A cached confirmation without its review row is re-confirmed so the review gets written.""" + + def test_a_cache_only_confirmation_is_reconfirmed_through_verify_grouping(self): + self._populate_target() + self._replay(execute=True) + rebuilt = self._target_track() + + with mock.patch( + "ami.main.models_future.validated_occurrences._has_current_grouping_review", return_value=False + ): + with mock.patch( + "ami.main.models_future.validated_occurrences.verify_grouping", wraps=verify_grouping + ) as spied: + _, report = self._replay(execute=True) + + spied.assert_called_once_with(rebuilt, self.reviewer, timestamp=CONFIRMED_AT) + track_outcome = next(o for o in report.outcomes if o.ref == str(self.track.pk)) + self.assertEqual((track_outcome.outcome, track_outcome.confirmed), (OUTCOME_APPLIED, True)) + self.assertEqual(self._target_track().grouping_verified_at, CONFIRMED_AT) + + def test_a_confirmation_with_its_review_is_left_alone(self): + self._populate_target() + self._replay(execute=True) + + with mock.patch("ami.main.models_future.validated_occurrences.verify_grouping") as spied: + _, report = self._replay(execute=True) + + spied.assert_not_called() + self.assertEqual(report.outcome_counts(), {OUTCOME_UNCHANGED: 2}) + + def test_on_a_schema_with_reviews_the_review_row_is_written(self): + if review_model() is None: + self.skipTest("This branch has no ValidationReview table") + self._populate_target() + self._replay(execute=True) + rebuilt = self._target_track() + reviews = review_model().objects.filter(occurrence=rebuilt, aspect="grouping", user=self.reviewer) + self.assertEqual(reviews.filter(is_current=True, withdrawn=False, timestamp=CONFIRMED_AT).count(), 1) + + # An occurrence confirmed before reviews existed: the cache is set, no review row. + reviews.delete() + _, report = self._replay(execute=True) + + self.assertEqual(reviews.filter(is_current=True, withdrawn=False, timestamp=CONFIRMED_AT).count(), 1) + track_outcome = next(o for o in report.outcomes if o.ref == str(self.track.pk)) + self.assertTrue(track_outcome.confirmed) + _, report = self._replay(execute=True) + self.assertEqual(report.outcome_counts(), {OUTCOME_UNCHANGED: 2}) From bbfaddd65a534cba18f376f9904308b0111f6a91 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Thu, 1 Oct 2026 21:53:31 -0700 Subject: [PATCH 4/5] fix: match each detection key once when a detection has a vector in both stores The matcher claims a target detection for one key, which is right between two boxes and wrong between two index rows of the same box. The vector import now matches the distinct keys and maps every row to its key's match. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01C7Xf6VPbwWtTumhjjF15g8 --- ami/main/models_future/embedding_transfer.py | 12 +++++++++--- ami/main/test_validated_occurrences.py | 7 +++++-- 2 files changed, 14 insertions(+), 5 deletions(-) diff --git a/ami/main/models_future/embedding_transfer.py b/ami/main/models_future/embedding_transfer.py index 57c4fe3f9..e0a643c3f 100644 --- a/ami/main/models_future/embedding_transfer.py +++ b/ami/main/models_future/embedding_transfer.py @@ -375,9 +375,15 @@ def import_embeddings( entries = read_index(directory) if len(entries) != manifest.count: raise ValueError(f"{INDEX_FILE} has {len(entries)} rows but the manifest says {manifest.count}.") - matches: list[DetectionMatch] = [] - for chunk in _chunks(entries, MATCH_CHUNK): - matches.extend(match_detections(project, [entry.key for entry in chunk], iou_threshold)) + # A detection with a vector in both stores appears twice in the index under one key. Match each + # key once: the matcher claims a target detection for one key only, which is right between + # different boxes and wrong between two rows of the same box. + unique_keys = list(dict.fromkeys(entry.key for entry in entries)) + match_by_key: dict[DetectionKey, DetectionMatch] = {} + for chunk in _chunks(unique_keys, MATCH_CHUNK): + for match in match_detections(project, chunk, iou_threshold): + match_by_key[match.key] = match + matches = [match_by_key[entry.key] for entry in entries] report = EmbeddingImportReport(entries=entries, matches=matches, execute=execute) if not execute: return report diff --git a/ami/main/test_validated_occurrences.py b/ami/main/test_validated_occurrences.py index e7225651b..5552e1546 100644 --- a/ami/main/test_validated_occurrences.py +++ b/ami/main/test_validated_occurrences.py @@ -396,8 +396,11 @@ def test_embeddings_are_exported_beside_the_features_and_rewritten_as_rows(self) ) self.assertEqual(rows.count(), 3) row = rows.get(detection__source_image=self.target_captures[0]) - np.testing.assert_array_equal( - np.asarray(list(row.vector), dtype=np.float32), self.embeddings[self.target_captures[0].path] + # The settled schema stores half precision; compare at that precision. + np.testing.assert_allclose( + np.asarray(list(row.vector), dtype=np.float32), + self.embeddings[self.target_captures[0].path], + rtol=2e-3, ) again = import_embeddings(self.target_project, directory, execute=True) From d374bf96ab7e4e3f60f833f5f67ed34e556f01c9 Mon Sep 17 00:00:00 2001 From: Michael Bunsen Date: Fri, 2 Oct 2026 13:15:09 -0700 Subject: [PATCH 5/5] refactor: let the vector and review code find out what the schema has instead of assuming it The outputs stack lands in pieces (results and reviews, then embeddings, then tracking), so a branch may have any subset of the tables and columns this code reads. The vector transfer now checks for the classification feature column as it already did for the embeddings table: an export notes a store it could not read in its manifest, and an import counts rows it has nowhere to put instead of failing. Embedding rows are written through the model's own insert-mostly writer where it has one. The confirmation check requires the grouping aspect to exist on the review model, not just the model, and falls back to the cached fields otherwise. The vector tests move to their own module so they run on a branch without the tracking code. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01C7Xf6VPbwWtTumhjjF15g8 --- .../management/commands/import_embeddings.py | 4 + ami/main/models_future/embedding_transfer.py | 98 +++++--- .../models_future/validated_occurrences.py | 37 ++- ami/main/test_embedding_transfer.py | 219 ++++++++++++++++++ ami/main/test_validated_occurrences.py | 158 +++---------- 5 files changed, 346 insertions(+), 170 deletions(-) create mode 100644 ami/main/test_embedding_transfer.py diff --git a/ami/main/management/commands/import_embeddings.py b/ami/main/management/commands/import_embeddings.py index 2561ae807..209226d21 100644 --- a/ami/main/management/commands/import_embeddings.py +++ b/ami/main/management/commands/import_embeddings.py @@ -72,5 +72,9 @@ def handle(self, *args, **options): f" replaced: {summary['replaced']}" f" classifier features without a classification on the target: {summary['skipped_no_classification']}" ) + if summary["skipped_no_field"]: + self.stdout.write( + self.style.WARNING(f" rows this branch has no table or column for: {summary['skipped_no_field']}") + ) else: self.stdout.write("Nothing was changed. Re-run with --execute to write the vectors.") diff --git a/ami/main/models_future/embedding_transfer.py b/ami/main/models_future/embedding_transfer.py index e0a643c3f..5532927b8 100644 --- a/ami/main/models_future/embedding_transfer.py +++ b/ami/main/models_future/embedding_transfer.py @@ -17,7 +17,9 @@ ``manifest.json`` (format version, algorithm, vector key, per-source count and dimensions). On import, embeddings become ``DetectionEmbedding`` rows and classifier features go back onto the matching classification when the target has one; a feature vector whose detection -has no classification from that algorithm is reported as skipped. +has no classification from that algorithm is reported as skipped. Both stores are optional +per branch: an export notes the ones it could not read in its manifest, and an import counts +rows it has nowhere to put instead of failing. """ from __future__ import annotations @@ -34,6 +36,7 @@ import numpy as np from django.apps import apps +from django.core.exceptions import FieldDoesNotExist from django.db.models import Model from ami.main.models import Classification, Project @@ -71,6 +74,7 @@ SOURCE_EMBEDDING = "embedding" # a DetectionEmbedding row SOURCE_CLASSIFIER_FEATURES = "classifier_features" # Classification.features_2048 SOURCES = (SOURCE_EMBEDDING, SOURCE_CLASSIFIER_FEATURES) +CLASSIFIER_FEATURES_FIELD = "features_2048" def embedding_model() -> type[Model] | None: @@ -85,6 +89,15 @@ def _model_field_names(model: type[Model]) -> set[str]: return {field.name for field in model._meta.get_fields()} +def classifier_features_available() -> bool: + """Whether classifications on this branch carry a feature vector (``features_2048``).""" + try: + Classification._meta.get_field(CLASSIFIER_FEATURES_FIELD) + except FieldDoesNotExist: + return False + return True + + def vectors_file(source: str) -> str: return f"vectors.{source}.npy" @@ -108,6 +121,8 @@ class EmbeddingManifest: sources: dict[str, dict] count: int dtype: str + # Sources the exporting branch could not read, with the reason, so an empty source is not a surprise. + skipped_sources: dict[str, str] = dataclasses.field(default_factory=dict) project_name: str | None = None exported_at: str = dataclasses.field(default_factory=lambda: datetime.datetime.now().isoformat()) version: int = EMBEDDINGS_VERSION @@ -163,10 +178,18 @@ def export_embeddings( within a source each detection appears once. Returns the manifest written. """ directory.mkdir(parents=True, exist_ok=True) - readers = { - SOURCE_EMBEDDING: _embedding_rows(project, algorithm, vector_key), - SOURCE_CLASSIFIER_FEATURES: _classifier_feature_rows(project, algorithm), - } + readers: dict[str, Iterator[tuple[DetectionKey, Any]]] = {} + skipped: dict[str, str] = {} + if embedding_model() is not None: + readers[SOURCE_EMBEDDING] = _embedding_rows(project, algorithm, vector_key) + else: + skipped[SOURCE_EMBEDDING] = "this branch has no DetectionEmbedding table" + if classifier_features_available(): + readers[SOURCE_CLASSIFIER_FEATURES] = _classifier_feature_rows(project, algorithm) + else: + skipped[ + SOURCE_CLASSIFIER_FEATURES + ] = f"classifications on this branch have no {CLASSIFIER_FEATURES_FIELD} field" sources: dict[str, dict] = {} with (directory / INDEX_FILE).open("w", newline="") as index_file: writer = csv.writer(index_file) @@ -205,6 +228,7 @@ def export_embeddings( sources=sources, count=sum(entry["count"] for entry in sources.values()), dtype=dtype, + skipped_sources=skipped, project_name=project.name, ) (directory / MANIFEST_FILE).write_text(json.dumps(manifest.as_dict(), indent=1)) @@ -255,6 +279,8 @@ class EmbeddingImportReport: replaced: Counter = dataclasses.field(default_factory=Counter) # per source # Classifier features whose detection has no classification from the algorithm on the target. skipped_no_classification: int = 0 + # Rows of a source the target branch has no column or table for. + skipped_no_field: Counter = dataclasses.field(default_factory=Counter) # per source def summary(self) -> dict: return { @@ -266,6 +292,7 @@ def summary(self) -> dict: "skipped_existing": dict(self.skipped_existing), "replaced": dict(self.replaced), "skipped_no_classification": self.skipped_no_classification, + "skipped_no_field": dict(self.skipped_no_field), } @@ -288,11 +315,33 @@ def _write_embeddings( replace: bool, report: EmbeddingImportReport, ) -> None: + """Store embedding rows, through the model's own insert-mostly writer where it has one. + + ``DetectionEmbedding.objects.store`` (the settled schema) leaves a row holding the same + vector alone and replaces one holding a different vector, so ``replace`` is implied there. + The draft table has no such writer; rows are compared by hand and ``replace`` decides. + """ model = embedding_model() if model is None: raise RuntimeError("This branch has no DetectionEmbedding table; embedding vectors cannot be imported here.") fields = _model_field_names(model) + + def build(row: int, detection_id: int): + values: dict[str, Any] = {"detection_id": detection_id, "algorithm": algorithm, "vector": matrix[row].tolist()} + # Fields the settled schema has; absent on the draft table. + if "key" in fields: + values["key"] = vector_key + if "project" in fields: + values["project"] = project + return model(**values) + + store = getattr(model.objects, "store", None) for chunk in _chunks(found, WRITE_BATCH): + if store is not None: + inserted, unchanged = store(build(row, detection_id) for row, detection_id in chunk) + report.written[SOURCE_EMBEDDING] += inserted + report.skipped_existing[SOURCE_EMBEDDING] += unchanged + continue detection_ids = [detection_id for _, detection_id in chunk] existing = model.objects.filter(detection_id__in=detection_ids, algorithm=algorithm) if "key" in fields: @@ -301,22 +350,8 @@ def _write_embeddings( if replace and existing_ids: existing.delete() report.replaced[SOURCE_EMBEDDING] += len(existing_ids) - rows = [] - for row, detection_id in chunk: - if detection_id in existing_ids and not replace: - report.skipped_existing[SOURCE_EMBEDDING] += 1 - continue - values: dict[str, Any] = { - "detection_id": detection_id, - "algorithm": algorithm, - "vector": matrix[row].tolist(), - } - # Fields the settled schema adds; absent on the draft table. - if "key" in fields: - values["key"] = vector_key - if "project" in fields: - values["project"] = project - rows.append(model(**values)) + rows = [build(row, detection_id) for row, detection_id in chunk if detection_id not in existing_ids or replace] + report.skipped_existing[SOURCE_EMBEDDING] += len(chunk) - len(rows) model.objects.bulk_create(rows, batch_size=WRITE_BATCH) report.written[SOURCE_EMBEDDING] += len(rows) @@ -368,7 +403,8 @@ def import_embeddings( Detections are found by natural key. A vector the target already holds (an embedding from the same algorithm and key, or a classification that already has features) is - skipped unless ``replace`` is set. A dry run only matches. + skipped unless ``replace`` is set. Rows of a source this branch has no table or column + for are counted as ``skipped_no_field``. A dry run only matches. """ manifest = EmbeddingManifest.read(directory) algorithm = algorithm or resolve_algorithm(manifest.algorithm_key) @@ -393,11 +429,17 @@ def import_embeddings( if match.found: found_by_source[entry.source].append((entry.row, match.detection_id)) if found_by_source[SOURCE_EMBEDDING]: - matrix = _load_matrix(directory, manifest, SOURCE_EMBEDDING) - _write_embeddings( - project, algorithm, manifest.vector_key, matrix, found_by_source[SOURCE_EMBEDDING], replace, report - ) + if embedding_model() is None: + report.skipped_no_field[SOURCE_EMBEDDING] += len(found_by_source[SOURCE_EMBEDDING]) + else: + matrix = _load_matrix(directory, manifest, SOURCE_EMBEDDING) + _write_embeddings( + project, algorithm, manifest.vector_key, matrix, found_by_source[SOURCE_EMBEDDING], replace, report + ) if found_by_source[SOURCE_CLASSIFIER_FEATURES]: - matrix = _load_matrix(directory, manifest, SOURCE_CLASSIFIER_FEATURES) - _write_classifier_features(algorithm, matrix, found_by_source[SOURCE_CLASSIFIER_FEATURES], replace, report) + if not classifier_features_available(): + report.skipped_no_field[SOURCE_CLASSIFIER_FEATURES] += len(found_by_source[SOURCE_CLASSIFIER_FEATURES]) + else: + matrix = _load_matrix(directory, manifest, SOURCE_CLASSIFIER_FEATURES) + _write_classifier_features(algorithm, matrix, found_by_source[SOURCE_CLASSIFIER_FEATURES], replace, report) return report diff --git a/ami/main/models_future/validated_occurrences.py b/ami/main/models_future/validated_occurrences.py index 26e02c058..643cec4bd 100644 --- a/ami/main/models_future/validated_occurrences.py +++ b/ami/main/models_future/validated_occurrences.py @@ -58,6 +58,7 @@ BUNDLE_FORMAT = "antenna-validated-occurrences" BUNDLE_VERSION = 1 +GROUPING_ASPECT = "grouping" # -- Bundle schema ------------------------------------------------------------------------ @@ -426,19 +427,35 @@ def review_model(): return None -def _has_current_grouping_review(occurrence: Occurrence, user: User, verified_at: datetime.datetime) -> bool | None: - """Whether a standing grouping review by ``user`` at ``verified_at`` exists; None where reviews do not exist.""" +def grouping_reviews_supported() -> bool: + """Whether this branch records grouping confirmations as reviews. + + The review table arrives before the tracking code that adds its ``grouping`` aspect, so + the table alone does not mean confirmations are kept there. + """ model = review_model() if model is None: + return False + choices = model._meta.get_field("aspect").choices or [] + return GROUPING_ASPECT in {value for value, _label in choices} + + +def _has_current_grouping_review(occurrence: Occurrence, user: User, verified_at: datetime.datetime) -> bool | None: + """Whether a standing grouping review by ``user`` at ``verified_at`` exists; None where grouping reviews do not.""" + if not grouping_reviews_supported(): return None - return model.objects.filter( - occurrence=occurrence, - aspect="grouping", - user=user, - timestamp=verified_at, - is_current=True, - withdrawn=False, - ).exists() + return ( + review_model() + .objects.filter( + occurrence=occurrence, + aspect=GROUPING_ASPECT, + user=user, + timestamp=verified_at, + is_current=True, + withdrawn=False, + ) + .exists() + ) def _is_confirmed_as(occurrence: Occurrence, user: User | None, verified_at: datetime.datetime | None) -> bool: diff --git a/ami/main/test_embedding_transfer.py b/ami/main/test_embedding_transfer.py new file mode 100644 index 000000000..3f48baf40 --- /dev/null +++ b/ami/main/test_embedding_transfer.py @@ -0,0 +1,219 @@ +"""Moving detection feature vectors between projects by detection key. + +Vectors live in two stores that come and go with the schema: ``DetectionEmbedding`` rows +and the ``features_2048`` column on classifications. These tests run on any branch: each +store's test skips where the branch lacks it, and the fallbacks for a missing store are +tested with the store detection patched out. +""" + +import datetime +import pathlib +import tempfile +from unittest import mock + +import numpy as np +from django.db import connection +from django.test import TestCase + +from ami.main.models import ( + Classification, + Detection, + Occurrence, + SourceImage, + Taxon, + TaxonRank, + group_images_into_events, +) +from ami.main.models_future.detection_matching import MATCH_EXACT +from ami.main.models_future.embedding_transfer import ( + SOURCE_CLASSIFIER_FEATURES, + SOURCE_EMBEDDING, + EmbeddingManifest, + classifier_features_available, + embedding_model, + export_embeddings, + import_embeddings, + read_index, + vectors_file, +) +from ami.ml.models import Algorithm +from ami.tests.fixtures.main import create_taxa, setup_test_project + +BOX = [10.0, 10.0, 40.0, 40.0] + + +def pgvector_is_available() -> bool: + with connection.cursor() as cursor: + cursor.execute("SELECT 1 FROM pg_extension WHERE extname = 'vector'") + return cursor.fetchone() is not None + + +class TestEmbeddingTransfer(TestCase): + def setUp(self) -> None: + if not pgvector_is_available(): + self.skipTest("pgvector is not installed in this database") + self.source_project, self.source_deployment = setup_test_project(reuse=False) + self.target_project, self.target_deployment = setup_test_project(reuse=False) + create_taxa(self.source_project) + self.taxon = ( + Taxon.objects.filter(rank=TaxonRank.SPECIES.name, projects=self.source_project).order_by("pk").first() + ) + self.source_captures = self._make_captures(self.source_deployment) + self.target_captures = self._make_captures(self.target_deployment) + self.algorithm = Algorithm.objects.create(key="replay-test-backbone", name="Replay Test Backbone") + rng = np.random.default_rng(7) + self.features: dict[str, np.ndarray] = {} + self.embeddings: dict[str, np.ndarray] = {} + for index, detection in enumerate(self._detections(self.source_project)): + if classifier_features_available(): + vector = rng.standard_normal(2048).astype(np.float32) + detection.classifications.create( + taxon=self.taxon, + score=0.5, + algorithm=self.algorithm, + timestamp=detection.timestamp, + features_2048=vector.tolist(), + ) + self.features[detection.source_image.path] = vector + # Where the embeddings table exists, half the detections also carry an extractor vector. + if embedding_model() is not None and index < 3: + embedding = rng.standard_normal(16).astype(np.float32) + self._store_embedding(detection, embedding) + self.embeddings[detection.source_image.path] = embedding + + def _make_captures(self, deployment) -> list[SourceImage]: + """Six captures one minute apart, each with one detection in its own occurrence.""" + start = datetime.datetime(2024, 6, 1, 22, 0) + captures = [ + SourceImage.objects.create( + deployment=deployment, + project=deployment.project, + timestamp=start + datetime.timedelta(minutes=i), + path=f"replay/capture-{i}.jpg", + width=640, + height=480, + ) + for i in range(6) + ] + group_images_into_events(deployment) + for capture in captures: + capture.refresh_from_db() + occurrence = Occurrence.objects.create( + event=capture.event, deployment=deployment, project=deployment.project + ) + Detection.objects.create( + source_image=capture, timestamp=capture.timestamp, bbox=BOX, occurrence=occurrence + ) + return captures + + def _detections(self, project): + return Detection.objects.filter(source_image__project=project).order_by("source_image__timestamp") + + def _store_embedding(self, detection, vector) -> None: + model = embedding_model() + fields = {f.name for f in model._meta.get_fields()} + values = {"detection": detection, "algorithm": self.algorithm, "vector": vector.tolist()} + if "key" in fields: + values["key"] = "embedding" + if "project" in fields: + values["project"] = detection.source_image.project + model.objects.create(**values) + + def _target_classifications(self) -> None: + for detection in self._detections(self.target_project): + detection.classifications.create( + taxon=self.taxon, score=0.4, algorithm=self.algorithm, timestamp=detection.timestamp + ) + + def test_classifier_features_are_exported_and_put_back_on_the_classification(self): + if not classifier_features_available(): + self.skipTest("Classifications on this branch have no feature vector column") + with tempfile.TemporaryDirectory() as directory: + directory = pathlib.Path(directory) + manifest = export_embeddings(self.source_project, self.algorithm, directory) + + self.assertEqual(manifest.sources[SOURCE_CLASSIFIER_FEATURES]["count"], 6) + self.assertEqual(manifest.sources[SOURCE_CLASSIFIER_FEATURES]["dimensions"], 2048) + self.assertEqual(EmbeddingManifest.read(directory).count, manifest.count) + matrix = np.load(directory / vectors_file(SOURCE_CLASSIFIER_FEATURES)) + for entry in read_index(directory): + if entry.source == SOURCE_CLASSIFIER_FEATURES: + np.testing.assert_array_equal(matrix[entry.row], self.features[entry.key.capture_path]) + + # No classification on the target yet: nowhere to put the features, so they are skipped, not lost. + report = import_embeddings(self.target_project, directory, execute=True) + self.assertEqual(report.summary()["detections"][MATCH_EXACT], manifest.count) + self.assertEqual(report.skipped_no_classification, 6) + self.assertEqual(report.written[SOURCE_CLASSIFIER_FEATURES], 0) + + self._target_classifications() + report = import_embeddings(self.target_project, directory, execute=True) + self.assertEqual(report.written[SOURCE_CLASSIFIER_FEATURES], 6) + stored = Classification.objects.get(detection__source_image=self.target_captures[0]) + # pgvector hands back decimal text parsed as float64; compare at the stored float32 precision. + np.testing.assert_array_equal( + np.asarray(list(stored.features_2048), dtype=np.float32), self.features[self.target_captures[0].path] + ) + + again = import_embeddings(self.target_project, directory, execute=True) + self.assertEqual( + (again.written[SOURCE_CLASSIFIER_FEATURES], again.skipped_existing[SOURCE_CLASSIFIER_FEATURES]), (0, 6) + ) + + def test_embeddings_are_exported_and_rewritten_as_rows(self): + if embedding_model() is None: + self.skipTest("This branch has no DetectionEmbedding table") + with tempfile.TemporaryDirectory() as directory: + directory = pathlib.Path(directory) + manifest = export_embeddings(self.source_project, self.algorithm, directory) + + self.assertEqual(manifest.sources[SOURCE_EMBEDDING]["count"], 3) + self.assertEqual(manifest.sources[SOURCE_EMBEDDING]["dimensions"], 16) + if classifier_features_available(): + self.assertEqual( + manifest.sources[SOURCE_CLASSIFIER_FEATURES]["count"], 6, "both stores, no double counting" + ) + self.assertEqual(manifest.count, 9) + else: + self.assertIn(SOURCE_CLASSIFIER_FEATURES, manifest.skipped_sources) + self.assertEqual(manifest.count, 3) + + report = import_embeddings(self.target_project, directory, execute=True) + self.assertEqual(report.written[SOURCE_EMBEDDING], 3) + rows = embedding_model().objects.filter( + detection__source_image__project=self.target_project, algorithm=self.algorithm + ) + self.assertEqual(rows.count(), 3) + row = rows.get(detection__source_image=self.target_captures[0]) + if {f.name for f in embedding_model()._meta.get_fields()} >= {"key", "project"}: + self.assertEqual((row.key, row.project_id), ("embedding", self.target_project.pk)) + # The settled schema stores half precision; compare at that precision. + np.testing.assert_allclose( + np.asarray(list(row.vector), dtype=np.float32), + self.embeddings[self.target_captures[0].path], + rtol=2e-3, + ) + + again = import_embeddings(self.target_project, directory, execute=True) + self.assertEqual((again.written[SOURCE_EMBEDDING], again.skipped_existing[SOURCE_EMBEDDING]), (0, 3)) + + def test_a_store_the_branch_lacks_is_noted_on_export_and_skipped_on_import(self): + if not classifier_features_available(): + self.skipTest("Needs a branch with classifier features to build the export from") + with tempfile.TemporaryDirectory() as directory: + directory = pathlib.Path(directory) + export_embeddings(self.source_project, self.algorithm, directory) + self._target_classifications() + + with mock.patch( + "ami.main.models_future.embedding_transfer.classifier_features_available", return_value=False + ): + report = import_embeddings(self.target_project, directory, execute=True) + self.assertEqual(report.skipped_no_field[SOURCE_CLASSIFIER_FEATURES], 6) + self.assertEqual(report.written[SOURCE_CLASSIFIER_FEATURES], 0) + + manifest = export_embeddings(self.source_project, self.algorithm, directory / "again") + self.assertNotIn(SOURCE_CLASSIFIER_FEATURES, manifest.sources) + self.assertIn("features_2048", manifest.skipped_sources[SOURCE_CLASSIFIER_FEATURES]) + if embedding_model() is None: + self.assertIn(SOURCE_EMBEDDING, manifest.skipped_sources) diff --git a/ami/main/test_validated_occurrences.py b/ami/main/test_validated_occurrences.py index 5552e1546..77223c980 100644 --- a/ami/main/test_validated_occurrences.py +++ b/ami/main/test_validated_occurrences.py @@ -9,15 +9,11 @@ import dataclasses import datetime -import pathlib -import tempfile from unittest import mock -import numpy as np from django.test import TestCase from ami.main.models import ( - Classification, Detection, Identification, Occurrence, @@ -28,16 +24,6 @@ group_images_into_events, ) from ami.main.models_future.detection_matching import MATCH_EXACT, MATCH_IOU, MATCH_NO_CANDIDATE, bbox_iou -from ami.main.models_future.embedding_transfer import ( - SOURCE_CLASSIFIER_FEATURES, - SOURCE_EMBEDDING, - EmbeddingManifest, - embedding_model, - export_embeddings, - import_embeddings, - read_index, - vectors_file, -) from ami.main.models_future.tracks import verify_grouping from ami.main.models_future.validated_occurrences import ( OUTCOME_APPLIED, @@ -46,13 +32,12 @@ Bundle, ImportOptions, build_bundle, + grouping_reviews_supported, import_bundle, occurrence_detection_keys, review_model, ) -from ami.ml.models import Algorithm from ami.tests.fixtures.main import create_taxa, setup_test_project -from ami.tests.fixtures.tracking import pgvector_is_available CONFIRMED_AT = datetime.datetime(2024, 6, 10, 9, 30, 0) IDENTIFIED_AT = datetime.datetime(2024, 6, 9, 8, 0, 0) @@ -294,119 +279,6 @@ def test_a_missing_reviewer_or_taxon_is_reported_and_skipped(self): self.assertEqual((track_outcome.identifications_applied, track_outcome.identifications_skipped), (1, 1)) -class TestEmbeddingTransfer(ReplayTestCase): - """Vectors from both stores travel by detection key and land back in their own store.""" - - def setUp(self) -> None: - if not pgvector_is_available(): - self.skipTest("pgvector is not installed in this database") - super().setUp() - self.algorithm = Algorithm.objects.create(key="replay-test-backbone", name="Replay Test Backbone") - rng = np.random.default_rng(7) - self.features = {} - self.embeddings = {} - for index, detection in enumerate( - Detection.objects.filter(source_image__project=self.source_project).order_by("source_image__timestamp") - ): - vector = rng.standard_normal(2048).astype(np.float32) - detection.classifications.create( - taxon=self.taxon, - score=0.5, - algorithm=self.algorithm, - timestamp=detection.timestamp, - features_2048=vector.tolist(), - ) - self.features[detection.source_image.path] = vector - # Where the embeddings table exists, half the detections also carry an extractor vector. - if embedding_model() is not None and index < 3: - embedding = rng.standard_normal(16).astype(np.float32) - self._store_embedding(detection, embedding) - self.embeddings[detection.source_image.path] = embedding - - def _store_embedding(self, detection, vector) -> None: - model = embedding_model() - fields = {f.name for f in model._meta.get_fields()} - values = {"detection": detection, "algorithm": self.algorithm, "vector": vector.tolist()} - if "key" in fields: - values["key"] = "embedding" - if "project" in fields: - values["project"] = detection.source_image.project - model.objects.create(**values) - - def _target_classifications(self) -> None: - for detection in Detection.objects.filter(source_image__project=self.target_project): - detection.classifications.create( - taxon=self.taxon, score=0.4, algorithm=self.algorithm, timestamp=detection.timestamp - ) - - def test_classifier_features_are_exported_and_put_back_on_the_classification(self): - self._populate_target() - with tempfile.TemporaryDirectory() as directory: - directory = pathlib.Path(directory) - manifest = export_embeddings(self.source_project, self.algorithm, directory) - - self.assertEqual(manifest.sources[SOURCE_CLASSIFIER_FEATURES]["count"], 6) - self.assertEqual(manifest.sources[SOURCE_CLASSIFIER_FEATURES]["dimensions"], 2048) - self.assertEqual(EmbeddingManifest.read(directory).count, manifest.count) - matrix = np.load(directory / vectors_file(SOURCE_CLASSIFIER_FEATURES)) - for entry in read_index(directory): - if entry.source == SOURCE_CLASSIFIER_FEATURES: - np.testing.assert_array_equal(matrix[entry.row], self.features[entry.key.capture_path]) - - # No classification on the target yet: nowhere to put the features, so they are skipped, not lost. - report = import_embeddings(self.target_project, directory, execute=True) - self.assertEqual(report.summary()["detections"][MATCH_EXACT], manifest.count) - self.assertEqual(report.skipped_no_classification, 6) - self.assertEqual(report.written[SOURCE_CLASSIFIER_FEATURES], 0) - - self._target_classifications() - report = import_embeddings(self.target_project, directory, execute=True) - self.assertEqual(report.written[SOURCE_CLASSIFIER_FEATURES], 6) - stored = Classification.objects.get( - detection__source_image=self.target_captures[0], detection__source_image__project=self.target_project - ) - np.testing.assert_array_equal( - np.asarray(list(stored.features_2048), dtype=np.float32), self.features[self.target_captures[0].path] - ) - - again = import_embeddings(self.target_project, directory, execute=True) - self.assertEqual( - (again.written[SOURCE_CLASSIFIER_FEATURES], again.skipped_existing[SOURCE_CLASSIFIER_FEATURES]), (0, 6) - ) - - def test_embeddings_are_exported_beside_the_features_and_rewritten_as_rows(self): - if embedding_model() is None: - self.skipTest("This branch has no DetectionEmbedding table") - self._populate_target() - with tempfile.TemporaryDirectory() as directory: - directory = pathlib.Path(directory) - manifest = export_embeddings(self.source_project, self.algorithm, directory) - - self.assertEqual(manifest.sources[SOURCE_EMBEDDING]["count"], 3) - self.assertEqual(manifest.sources[SOURCE_EMBEDDING]["dimensions"], 16) - self.assertEqual( - manifest.sources[SOURCE_CLASSIFIER_FEATURES]["count"], 6, "both stores, no double counting" - ) - self.assertEqual(manifest.count, 9) - - report = import_embeddings(self.target_project, directory, execute=True) - self.assertEqual(report.written[SOURCE_EMBEDDING], 3) - rows = embedding_model().objects.filter( - detection__source_image__project=self.target_project, algorithm=self.algorithm - ) - self.assertEqual(rows.count(), 3) - row = rows.get(detection__source_image=self.target_captures[0]) - # The settled schema stores half precision; compare at that precision. - np.testing.assert_allclose( - np.asarray(list(row.vector), dtype=np.float32), - self.embeddings[self.target_captures[0].path], - rtol=2e-3, - ) - - again = import_embeddings(self.target_project, directory, execute=True) - self.assertEqual((again.written[SOURCE_EMBEDDING], again.skipped_existing[SOURCE_EMBEDDING]), (0, 3)) - - class TestConfirmationNeedsItsReview(ReplayTestCase): """A cached confirmation without its review row is re-confirmed so the review gets written.""" @@ -438,9 +310,31 @@ def test_a_confirmation_with_its_review_is_left_alone(self): spied.assert_not_called() self.assertEqual(report.outcome_counts(), {OUTCOME_UNCHANGED: 2}) - def test_on_a_schema_with_reviews_the_review_row_is_written(self): - if review_model() is None: - self.skipTest("This branch has no ValidationReview table") + def test_a_review_table_without_the_grouping_aspect_falls_back_to_the_cache(self): + class aspect_field: + choices = [("identification", "Identification"), ("comment", "Comment")] + + class meta: + @staticmethod + def get_field(name): + return aspect_field + + class review_table: + _meta = meta + + with mock.patch("ami.main.models_future.validated_occurrences.review_model", return_value=review_table): + self.assertFalse(grouping_reviews_supported()) + self._populate_target() + self._replay(execute=True) + with mock.patch("ami.main.models_future.validated_occurrences.verify_grouping") as spied: + _, report = self._replay(execute=True) + + spied.assert_not_called() + self.assertEqual(report.outcome_counts(), {OUTCOME_UNCHANGED: 2}) + + def test_on_a_schema_with_grouping_reviews_the_review_row_is_written(self): + if not grouping_reviews_supported(): + self.skipTest("This branch does not keep grouping confirmations as reviews") self._populate_target() self._replay(execute=True) rebuilt = self._target_track()