|
| 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.")) |
0 commit comments