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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
64 changes: 64 additions & 0 deletions ami/main/management/commands/reset_tracking.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,64 @@
"""
Undo tracking for one or more sessions, so they can be tracked again from scratch.

Splits every multi-detection occurrence back into one occurrence per detection, clears
the chain links and grouping confirmations, and prints the counts before and after. See
``ami/main/models_future/session_reset.py`` for exactly what changes. Each session is
reset in its own transaction.
"""

import dataclasses
import logging

from django.core.management.base import BaseCommand, CommandError

from ami.main.models import Event
from ami.main.models_future.session_reset import SessionResetRefused, reset_session_tracking

logger = logging.getLogger(__name__)


class Command(BaseCommand):
help = "Split a session's tracked occurrences back into one occurrence per detection."

def add_arguments(self, parser):
parser.add_argument("--project", type=int, required=True, help="Project the sessions belong to.")
parser.add_argument(
"--event",
type=int,
action="append",
dest="events",
required=True,
help="Session (event) ID to reset. Repeat to reset several.",
)
parser.add_argument("--dry-run", action="store_true", help="Report what would change and write nothing.")
parser.add_argument(
"--force",
action="store_true",
help="Reset even when occurrences in the session carry human identifications.",
)

def handle(self, *args, **options):
project_id: int = options["project"]
event_ids: list[int] = options["events"]
events = {event.pk: event for event in Event.objects.filter(project_id=project_id, pk__in=event_ids)}
missing = [pk for pk in event_ids if pk not in events]
if missing:
raise CommandError(f"Session(s) {missing} not found in project {project_id}")

refused = []
for pk in event_ids:
try:
result = reset_session_tracking(events[pk], force=options["force"], dry_run=options["dry_run"])
except SessionResetRefused as error:
self.stderr.write(str(error))
refused.append(pk)
continue
summary = {k: v for k, v in dataclasses.asdict(result).items() if k not in ("before", "after")}
logger.info(f"reset_tracking: {summary}")
self.stdout.write(f"Session {pk}{' (dry run)' if result.dry_run else ''}: {summary}")
self.stdout.write(f" before: {dataclasses.asdict(result.before)}")
if result.after is not None:
self.stdout.write(f" after: {dataclasses.asdict(result.after)}")
if refused:
raise CommandError(f"Refused to reset session(s) {refused}; see above.")
330 changes: 330 additions & 0 deletions ami/main/models_future/session_reset.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,330 @@
"""Undo tracking for a session: put every detection back in an occurrence of its own.

Tracking only runs on a session whose occurrences each hold a single detection (see
``tracking_task.event_is_fresh``). Resetting returns a tracked session to that state,
so the same night can be tracked again with other settings, or looked at as it was
before tracking. It is a staff tool: it throws away every grouping in the session,
including groupings a person confirmed.

After a reset, for every detection of the session that shared an occurrence:

- The earliest detection (in capture order) keeps the occurrence; each other
detection in the session moves to a new occurrence of its own. An occurrence that
also holds detections of another session stays with those detections, untouched
there, and every detection of this session leaves it.
- Every occurrence the split touched takes its determination from its own
detections' best prediction, the rule ``Occurrence.best_prediction`` applies, or
none if they have no scored prediction. An
occurrence with a human identification keeps the determination it has.
- Chain links between two detections of the session are cleared. A link that crosses
into another session records one animal across a regroup boundary, which tracking
never follows, so it is kept.
- Grouping verification is cleared on every occurrence of the session.
- Classifications the tracking task recorded on the session's detections are deleted.
Each is a copy of a prediction on the same detection, made when a merge changed a
determination, so it describes a merge that no longer exists.
- For the same reason, the tracking results in the history of the session's
occurrences are deleted. Reviews stay, since the history keeps every review. An
occurrence that also holds detections of another session keeps its whole history,
because a tracking result does not say which session's run wrote it.
"""

from __future__ import annotations

import dataclasses
from collections.abc import Iterable

from django.db import transaction
from django.db.models import Count, Prefetch, Q

from ami.main.models import (
Classification,
Detection,
Event,
Identification,
Occurrence,
OccurrenceHistoryRecord,
update_calculated_fields_for_sessions_and_stations,
)
from ami.main.models_future.occurrence import best_prediction_from_prefetch
from ami.main.models_future.track_stats import refresh_track_stats_for_ids
from ami.main.models_future.tracks import _capture_order_key

# Occurrences whose determination is recomputed per round trip. Each batch reads the
# occurrences, their detections and those detections' classifications: three queries.
DETERMINATION_BATCH_SIZE = 2000

# The Algorithm key the tracking task records its determination changes under.
TRACKING_ALGORITHM_KEY = "tracking"


class SessionResetRefused(ValueError):
"""The session holds work a reset would discard, and the caller did not force it."""


@dataclasses.dataclass
class SessionTrackingCounts:
"""How grouped a session is, read through its detections' captures."""

occurrences: int
multi_detection_occurrences: int
determinations: int
links: int
grouping_verified: int


@dataclasses.dataclass
class SessionResetResult:
event_id: int
dry_run: bool
identifications: int
occurrences_split: int
occurrences_created: int
links_cleared: int
verifications_cleared: int
tracking_classifications_deleted: int
tracking_history_deleted: int
determinations_updated: int
before: SessionTrackingCounts
after: SessionTrackingCounts | None


def _session_detections(event: Event):
return Detection.objects.filter(source_image__event=event)


def _session_occurrence_ids(event: Event):
"""Occurrences reached through the session's captures, as ``event_is_fresh`` finds them."""
return _session_detections(event).filter(occurrence__isnull=False).order_by().values("occurrence_id")


def session_tracking_counts(event: Event) -> SessionTrackingCounts:
"""Occurrences, multi-frame occurrences, distinct determinations, links and confirmations. Three queries."""
occurrences = Occurrence.objects.filter(pk__in=_session_occurrence_ids(event))
totals = occurrences.aggregate(
occurrences=Count("pk"),
determinations=Count("determination", distinct=True),
grouping_verified=Count("pk", filter=Q(grouping_verified_at__isnull=False)),
)
multi = occurrences.annotate(_n=Count("detections")).filter(_n__gt=1).count()
links = _session_detections(event).filter(next_detection__source_image__event=event).count()
return SessionTrackingCounts(
occurrences=totals["occurrences"],
multi_detection_occurrences=multi,
determinations=totals["determinations"],
links=links,
grouping_verified=totals["grouping_verified"],
)


def _plan_split(event: Event) -> tuple[list[int], list[int], dict[int, int]]:
"""The multi-detection occurrences of the session, the detections that leave them, and
the occurrences whose session changes.

An occurrence that lies wholly in this session keeps its first detection in capture
order, and its other detections leave. An occurrence that also holds detections of
another session stays whole there: every detection of this session leaves it, and if
it was filed under this session it moves to the session of its first remaining
detection.
"""
multi = dict(
Occurrence.objects.filter(pk__in=_session_occurrence_ids(event))
.annotate(_n=Count("detections"))
.filter(_n__gt=1)
.values_list("pk", "event_id")
)
if not multi:
return [], [], {}
rows = list(
Detection.objects.filter(occurrence_id__in=multi)
.order_by()
.values("pk", "occurrence_id", "source_image_id", "source_image__timestamp", "source_image__event_id")
)
by_occurrence: dict[int, list[dict]] = {}
for row in rows:
by_occurrence.setdefault(row["occurrence_id"], []).append(row)
movers: list[int] = []
new_sessions: dict[int, int] = {}
for occurrence_id, members in by_occurrence.items():
members.sort(key=_capture_order_key)
outside = [row for row in members if row["source_image__event_id"] != event.pk]
if outside:
movers.extend(row["pk"] for row in members if row["source_image__event_id"] == event.pk)
if multi[occurrence_id] == event.pk:
new_sessions[occurrence_id] = outside[0]["source_image__event_id"]
else:
movers.extend(row["pk"] for row in members[1:])
return list(multi), movers, new_sessions


def _refresh_determinations(occurrence_ids: Iterable[int]) -> int:
"""Set each occurrence's determination to its best prediction, in batches; returns how many changed.

The batched form of ``update_occurrence_determination`` for occurrences with no
identifications, using the same choice as ``Occurrence.best_prediction``. Unlike it,
an occurrence with no scored prediction gets no determination.
"""
ids = list(occurrence_ids)
classifications = Classification.objects.select_related(None).only(
"pk", "detection_id", "algorithm_id", "taxon_id", "score", "terminal"
)
detections = (
Detection.objects.select_related(None)
.only("pk", "occurrence_id")
.prefetch_related(Prefetch("classifications", queryset=classifications))
)
changed: list[Occurrence] = []
for start in range(0, len(ids), DETERMINATION_BATCH_SIZE):
batch = (
Occurrence.objects.select_related(None)
.filter(pk__in=ids[start : start + DETERMINATION_BATCH_SIZE])
.only("pk", "determination_id", "determination_score")
)
for occurrence in batch.prefetch_related(Prefetch("detections", queryset=detections)):
best = best_prediction_from_prefetch(occurrence)
# With no scored prediction left, a determination inherited from the merged
# track would describe detections that have moved away, so it is cleared.
wanted = (best.taxon_id, best.score) if best is not None and best.taxon_id else (None, None)
if (occurrence.determination_id, occurrence.determination_score) != wanted:
occurrence.determination_id, occurrence.determination_score = wanted
changed.append(occurrence)
Occurrence.objects.bulk_update(changed, ["determination", "determination_score"], batch_size=1000)
return len(changed)


@dataclasses.dataclass
class _ResetPlan:
before: SessionTrackingCounts
identified: set[int]
identification_count: int
multi_ids: list[int]
movers: list[int]
new_sessions: dict[int, int]


def _plan_reset(event: Event, force: bool) -> _ResetPlan:
"""Read the session's state and decide what the reset changes; raises ``SessionResetRefused``."""
before = session_tracking_counts(event)
identified = set(
Identification.objects.filter(occurrence_id__in=_session_occurrence_ids(event)).values_list(
"occurrence_id", flat=True
)
)
identification_count = Identification.objects.filter(occurrence_id__in=identified).count() if identified else 0
if identification_count and not force:
raise SessionResetRefused(
f"Session {event.pk} has {identification_count} identification(s) on {len(identified)} occurrence(s). "
"Resetting would split occurrences a person identified; pass force to reset anyway."
)
multi_ids, movers, new_sessions = _plan_split(event)
return _ResetPlan(before, identified, identification_count, multi_ids, movers, new_sessions)


def _lock_session_occurrences(event: Event) -> None:
"""Lock the session's occurrence rows until the transaction ends.

A tracking run, a track edit or a new identification on these occurrences waits for
the reset instead of changing them between the plan and the writes.
"""
list(
Occurrence.objects.select_related(None)
.select_for_update(of=("self",))
.filter(pk__in=_session_occurrence_ids(event))
.order_by("pk")
.values_list("pk", flat=True)
)


def reset_session_tracking(event: Event, *, force: bool = False, dry_run: bool = False) -> SessionResetResult:
"""Return ``event`` to one occurrence per detection, with no links and no confirmations.

Raises ``SessionResetRefused`` when the session's occurrences carry human
identifications, unless ``force`` is set; forced, each identification stays on the
occurrence it was made on, which keeps its first detection in the session (or its
detections in another session). The plan and the writes run in one transaction that
holds the session's occurrences locked. With ``dry_run`` the counts are those a reset
would produce and nothing is written or locked. The query count depends on the number
of occurrences only through the determination batches, not per row.
"""
links = _session_detections(event).filter(next_detection__source_image__event=event)
tracking_classifications = Classification.objects.filter(
detection__source_image__event=event, algorithm__key=TRACKING_ALGORITHM_KEY
)
verified = Occurrence.objects.filter(pk__in=_session_occurrence_ids(event), grouping_verified_at__isnull=False)
spanning = (
Detection.objects.filter(occurrence_id__in=_session_occurrence_ids(event))
.exclude(source_image__event=event)
.order_by()
.values("occurrence_id")
)
tracking_history = OccurrenceHistoryRecord.objects.filter(
occurrence_id__in=_session_occurrence_ids(event),
kind=OccurrenceHistoryRecord.Kind.ALGORITHM_RESULT,
subtype=TRACKING_ALGORITHM_KEY,
).exclude(occurrence_id__in=spanning)

if dry_run:
plan = _plan_reset(event, force)
return SessionResetResult(
event_id=event.pk,
dry_run=True,
identifications=plan.identification_count,
occurrences_split=len(plan.multi_ids),
occurrences_created=len(plan.movers),
links_cleared=plan.before.links,
verifications_cleared=plan.before.grouping_verified,
tracking_classifications_deleted=tracking_classifications.count(),
tracking_history_deleted=tracking_history.count(),
determinations_updated=0,
before=plan.before,
after=None,
)

with transaction.atomic():
_lock_session_occurrences(event)
plan = _plan_reset(event, force)
movers, new_sessions = plan.movers, plan.new_sessions
links_cleared = links.update(next_detection=None)
verifications_cleared = verified.update(grouping_verified_at=None, grouping_verified_by=None)
_, deleted_by_model = tracking_classifications.delete()
tracking_deleted = deleted_by_model.get(Classification._meta.label, 0)
# Before the split, while each occurrence still holds the detections that decide
# whether it reaches into another session.
history_deleted, _ = tracking_history.delete()

new_occurrences = Occurrence.objects.bulk_create(
[Occurrence(event=event, deployment_id=event.deployment_id, project_id=event.project_id) for _ in movers],
batch_size=1000,
)
Detection.objects.bulk_update(
[
Detection(pk=detection_id, occurrence_id=occurrence.pk)
for detection_id, occurrence in zip(movers, new_occurrences)
],
["occurrence"],
batch_size=1000,
)

Occurrence.objects.bulk_update(
[Occurrence(pk=pk, event_id=session_id) for pk, session_id in new_sessions.items()], ["event"]
)

touched = [*plan.multi_ids, *(o.pk for o in new_occurrences)]
determinations_updated = _refresh_determinations(pk for pk in touched if pk not in plan.identified)
refresh_track_stats_for_ids(touched)

update_calculated_fields_for_sessions_and_stations([event.pk, *set(new_sessions.values())], stations_async=False)
return SessionResetResult(
event_id=event.pk,
dry_run=False,
identifications=plan.identification_count,
occurrences_split=len(plan.multi_ids),
occurrences_created=len(new_occurrences),
links_cleared=links_cleared,
verifications_cleared=verifications_cleared,
tracking_classifications_deleted=tracking_deleted,
tracking_history_deleted=history_deleted,
determinations_updated=determinations_updated,
before=plan.before,
after=session_tracking_counts(event),
)
Loading