Skip to content

Commit 558fcaf

Browse files
mihowclaude
andcommitted
feat(post-processing): implement class masking + rank rollup on the admin framework
Rework of #999 onto the #1289 post-processing framework. The branch was cut from an old main with its own hand-rolled admin action; ClassMaskingTask and RankRollupTask now subclass BasePostProcessingTask with pydantic config schemas and are triggered through make_post_processing_action (collection scope on SourceImageCollection, single-occurrence scope on Occurrence for class masking). Correctness fixes from the review threads: - Class masking selects the top class from an -inf-masked softmax, so a class excluded by the taxa list can never win even when it had the highest logit; raises when the taxa list excludes every class in the category map. Stored logits stay raw (JSON-safe) and the mask is captured in scores (excluded -> 0). - The masked-output Algorithm is one per (source algorithm, taxa list) and its category map is persisted (previously set in memory only, so masked classifications referenced a null map). - applied_to is populated on new masked classifications (the provenance the API exposes was left blank). - Rank rollup preloads category-map labels in two queries and select_relates the per-row relations instead of dereferencing category_map per classification. Surfaces provenance in the API: applied_to is added to the Classification serializers, and applied_to__algorithm is prefetched in the occurrence list/detail prefetch and the classification viewset to avoid an N+1 on render. Tests: pydantic config validation, admin trigger for both scopes, the masking maths (including the excluded-class guarantee and the all-excluded error), ClassMaskingTask.run() end to end for both scopes, and rank rollup. 20 new tests; full post_processing suite and occurrence query-count tests pass. Co-Authored-By: Claude <noreply@anthropic.com>
1 parent f5922ec commit 558fcaf

13 files changed

Lines changed: 1119 additions & 3 deletions

File tree

‎ami/main/admin.py‎

Lines changed: 29 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,11 @@
1414
from ami.jobs.models import Job
1515
from ami.ml.models.project_pipeline_config import ProjectPipelineConfig
1616
from ami.ml.post_processing.admin.actions import make_post_processing_action
17+
from ami.ml.post_processing.admin.class_masking_form import ClassMaskingActionForm
18+
from ami.ml.post_processing.admin.rank_rollup_form import RankRollupActionForm
1719
from ami.ml.post_processing.admin.small_size_filter_form import SmallSizeFilterActionForm
20+
from ami.ml.post_processing.class_masking import ClassMaskingTask
21+
from ami.ml.post_processing.rank_rollup import RankRollupTask
1822
from ami.ml.post_processing.small_size_filter import SmallSizeFilterTask
1923
from ami.ml.tasks import remove_duplicate_classifications
2024

@@ -528,6 +532,12 @@ def detections_count(self, obj) -> int:
528532
scope_resolver=lambda occurrence: {"occurrence_id": occurrence.pk},
529533
name_resolver=lambda task_cls, occurrence: (f"Post-processing: {task_cls.name} on Occurrence {occurrence.pk}"),
530534
)
535+
run_class_masking = make_post_processing_action(
536+
ClassMaskingTask,
537+
ClassMaskingActionForm,
538+
scope_resolver=lambda occurrence: {"occurrence_id": occurrence.pk},
539+
name_resolver=lambda task_cls, occurrence: (f"Post-processing: {task_cls.name} on Occurrence {occurrence.pk}"),
540+
)
531541

532542
@admin.action(description="Recompute determination from current classifications and identifications")
533543
def recompute_determination(self, request: HttpRequest, queryset: QuerySet[Any]) -> None:
@@ -544,7 +554,7 @@ def recompute_determination(self, request: HttpRequest, queryset: QuerySet[Any])
544554
count += 1
545555
self.message_user(request, f"Recomputed determination for {count} occurrence(s).")
546556

547-
actions = [run_small_size_filter, recompute_determination]
557+
actions = [run_small_size_filter, run_class_masking, recompute_determination]
548558

549559
# Order by -id (the indexed primary key) rather than -created_at, which has no
550560
# index and would force a full sort of the table to find the newest page. id
@@ -841,11 +851,29 @@ def populate_collection_async(self, request: HttpRequest, queryset: QuerySet[Sou
841851
f"Post-processing: {task_cls.name} on Capture Set {collection.pk}"
842852
),
843853
)
854+
run_class_masking = make_post_processing_action(
855+
ClassMaskingTask,
856+
ClassMaskingActionForm,
857+
scope_resolver=lambda collection: {"source_image_collection_id": collection.pk},
858+
name_resolver=lambda task_cls, collection: (
859+
f"Post-processing: {task_cls.name} on Capture Set {collection.pk}"
860+
),
861+
)
862+
run_rank_rollup = make_post_processing_action(
863+
RankRollupTask,
864+
RankRollupActionForm,
865+
scope_resolver=lambda collection: {"source_image_collection_id": collection.pk},
866+
name_resolver=lambda task_cls, collection: (
867+
f"Post-processing: {task_cls.name} on Capture Set {collection.pk}"
868+
),
869+
)
844870

845871
actions = [
846872
populate_collection,
847873
populate_collection_async,
848874
run_small_size_filter,
875+
run_class_masking,
876+
run_rank_rollup,
849877
]
850878

851879
# Hide images many-to-many field from form. This would list all source images in the database.

‎ami/main/api/serializers.py‎

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -944,10 +944,26 @@ class ClassificationPredictionItemSerializer(serializers.Serializer):
944944
logit = serializers.FloatField(read_only=True)
945945

946946

947+
class ClassificationAppliedToSerializer(serializers.ModelSerializer):
948+
"""Lightweight nested representation of the parent classification this was derived from.
949+
950+
Post-processing tasks (class masking, rank rollup) record provenance via
951+
``Classification.applied_to``; this exposes just enough to show what a result
952+
was derived from without recursing back into the full classification.
953+
"""
954+
955+
algorithm = AlgorithmSerializer(read_only=True)
956+
957+
class Meta:
958+
model = Classification
959+
fields = ["id", "created_at", "algorithm"]
960+
961+
947962
class ClassificationSerializer(DefaultSerializer):
948963
taxon = TaxonNestedSerializer(read_only=True)
949964
algorithm = AlgorithmSerializer(read_only=True)
950965
top_n = ClassificationPredictionItemSerializer(many=True, read_only=True)
966+
applied_to = ClassificationAppliedToSerializer(read_only=True)
951967

952968
class Meta:
953969
model = Classification
@@ -960,6 +976,7 @@ class Meta:
960976
"scores",
961977
"logits",
962978
"top_n",
979+
"applied_to",
963980
"created_at",
964981
"updated_at",
965982
]
@@ -982,6 +999,8 @@ class Meta(ClassificationSerializer.Meta):
982999

9831000

9841001
class ClassificationListSerializer(DefaultSerializer):
1002+
applied_to = ClassificationAppliedToSerializer(read_only=True)
1003+
9851004
class Meta:
9861005
model = Classification
9871006
fields = [
@@ -990,6 +1009,7 @@ class Meta:
9901009
"taxon",
9911010
"score",
9921011
"algorithm",
1012+
"applied_to",
9931013
"created_at",
9941014
"updated_at",
9951015
]
@@ -1009,6 +1029,7 @@ class Meta:
10091029
"score",
10101030
"terminal",
10111031
"algorithm",
1032+
"applied_to",
10121033
"created_at",
10131034
]
10141035

‎ami/main/api/views.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2053,7 +2053,7 @@ class ClassificationViewSet(DefaultViewSet, ProjectMixin):
20532053
"""
20542054

20552055
require_project_for_list = True # Unfiltered list scans are too expensive on this table
2056-
queryset = Classification.objects.all().select_related("taxon", "algorithm") # , "detection")
2056+
queryset = Classification.objects.all().select_related("taxon", "algorithm", "applied_to__algorithm")
20572057
serializer_class = ClassificationSerializer
20582058
filterset_fields = [
20592059
# Docs about slow loading API browser because of large choice fields

‎ami/main/models_future/occurrence.py‎

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -58,7 +58,10 @@ def _detections_prefetch(*, ordering: tuple[str, ...], with_source_image: bool)
5858
qs = Detection.objects.prefetch_related(
5959
Prefetch(
6060
"classifications",
61-
queryset=Classification.objects.select_related("taxon", "algorithm"),
61+
# applied_to__algorithm: post-processed classifications (class masking,
62+
# rank rollup) serialize their provenance parent; pull it here so the
63+
# nested applied_to render doesn't issue a query per classification.
64+
queryset=Classification.objects.select_related("taxon", "algorithm", "applied_to__algorithm"),
6265
)
6366
).order_by(*ordering)
6467
if with_source_image:
Lines changed: 83 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,83 @@
1+
from django.core.management.base import BaseCommand, CommandError
2+
3+
from ami.main.models import SourceImageCollection, TaxaList
4+
from ami.ml.models.algorithm import Algorithm
5+
from ami.ml.post_processing.class_masking import ClassMaskingTask
6+
7+
8+
class Command(BaseCommand):
9+
help = (
10+
"Run class masking post-processing on a source image collection. "
11+
"Masks classifier logits for species not in the given taxa list and recalculates softmax scores."
12+
)
13+
14+
def add_arguments(self, parser):
15+
parser.add_argument("--collection-id", type=int, required=True, help="SourceImageCollection ID to process")
16+
parser.add_argument("--taxa-list-id", type=int, required=True, help="TaxaList ID to use as the species mask")
17+
parser.add_argument(
18+
"--algorithm-id", type=int, required=True, help="Algorithm ID whose classifications to mask"
19+
)
20+
parser.add_argument("--dry-run", action="store_true", help="Show what would be done without making changes")
21+
22+
def handle(self, *args, **options):
23+
collection_id = options["collection_id"]
24+
taxa_list_id = options["taxa_list_id"]
25+
algorithm_id = options["algorithm_id"]
26+
dry_run = options["dry_run"]
27+
28+
# Validate inputs
29+
try:
30+
collection = SourceImageCollection.objects.get(pk=collection_id)
31+
except SourceImageCollection.DoesNotExist:
32+
raise CommandError(f"SourceImageCollection {collection_id} does not exist.")
33+
34+
try:
35+
taxa_list = TaxaList.objects.get(pk=taxa_list_id)
36+
except TaxaList.DoesNotExist:
37+
raise CommandError(f"TaxaList {taxa_list_id} does not exist.")
38+
39+
try:
40+
algorithm = Algorithm.objects.get(pk=algorithm_id)
41+
except Algorithm.DoesNotExist:
42+
raise CommandError(f"Algorithm {algorithm_id} does not exist.")
43+
44+
if not algorithm.category_map:
45+
raise CommandError(f"Algorithm '{algorithm.name}' does not have a category map.")
46+
47+
from ami.main.models import Classification
48+
49+
classification_count = (
50+
Classification.objects.filter(
51+
detection__source_image__collections=collection,
52+
terminal=True,
53+
algorithm=algorithm,
54+
scores__isnull=False,
55+
)
56+
.distinct()
57+
.count()
58+
)
59+
60+
taxa_count = taxa_list.taxa.count()
61+
62+
self.stdout.write(
63+
f"Collection: {collection.name} (id={collection.pk})\n"
64+
f"Taxa list: {taxa_list.name} (id={taxa_list.pk}, {taxa_count} taxa)\n"
65+
f"Algorithm: {algorithm.name} (id={algorithm.pk})\n"
66+
f"Classifications to process: {classification_count}"
67+
)
68+
69+
if classification_count == 0:
70+
raise CommandError("No terminal classifications with scores found for this collection/algorithm.")
71+
72+
if dry_run:
73+
self.stdout.write(self.style.WARNING("Dry run — no changes made."))
74+
return
75+
76+
self.stdout.write("Running class masking...")
77+
task = ClassMaskingTask(
78+
source_image_collection_id=collection_id,
79+
taxa_list_id=taxa_list_id,
80+
algorithm_id=algorithm_id,
81+
)
82+
task.run()
83+
self.stdout.write(self.style.SUCCESS("Class masking completed."))

‎ami/ml/post_processing/__init__.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1 +1,3 @@
1+
from . import class_masking # noqa: F401
2+
from . import rank_rollup # noqa: F401
13
from . import small_size_filter # noqa: F401
Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,38 @@
1+
from __future__ import annotations
2+
3+
from django import forms
4+
5+
from ami.main.models import TaxaList
6+
from ami.ml.models import Algorithm
7+
from ami.ml.models.algorithm import AlgorithmTaskType
8+
from ami.ml.post_processing.admin.forms import BasePostProcessingActionForm
9+
10+
11+
class ClassMaskingActionForm(BasePostProcessingActionForm):
12+
"""Knobs surfaced when an admin triggers Class masking.
13+
14+
The operator picks the source classifier and the taxa list to keep; the
15+
scope (which collection or occurrence) is supplied by the admin entry point,
16+
not the form. Selections are model instances, so ``to_config`` hands the
17+
schema their primary keys (``ClassMaskingConfig`` expects ``*_id`` ints).
18+
"""
19+
20+
algorithm_id = forms.ModelChoiceField(
21+
queryset=Algorithm.objects.filter(task_type=AlgorithmTaskType.CLASSIFICATION.value).order_by("name"),
22+
label="Source classifier",
23+
help_text="The classification algorithm whose terminal predictions will be re-scored.",
24+
)
25+
taxa_list_id = forms.ModelChoiceField(
26+
queryset=TaxaList.objects.all().order_by("name"),
27+
label="Taxa list to keep",
28+
help_text=(
29+
"Classes whose taxon is not in this list are masked out; each "
30+
"classification's softmax is renormalised over the classes that remain."
31+
),
32+
)
33+
34+
def to_config(self) -> dict:
35+
return {
36+
"algorithm_id": self.cleaned_data["algorithm_id"].pk,
37+
"taxa_list_id": self.cleaned_data["taxa_list_id"].pk,
38+
}
Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,13 @@
1+
from __future__ import annotations
2+
3+
from ami.ml.post_processing.admin.forms import BasePostProcessingActionForm
4+
5+
6+
class RankRollupActionForm(BasePostProcessingActionForm):
7+
"""Knob form for Rank rollup.
8+
9+
Rank rollup runs with the per-rank score thresholds and rollup order defined
10+
on ``RankRollupConfig``. There are no per-run knobs yet, so the form only
11+
confirms the selected capture set(s); the empty ``cleaned_data`` lets the
12+
schema apply its defaults. Threshold overrides can be added here later.
13+
"""

0 commit comments

Comments
 (0)