Skip to content
Open
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
11 changes: 11 additions & 0 deletions ami/ml/post_processing/admin/class_masking_form.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,16 @@ class ClassMaskingActionForm(BasePostProcessingActionForm):
"Off = keep the model's raw absolute scores; the chosen species is unchanged either way."
),
)
only_when_taxon_changes = forms.BooleanField(
required=False,
initial=True,
label="Only record a new identification when the species changes",
help_text=(
"Masking nudges the scores of nearly every prediction, so leaving this off records a new "
"identification for most of them, usually repeating the species already shown. "
"Off = record every re-score, which is useful for comparing confidence across taxa lists."
),
)

def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
Expand Down Expand Up @@ -73,4 +83,5 @@ def to_config(self) -> dict:
"algorithm_id": self.cleaned_data["algorithm_id"].pk,
"taxa_list_id": self.cleaned_data["taxa_list_id"].pk,
"reweight": self.cleaned_data["reweight"],
"only_when_taxon_changes": self.cleaned_data["only_when_taxon_changes"],
}
32 changes: 25 additions & 7 deletions ami/ml/post_processing/class_masking.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,9 @@ class ClassMaskingConfig(pydantic.BaseModel):
# masking. When False, the kept classes retain their original absolute scores and
# the excluded classes are zeroed; the chosen species is identical either way.
reweight: bool = True
# When True (default), only record a re-score that changes the winning taxon;
# recording the rest repeats the taxon already shown. See #1377.
only_when_taxon_changes: bool = True

@pydantic.root_validator(skip_on_failure=True)
def _exactly_one_scope(cls, values: dict) -> dict:
Expand All @@ -49,6 +52,7 @@ def make_classifications_filtered_by_taxa_list(
*,
batch_size: int = 200,
reweight: bool = True,
only_when_taxon_changes: bool = True,
task_logger: logging.Logger = logger,
on_setup: Callable[[int], None] | None = None,
on_batch: Callable[[dict], None] | None = None,
Expand All @@ -62,6 +66,10 @@ def make_classifications_filtered_by_taxa_list(
(attributed to ``new_algorithm``, linked back via ``applied_to``) records the
masked prediction. The original classification is demoted to non-terminal.

With ``only_when_taxon_changes`` (the default), a re-score that leaves the same
winning taxon is not recorded and the source prediction keeps its place. Set it
to False to keep every re-score. See #1377.

Commits in batches of ``batch_size`` so memory stays bounded and the job
health-check reaper sees regular heartbeats. ``on_batch`` is called after every
flush with running counters:
Expand Down Expand Up @@ -156,14 +164,19 @@ def make_classifications_filtered_by_taxa_list(
new_scores = new_scores_np.tolist()
score = float(new_scores_np[top_index])

# No-change short-circuit: if masking shifted no probability (the classes
# this taxa list drops carried ~zero probability here), leave the row
# untouched. Compared against the unmasked softmax, not the stored scores.
if np.allclose(full_softmax, new_scores_np, atol=1e-9):
task_logger.debug(f"Classification {classification.pk} unchanged by masking; skipping")
else:
top_taxon = index_to_taxon.get(top_index) # guaranteed in taxa_in_list (top_index is kept)
top_taxon = index_to_taxon.get(top_index) # guaranteed in taxa_in_list (top_index is kept)

# Leave the row untouched when masking shifted no probability, which
# happens when every class this taxa list drops already scored ~zero.
# Compared against the unmasked softmax, not the stored scores.
unchanged_scores = bool(np.allclose(full_softmax, new_scores_np, atol=1e-9))
# Compare ids rather than instances: the scope query does not fetch the
# taxon, so touching ``classification.taxon`` would load one per row.
unchanged_taxon = classification.taxon_id == (top_taxon.pk if top_taxon else None)

if unchanged_scores or (only_when_taxon_changes and unchanged_taxon):
task_logger.debug(f"Classification {classification.pk} keeps its taxon under this mask; skipping")
else:
classification.terminal = False
classification.updated_at = timestamp

Expand Down Expand Up @@ -252,6 +265,10 @@ def _get_or_create_masking_algorithm(
them apart. Its category map is the source map (indices still align with
the masked score vector) and is persisted — earlier code set it in memory
only, so masked classifications referenced a null map.

``only_when_taxon_changes`` is deliberately not part of the identity: it
decides how many predictions get recorded, not what a recorded prediction
means, so both settings produce interchangeable rows.
"""
mode = "reweighted" if reweight else "absolute"
algorithm, created = Algorithm.objects.get_or_create(
Expand Down Expand Up @@ -351,6 +368,7 @@ def _on_batch(m: dict) -> None:
algorithm=source_algorithm,
new_algorithm=masking_algorithm,
reweight=config.reweight,
only_when_taxon_changes=config.only_when_taxon_changes,
task_logger=self.logger,
on_setup=_on_setup,
on_batch=_on_batch,
Expand Down
106 changes: 105 additions & 1 deletion ami/ml/post_processing/tests/test_class_masking.py
Original file line number Diff line number Diff line change
Expand Up @@ -367,7 +367,10 @@ def test_occurrences_updated_counts_only_changed_determinations(self):

occ1: original winner is index 2 (excluded) — masking flips determination to index 0.
occ2: original winner is index 0 (kept) — masking reassigns scores but determination stays index 0.
Only occ1 should count."""
Only occ1 should count.

Runs with ``only_when_taxon_changes=False`` so occ2 reaches the counting step
rather than being skipped."""
taxa_list = TaxaList.objects.create(name="Changed-only count test")
taxa_list.taxa.set(self.species_taxa[:2]) # excludes species_taxa[2] (index 2)

Expand Down Expand Up @@ -402,6 +405,7 @@ def test_occurrences_updated_counts_only_changed_determinations(self):
taxa_list=taxa_list,
algorithm=self.algorithm,
new_algorithm=new_algorithm,
only_when_taxon_changes=False,
)

self.assertEqual(metrics["classifications_masked"], 2, "Both classifications are modified by masking")
Expand Down Expand Up @@ -574,3 +578,103 @@ def test_scope_size_is_reported_before_the_first_batch(self):

self.assertEqual(events[0], ("setup", 1), "Scope size is reported once, before the first batch")
self.assertEqual([name for name, _ in events], ["setup", "batch"])

# ----- skipping re-scores that keep the same taxon ---------------------

def test_no_new_classification_when_the_taxon_does_not_change(self):
"""By default, a re-score that keeps the same winning taxon writes nothing,
so the identification history does not gain a repeat of the taxon already
there. See #1377."""
# Index 0 has the highest logit and is in the taxa list, so it wins both
# before and after masking. Index 2 is dropped and held real probability,
# so the scores do change.
logits = [5.0, 1.0, 3.0]
taxa_list = TaxaList.objects.create(name="Winner survives the mask")
taxa_list.taxa.set(self.species_taxa[:2])

det, occ = self._detection_with_occurrence()
original = self._create_classification_with_logits(det, self.species_taxa[0], _softmax(logits), logits)
occ.save(update_determination=True) # settle the pre-masking determination

ClassMaskingTask(
source_image_collection_id=self.collection.pk,
taxa_list_id=taxa_list.pk,
algorithm_id=self.algorithm.pk,
).run()

original.refresh_from_db()
self.assertTrue(original.terminal, "The source prediction keeps its place as the current one")
self.assertEqual(Classification.objects.filter(detection=det).count(), 1, "No second row for the same taxon")
occ.refresh_from_db()
self.assertEqual(occ.determination, self.species_taxa[0], "The occurrence keeps the determination it had")

def test_new_classification_when_the_taxon_changes(self):
"""The taxon check must not swallow the case masking exists for: a guard that
always skipped would pass the test above, so the transition it is meant to
allow is pinned here."""
logits = [2.0, 1.0, 5.0] # index 2 wins before masking and is dropped
taxa_list = TaxaList.objects.create(name="Winner is dropped")
taxa_list.taxa.set(self.species_taxa[:2])

det, _ = self._detection_with_occurrence()
original = self._create_classification_with_logits(det, self.species_taxa[2], _softmax(logits), logits)

ClassMaskingTask(
source_image_collection_id=self.collection.pk,
taxa_list_id=taxa_list.pk,
algorithm_id=self.algorithm.pk,
).run()

original.refresh_from_db()
self.assertFalse(original.terminal)
new_clf = Classification.objects.get(detection=det, terminal=True)
self.assertEqual(new_clf.taxon, self.species_taxa[0])
self.assertEqual(new_clf.applied_to, original)

def test_rescore_is_recorded_when_the_operator_asks_for_it(self):
"""With ``only_when_taxon_changes=False`` the re-scored prediction is stored
even though the taxon is unchanged, which is what an operator comparing
confidence before and after a taxa list needs."""
logits = [5.0, 1.0, 3.0]
scores = _softmax(logits)
taxa_list = TaxaList.objects.create(name="Record every re-score")
taxa_list.taxa.set(self.species_taxa[:2])

det, _ = self._detection_with_occurrence()
original = self._create_classification_with_logits(det, self.species_taxa[0], scores, logits)

ClassMaskingTask(
source_image_collection_id=self.collection.pk,
taxa_list_id=taxa_list.pk,
algorithm_id=self.algorithm.pk,
only_when_taxon_changes=False,
).run()

original.refresh_from_db()
self.assertFalse(original.terminal)
new_clf = Classification.objects.get(detection=det, terminal=True)
self.assertEqual(new_clf.taxon, self.species_taxa[0], "Same taxon as the source")
self.assertGreater(new_clf.score, scores[0], "Dropping a class raises the winner's confidence")

def test_both_skip_settings_share_one_masking_algorithm(self):
"""Both skip settings resolve to the same masking algorithm, unlike the
reweight mode, which changes score semantics and so forms part of the
algorithm's identity. See #1377."""
logits = [2.0, 1.0, 5.0]
taxa_list = TaxaList.objects.create(name="Shared algorithm list")
taxa_list.taxa.set(self.species_taxa[:2])

for skip in (True, False):
det, _ = self._detection_with_occurrence()
self._create_classification_with_logits(det, self.species_taxa[2], _softmax(logits), logits)
ClassMaskingTask(
source_image_collection_id=self.collection.pk,
taxa_list_id=taxa_list.pk,
algorithm_id=self.algorithm.pk,
only_when_taxon_changes=skip,
).run()

self.assertEqual(
Algorithm.objects.filter(key__startswith=f"{self.algorithm.key}_filtered_by_taxa_list_").count(),
1,
)
45 changes: 45 additions & 0 deletions ami/ml/post_processing/tests/test_class_masking_admin.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,16 @@ def test_reweight_can_be_set_false(self):
config = ClassMaskingConfig(source_image_collection_id=1, taxa_list_id=2, algorithm_id=3, reweight=False)
self.assertFalse(config.reweight)

def test_only_when_taxon_changes_defaults_to_true(self):
config = ClassMaskingConfig(source_image_collection_id=1, taxa_list_id=2, algorithm_id=3)
self.assertTrue(config.only_when_taxon_changes)

def test_only_when_taxon_changes_can_be_set_false(self):
config = ClassMaskingConfig(
source_image_collection_id=1, taxa_list_id=2, algorithm_id=3, only_when_taxon_changes=False
)
self.assertFalse(config.only_when_taxon_changes)


class _PostProcessingAdminCase(TestCase):
@classmethod
Expand Down Expand Up @@ -175,6 +185,41 @@ def test_to_config_includes_reweight_false_when_unchecked(self):
self.assertFalse(job.params["config"]["reweight"])


class TestClassMaskingFormOnlyWhenTaxonChanges(_PostProcessingAdminCase):
"""The admin form exposes the "only when the species changes" toggle and passes
it through to the job config."""

def _post_collection(self, include_flag: bool):
data = {
"action": "run_class_masking",
django_admin.helpers.ACTION_CHECKBOX_NAME: [str(self.collection.pk)],
"confirm": "yes",
"taxa_list_id": str(self.taxa_list.pk),
"algorithm_id": str(self.algorithm.pk),
"reweight": "on",
}
if include_flag:
data["only_when_taxon_changes"] = "on"
return self.client.post(reverse("admin:main_sourceimagecollection_changelist"), data=data)

def test_form_has_the_field_checked_by_default(self):
form = ClassMaskingActionForm()
self.assertIn("only_when_taxon_changes", form.fields)
self.assertTrue(form.fields["only_when_taxon_changes"].initial)

def test_checked_box_yields_true_in_the_job_config(self):
response = self._post_collection(include_flag=True)
self.assertEqual(response.status_code, 302)
job = Job.objects.get(project=self.project, job_type_key="post_processing")
self.assertTrue(job.params["config"]["only_when_taxon_changes"])

def test_unchecked_box_yields_false_in_the_job_config(self):
response = self._post_collection(include_flag=False)
self.assertEqual(response.status_code, 302)
job = Job.objects.get(project=self.project, job_type_key="post_processing")
self.assertFalse(job.params["config"]["only_when_taxon_changes"])


class TestClassMaskingFormScopeFiltering(TestCase):
"""The class-mask form offers only classifiers that actually produced
classifications within the selected scope, so an operator cannot pick an
Expand Down
Loading