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
12 changes: 7 additions & 5 deletions ami/main/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -3828,8 +3828,6 @@ def update_occurrence_determination(
The `occurrence` object may already have a different un-saved determination set
so it is necessary to retrieve the current determination from the database, but
this can also be passed in as an argument to avoid an extra database query.

@TODO Add tests for this important method!
"""
needs_update = False

Expand All @@ -3850,13 +3848,17 @@ def update_occurrence_determination(
new_determination = None
new_score = None

# Pick the winner first and compare each cached field to it separately. The score
# can change while the taxon stays the same (class masking re-scores a classification
# without changing its winner), and a stale score hides the occurrence below the
# project's score threshold. See #1461.
top_identification = occurrence.best_identification
if top_identification and top_identification.taxon and top_identification.taxon != current_determination:
if top_identification and top_identification.taxon:
new_determination = top_identification.taxon
new_score = top_identification.score
elif not top_identification:
top_prediction = occurrence.best_prediction
if top_prediction and top_prediction.taxon and top_prediction.taxon != current_determination:
if top_prediction and top_prediction.taxon:
new_determination = top_prediction.taxon
new_score = top_prediction.score

Expand All @@ -3865,7 +3867,7 @@ def update_occurrence_determination(
occurrence.determination = new_determination
needs_update = True

if new_score and new_score != occurrence.determination_score:
if new_score is not None and new_score != occurrence.determination_score:
logger.debug(f"Changing det. score of {occurrence} from {occurrence.determination_score} to {new_score}")
occurrence.determination_score = new_score
needs_update = True
Expand Down
96 changes: 96 additions & 0 deletions ami/main/tests.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,7 +41,9 @@
Taxon,
TaxonRank,
group_images_into_events,
update_occurrence_determination,
)
from ami.ml.models import Algorithm
from ami.ml.models.pipeline import Pipeline
from ami.ml.models.processing_service import ProcessingService
from ami.ml.models.project_pipeline_config import ProjectPipelineConfig
Expand Down Expand Up @@ -8414,3 +8416,97 @@ def test_identifications_browsable_page(self):
html = self._get_html("/api/v2/identifications/")
self._assert_number_input(html, "occurrence")
self._assert_number_input(html, "taxon")


class TestOccurrenceDeterminationRefresh(TestCase):
"""Pin how ``update_occurrence_determination`` keeps the cached determination
fields in step with the best identification or prediction.

The score must follow the winner even when the winning taxon is unchanged:
class masking re-scores a classification under a new algorithm without
changing its taxon, and a stale lower score hides the occurrence below the
project's score threshold. See #1461.
"""

def setUp(self):
self.project, self.deployment = setup_test_project()
create_taxa(project=self.project)
create_captures(deployment=self.deployment)
taxa = list(Taxon.objects.filter(projects=self.project, rank=TaxonRank.SPECIES.name).order_by("name")[:2])
self.taxon_a, self.taxon_b = taxa
create_occurrences(deployment=self.deployment, num=1, taxon=self.taxon_a, determination_score=0.38)
self.detection = Detection.objects.filter(source_image__deployment=self.deployment).latest("pk")
self.occurrence = self.detection.occurrence
assert self.occurrence is not None
self.assertEqual(self.occurrence.determination, self.taxon_a)
self.assertEqual(self.occurrence.determination_score, 0.38)
self.masking_algorithm = Algorithm.objects.create(name="Masked classifier", key="masked_classifier_test")

def _rescore(self, taxon: Taxon, score: float) -> Classification:
"""Mimic class masking: demote the current terminal classification and add a new
terminal one from another algorithm."""
self.detection.classifications.update(terminal=False)
return self.detection.classifications.create(
taxon=taxon,
score=score,
terminal=True,
algorithm=self.masking_algorithm,
timestamp=datetime.datetime.now(),
)

def _refresh(self) -> bool:
needs_update = update_occurrence_determination(
self.occurrence, current_determination=self.occurrence.determination
)
self.occurrence.refresh_from_db()
return needs_update

def test_score_follows_a_rescored_prediction_with_the_same_taxon(self):
"""A new best prediction with the same taxon but a different score updates the score."""
self._rescore(self.taxon_a, 0.55)

self.assertTrue(self._refresh())
self.assertEqual(self.occurrence.determination, self.taxon_a)
self.assertEqual(self.occurrence.determination_score, 0.55)

def test_taxon_and_score_follow_a_new_best_prediction(self):
"""A new best prediction with a different taxon updates both fields."""
self._rescore(self.taxon_b, 0.55)

self.assertTrue(self._refresh())
self.assertEqual(self.occurrence.determination, self.taxon_b)
self.assertEqual(self.occurrence.determination_score, 0.55)

def test_no_save_when_nothing_changed(self):
"""When the winner and its score already match, the occurrence is not written."""
updated_at = self.occurrence.updated_at

self.assertFalse(self._refresh())
self.assertEqual(self.occurrence.determination, self.taxon_a)
self.assertEqual(self.occurrence.determination_score, 0.38)
self.assertEqual(self.occurrence.updated_at, updated_at)

def test_human_identification_wins_over_a_rescored_prediction(self):
"""A human identification sets the determination and its own score, and a later
machine re-score does not overwrite either."""
user = User.objects.create_user(email="identifier@insectai.org") # type: ignore
identification = Identification.objects.create(occurrence=self.occurrence, user=user, taxon=self.taxon_b)
self.occurrence.refresh_from_db()
self.assertEqual(self.occurrence.determination, self.taxon_b)
self.assertEqual(self.occurrence.determination_score, identification.score)

self._rescore(self.taxon_a, 0.99)

self.assertFalse(self._refresh())
self.assertEqual(self.occurrence.determination, self.taxon_b)
self.assertEqual(self.occurrence.determination_score, identification.score)

def test_human_identification_with_the_same_taxon_sets_its_own_score(self):
"""Confirming the machine's taxon still replaces the machine score with the
identification's score, the same as when the human picks a different taxon."""
user = User.objects.create_user(email="identifier@insectai.org") # type: ignore
identification = Identification.objects.create(occurrence=self.occurrence, user=user, taxon=self.taxon_a)

self.occurrence.refresh_from_db()
self.assertEqual(self.occurrence.determination, self.taxon_a)
self.assertEqual(self.occurrence.determination_score, identification.score)
42 changes: 42 additions & 0 deletions ami/ml/post_processing/tests/test_class_masking.py
Original file line number Diff line number Diff line change
Expand Up @@ -458,6 +458,48 @@ def test_occurrences_updated_counts_only_changed_determinations(self):
"Only the occurrence whose determination changed (occ1) counts",
)

def test_rescored_unchanged_winner_keeps_the_occurrence_visible(self):
"""When masking re-scores a classification without changing its winner, the
occurrence takes the new score, so a winner that rises above the project's
score threshold is no longer hidden by the stale lower score. See #1461."""
self.project.default_filters_score_threshold = 0.5
self.project.save()
taxa_list = TaxaList.objects.create(name="Visibility test")
taxa_list.taxa.set(self.species_taxa[:2]) # excludes index 2

new_algorithm = Algorithm.objects.create(
name="masked_visibility",
key="masked_visibility_test",
task_type=AlgorithmTaskType.CLASSIFICATION.value,
category_map=self.algorithm.category_map,
)

# Index 0 wins before and after masking. Its original probability is below the
# threshold (about 0.44); dropping the excluded index 2 and renormalising lifts
# it above (about 0.73).
logits = [1.0, 0.0, 0.9]
det, occ = self._detection_with_occurrence()
clf = self._create_classification_with_logits(det, self.species_taxa[0], _softmax(logits), logits)
occ.save(update_determination=True)
self.assertEqual(occ.determination, self.species_taxa[0])
self.assertLess(occ.determination_score, 0.5)
visible = Occurrence.objects.filter(pk=occ.pk).apply_default_filters(project=self.project, request=None)
self.assertFalse(visible.exists(), "Sanity: hidden by the threshold before masking")

make_classifications_filtered_by_taxa_list(
classifications=Classification.objects.filter(pk=clf.pk),
taxa_list=taxa_list,
algorithm=self.algorithm,
new_algorithm=new_algorithm,
)

masked = Classification.objects.get(detection=det, algorithm=new_algorithm, terminal=True)
self.assertGreater(masked.score, 0.5)
occ.refresh_from_db()
self.assertEqual(occ.determination, self.species_taxa[0], "Winner is unchanged")
self.assertAlmostEqual(occ.determination_score, masked.score)
self.assertTrue(visible.exists(), "The re-scored occurrence passes the default score filter")

# ----- reweight toggle ------------------------------------------------

def test_reweight_false_winner_identical_scores_differ(self):
Expand Down
Loading