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..209226d21 --- /dev/null +++ b/ami/main/management/commands/import_embeddings.py @@ -0,0 +1,80 @@ +""" +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 {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']}" + ) + 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/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..5532927b8 --- /dev/null +++ b/ami/main/models_future/embedding_transfer.py @@ -0,0 +1,445 @@ +"""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. + +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. 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 + +import csv +import dataclasses +import datetime +import json +import logging +import pathlib +from collections import Counter +from collections.abc import Iterable, Iterator +from typing import Any + +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 +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" +INDEX_FILE = "index.csv" +MANIFEST_FILE = "manifest.json" +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) +CLASSIFIER_FEATURES_FIELD = "features_2048" + + +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 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" + + +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 + # Per source: {"file": ..., "count": ..., "dimensions": ...} + 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 + + 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 _embedding_rows(project: Project, algorithm: Algorithm, vector_key: str) -> Iterator[tuple[DetectionKey, Any]]: + model = embedding_model() + 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()) + + +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 + ) + .select_related("detection__source_image__deployment", "detection__detection_algorithm") + .order_by("detection_id", "-timestamp", "-pk") + .distinct("detection_id") + ) + return ((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``, 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) + 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) + writer.writerow(INDEX_COLUMNS) + 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, + 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)) + return manifest + + +@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): + 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 entries + + +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: + entries: list[IndexEntry] + matches: list[DetectionMatch] + 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 + # 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 { + "mode": "execute" if self.execute else "dry-run", + "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, + "skipped_no_field": dict(self.skipped_no_field), + } + + +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 + + +def _write_embeddings( + project: Project, + algorithm: Algorithm, + vector_key: str, + matrix: np.ndarray, + found: list[tuple[int, int]], + 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: + 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[SOURCE_EMBEDDING] += len(existing_ids) + 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) + + +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. 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) + entries = read_index(directory) + if len(entries) != manifest.count: + raise ValueError(f"{INDEX_FILE} has {len(entries)} rows but the manifest says {manifest.count}.") + # 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 + + 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]: + 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]: + 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/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 new file mode 100644 index 000000000..643cec4bd --- /dev/null +++ b/ami/main/models_future/validated_occurrences.py @@ -0,0 +1,731 @@ +"""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.apps import apps +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 +GROUPING_ASPECT = "grouping" + + +# -- 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, 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, 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 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 ( + 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: + """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: + """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_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 new file mode 100644 index 000000000..77223c980 --- /dev/null +++ b/ami/main/test_validated_occurrences.py @@ -0,0 +1,352 @@ +"""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 +from unittest import mock + +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.tracks import verify_grouping +from ami.main.models_future.validated_occurrences import ( + OUTCOME_APPLIED, + OUTCOME_PARTIAL, + OUTCOME_UNCHANGED, + Bundle, + ImportOptions, + build_bundle, + grouping_reviews_supported, + import_bundle, + occurrence_detection_keys, + review_model, +) +from ami.tests.fixtures.main import create_taxa, setup_test_project + +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, 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) + 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_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() + + 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 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_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() + 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})