Commit fdc00bd2 authored by Open Science Conservation Fund's avatar Open Science Conservation Fund
Browse files

more robust ai_pipeline celery tasks

parent 1d0d20dc
Loading
Loading
Loading
Loading
+8 −5
Original line number Diff line number Diff line
@@ -153,11 +153,11 @@ class ProjectRoleInline(admin.TabularInline):
class ProjectCollectionInline(admin.TabularInline):
    model = ClassificationProjectCollection
    extra = 0
    raw_id_fields = ["collection"]
    readonly_fields = [
    raw_id_fields = ("collection",)
    readonly_fields = (
        "resources_count",
        "collection_storage",
    ]
    )


class ClassificationProjectAdmin(admin.ModelAdmin):
@@ -179,9 +179,9 @@ class ClassificationProjectAdmin(admin.ModelAdmin):
        "disabled_at",
        "disabled_by",
    )
    exclude = ("deployments",)
    list_filter = ("research_project", "citizen_science_status")
    readonly_fields = ("slug_id", "slug", "date_created")
    autocomplete_fields = ("deployments",)


class ClassificationProjectRoleAdmin(admin.ModelAdmin):
@@ -223,9 +223,12 @@ class ClassificationProjectCollectionAdmin(admin.ModelAdmin):
        "project",
        "collection",
        "is_active",
        "ai_pipeline_status",
        "resources_count",
        "collection_storage",
    )
    search_fields = ("collection__collection__name",)
    list_filter = ("project", "is_active")
    list_filter = ("project", "is_active", "ai_pipeline_status")
    actions = ["rerun_ai_pipeline"]

    @admin.action(
+10 −7
Original line number Diff line number Diff line
import json
import logging
import uuid
from dataclasses import dataclass
from enum import Enum
@@ -28,6 +29,8 @@ from trapper.apps.media_classification.models import (
    AIClassificationJobAdditionalResource,
)

logger = logging.getLogger(__name__)


@dataclass_json(undefined=Undefined.EXCLUDE)
@dataclass
@@ -399,12 +402,12 @@ class TrapperAIProviderManager(BaseAIProviderManager):
                    # Labels and bboxes from Object Detection model
                    bboxes=classification.get_bboxes(),
                )
                print("[DEBUG] TrapperAISample with bboxes (crop_image)")
                logger.debug("TrapperAISample with bboxes (crop_image)")
            else:
                trapper_ai_sample = TrapperAISample(
                    sample_id=classification.pk, sample_url=url
                )
                print("[DEBUG] TrapperAISample without bboxes (native image)")
                logger.debug("TrapperAISample without bboxes (native image)")

            data.append(trapper_ai_sample)

@@ -513,20 +516,20 @@ class TrapperAIProviderManager(BaseAIProviderManager):
        task.save()

    def _fetch_and_process_batch_results(self, job: AIClassificationJob, batch_id: str):
        print("_fetch_batch_results")
        logger.debug("Fetching batch results")
        results = self.trapper_ai_service.fetch_detections(batch_id=batch_id)

        detection_category_mapping = {
            ot.value: ot.observation_type
            for ot in self.config.observation_type_mappings.all()
        }
        print("Detection category mapping", detection_category_mapping)
        logger.debug("Detection category mapping: %s", detection_category_mapping)

        species_category_mapping = {
            ot.value: ot.species_id
            for ot in self.config.species_ai_label_mappings.all()
        }
        print("Species category mapping", species_category_mapping)
        logger.debug("Species category mapping: %s", species_category_mapping)

        ai_classifications = []
        all_detected_objects = {}
@@ -540,7 +543,7 @@ class TrapperAIProviderManager(BaseAIProviderManager):
            h = y2 - y
            return [x, y, w, h]

        print(f"Found {len(results)} observations")
        logger.info("Found %d observations", len(results))

        for result in results:
            classification_id = result.sample_id
@@ -621,7 +624,7 @@ class TrapperAIProviderManager(BaseAIProviderManager):
            all_detected_objects[classification_id] = detected_objects

        ai_classifications = self.create_ai_classification_objects(ai_classifications)
        print(f"Created {len(ai_classifications)} AIClassification objects")
        logger.info("Created %d AIClassification objects", len(ai_classifications))

        ai_classification_dynamic_objects = []
        for ai_classification in ai_classifications:
+18 −0
Original line number Diff line number Diff line
# Generated by Django 4.2.17 on 2025-08-12 13:35

from django.db import migrations, models


class Migration(migrations.Migration):

    dependencies = [
        ('media_classification', '0085_remove_classificationproject_default_ai_model'),
    ]

    operations = [
        migrations.AddField(
            model_name='classificationprojectcollection',
            name='ai_pipeline_status',
            field=models.CharField(blank=True, choices=[('STARTED', 'STARTED'), ('SUCCESS', 'SUCCESS'), ('REJECTED', 'REJECTED'), ('FAILURE', 'FAILURE')], max_length=10, null=True),
        ),
    ]
+4 −0
Original line number Diff line number Diff line
@@ -26,6 +26,7 @@ from polymorphic.managers import PolymorphicManager
from polymorphic.models import PolymorphicModel

from trapper.apps.accounts.models import UserRemoteTask, User
from trapper.apps.accounts.taxonomy import UserRemoteTaskStatus
from trapper.apps.common.fields import SafeTextField, ResizedImageField
from trapper.apps.common.utils.models import checkbox2select
from trapper.apps.extra_tables.models import Species
@@ -705,6 +706,9 @@ class ClassificationProjectCollection(models.Model):
    )
    collection = models.ForeignKey(ResearchProjectCollection, on_delete=models.PROTECT)
    is_active = models.BooleanField(_("Active"), default=True)
    ai_pipeline_status = models.CharField(
        choices=UserRemoteTaskStatus.CHOICES, max_length=10, null=True, blank=True
    )

    objects = ClassificationProjectCollectionManager()

+95 −45
Original line number Diff line number Diff line
@@ -3,7 +3,6 @@ import logging
from celery import shared_task

from django.contrib.auth import get_user_model
from django.core.cache import caches

from trapper.apps.accounts.taxonomy import UserRemoteTaskStatus
from trapper.apps.media_classification.ai_providers.ai_provider_factory import (
@@ -24,6 +23,7 @@ from trapper.apps.media_classification.models import (
    ClassificationProjectCollection,
    AIProvider,
)
from trapper.apps.media_classification.taxonomy import ObservationType
from trapper.apps.media_classification.tasks import (
    celery_approve_ai_classifications,
    celery_blur_humans,
@@ -59,6 +59,7 @@ def start_ai_object_detection(
    target_fps = project.target_fps

    if not ai_model:
        collection.ai_pipeline_status = UserRemoteTaskStatus.REJECTED
        logger.error("AI provider is not set for Citizen Science")

    elif classification_ids:
@@ -79,18 +80,25 @@ def start_ai_object_detection(
                (job_id, project.pk, collection_pk, user_pk, overwrite),
                countdown=7,
            )
            collection.ai_pipeline_status = UserRemoteTaskStatus.STARTED

        except AIProviderException:
            collection.ai_pipeline_status = UserRemoteTaskStatus.FAILURE
            logger.error(
                f"Failed to start AI object detection for project {project.pk} "
                "and collection {collection_pk}"
                f"and collection {collection_pk}"
            )

    else:
        collection.ai_pipeline_status = UserRemoteTaskStatus.REJECTED
        logger.info(
            f"No classifications to process for project {project.pk} "
            f"and collection {collection_pk}. Skipping AI object detection."
        )

    # update collection status
    collection.save(update_fields=["ai_pipeline_status"])


@shared_task(serializer="json")
def check_ai_object_detection_finished(
@@ -111,35 +119,41 @@ def check_ai_object_detection_finished(
        classification_job = AIClassificationJob.objects.get(pk=classification_job_pk)
    except AIClassificationJob.DoesNotExist:
        return

    status = classification_job.user_remote_task.status
    collection = ClassificationProjectCollection.objects.get(pk=collection_pk)

    # The AI object detection job was successful
    if status == UserRemoteTaskStatus.SUCCESS:

        ai_classification_pks = AIClassification.objects.filter(
            classification__project_id=classification_project_pk,
            classification__collection_id=collection_pk,
            classification__approved_source_ai__isnull=True,
        # First, get Classification objects
        if overwrite:
            qs = collection.classifications.filter()
        else:
            qs = collection.classifications.filter(approved_source_ai__isnull=True)
        classification_ids = list(qs.values_list("pk", flat=True))

        # Get AIClassification objects related to the classifications
        # and the AI provider used for the classification job
        ai_classifications = AIClassification.objects.filter(
            classification__in=classification_ids,
            model_id=classification_job.ai_provider_id,
        ).values_list("pk", flat=True)
        )
        ai_classifications_ids = list(ai_classifications.values_list("pk", flat=True))

        project = ClassificationProject.objects.get(pk=classification_project_pk)
        user = User.objects.get(pk=user_pk)

        fields_to_copy = ["observation_type"]
        should_mark_approved = not (
            project.copy_ai_classifications and project.species_ai_model
        )

        logger.info(
            f"Running `celery_approve_ai_classifications` with mark_as_approved={should_mark_approved}."
        )
        # Run `celery_approve_ai_classifications` synchronously because blurring requires
        # `Classification` objects to have the `observation_type` already set
        logger.info("Running `celery_approve_ai_classifications`")
        msg = celery_approve_ai_classifications(
            user=user,
            project_id=classification_project_pk,
            ai_classification_pks=ai_classification_pks,
            ai_classification_pks=ai_classifications_ids,
            fields_to_copy=fields_to_copy,
            mark_as_approved=should_mark_approved,
            mark_as_approved=True,
            minimum_confidence=classification_job.ai_provider.minimum_confidence,
            overwrite_attrs=True,
            copy_bboxes=True,
@@ -161,42 +175,39 @@ def check_ai_object_detection_finished(
                AIProvider.ANONIMIZE_HUMAN,
                AIProvider.ANONIMIZE_HUMAN_AND_VEHICLE,
            ]:
                logger.info("Running celery_blur_humans.")
                msg = celery_blur_humans(**celery_blur_params)
                logger.info("Running `celery_blur_humans`")
                msg = celery_blur_humans.delay(**celery_blur_params)
                logger.info(msg)
            else:
                logger.info("Human blur filter was ignored")
                logger.info("Blurring human observations was skipped")

            # Execute Vehicle blur filter if AIProvider configuration support it
            if classification_job.ai_provider.anonimize_classes in [
                AIProvider.ANONIMIZE_VEHICLE,
                AIProvider.ANONIMIZE_HUMAN_AND_VEHICLE,
            ]:
                logger.info("Running celery_blur_vehicles.")
                msg = celery_blur_vehicles(**celery_blur_params)
                logger.info("Running `celery_blur_vehicles`")
                msg = celery_blur_vehicles.delay(**celery_blur_params)
                logger.info(msg)
            else:
                logger.info("Vehicle blur filter was ignored")
                logger.info("Blurring vehicle observations was skipped")
        else:
            logger.info("Human and Vehicle blur filter was ignored")
            logger.info("Blurring human and vehicle observations was skipped")

        # TODO: clear cache only for the current user (!)
        # Clear cache because TrapperPaginator caches count result, and doesn't
        # refresh queryset,
        caches["default"].clear()
        collection.ai_pipeline_status = UserRemoteTaskStatus.SUCCESS

        # Run Trapper AI using the species model if specified in the project's configuration
        # The species model is run only for animal observations (!)
        if project.species_ai_model:

            collection = ClassificationProjectCollection.objects.get(pk=collection_pk)

            if overwrite:
                qs = collection.classifications.all()
            else:
                qs = collection.classifications.filter(approved_source_ai__isnull=True)
            classification_ids = list(qs.values_list("pk", flat=True))
            qs = ai_classifications.filter(
                dynamic_attrs__observation_type=ObservationType.ANIMAL
            ).distinct()
            classification_ids = list(qs.values_list("classification__pk", flat=True))

            if not classification_ids:
                # update collection status
                collection.save(update_fields=["ai_pipeline_status"])
                logger.info(
                    f"No classifications to process for project {classification_project_pk} "
                    f"and collection {collection_pk}. Skipping species classification."
@@ -204,8 +215,8 @@ def check_ai_object_detection_finished(
                return

            logger.info(
                f"Running TrapperAI species classification for project {classification_project_pk}"
                + f" and collection {collection_pk}."
                f"Running TrapperAI species classification for project {classification_project_pk},"
                + f" collection {collection_pk} and {len(classification_ids)} animal observations."
            )
            manager: TrapperAIProviderManager = get_ai_provider_manager(
                project.species_ai_model
@@ -226,22 +237,42 @@ def check_ai_object_detection_finished(
                )

                check_ai_species_classification_finished.apply_async(
                    (job_id, classification_project_pk, collection_pk, user_pk),
                    (
                        job_id,
                        classification_project_pk,
                        classification_ids,
                        collection_pk,
                        user_pk,
                    ),
                    countdown=7,
                )

                collection.ai_pipeline_status = UserRemoteTaskStatus.STARTED
                # persist status change when scheduling species step
                collection.save(update_fields=["ai_pipeline_status"])

            except AIProviderException as e:
                collection.ai_pipeline_status = UserRemoteTaskStatus.FAILURE
                logger.error(
                    f"Failed to start AI species classification for project {classification_project_pk} "
                    f"and collection {collection_pk}: {e}"
                )
        else:
            # No species model configured; finalize SUCCESS now
            collection.save(update_fields=["ai_pipeline_status"])
            logger.info(
                f"No species AI model configured for project {classification_project_pk}. "
                "Skipping species classification."
            )

    elif status in [UserRemoteTaskStatus.REJECTED, UserRemoteTaskStatus.FAILURE]:
        # update collection status
        collection.ai_pipeline_status = status
        collection.save(update_fields=["ai_pipeline_status"])
        logger.error(
            f"Classification job {classification_job_pk} failed or rejected. Unable to complete "
            "collection processing."
        )
        return

    else:
        check_ai_object_detection_finished.apply_async(
@@ -260,6 +291,7 @@ def check_ai_object_detection_finished(
def check_ai_species_classification_finished(
    classification_job_pk: int,
    classification_project_pk: int,
    classification_ids: list[int],
    collection_pk: int,
    user_pk: int,
):
@@ -274,27 +306,35 @@ def check_ai_species_classification_finished(
    except AIClassificationJob.DoesNotExist:
        return
    status = classification_job.user_remote_task.status
    collection = ClassificationProjectCollection.objects.get(pk=collection_pk)

    # The AI species classification job was successful
    if status == UserRemoteTaskStatus.SUCCESS:

        project = ClassificationProject.objects.get(pk=classification_project_pk)

        # The `copy_ai_classifications` setting indicates that the AIClassification attributes
        # (specifically the `species` attribute and corresponding `bboxes`)
        # should be merged with the attributes already approved for Classification objects.
        if not project.copy_ai_classifications:
            collection.ai_pipeline_status = UserRemoteTaskStatus.SUCCESS
            collection.save(update_fields=["ai_pipeline_status"])
            return

        user = User.objects.get(pk=user_pk)
        iou_threshold = project.species_matching_iou_threshold

        ai_classifications = AIClassification.objects.filter(
            classification__project_id=classification_project_pk,
            classification__collection_id=collection_pk,
        ai_classification_ids = list(
            AIClassification.objects.filter(
                classification_id__in=classification_ids,
                model_id=classification_job.ai_provider_id,
            ).values_list("pk", flat=True)
        )

        msg = celery_approve_ai_classifications(
        msg = celery_approve_ai_classifications.delay(
            user=user,
            project_id=classification_project_pk,
            ai_classification_pks=ai_classifications.values_list("pk", flat=True),
            ai_classification_pks=ai_classification_ids,
            fields_to_copy=["species"],
            mark_as_approved=True,
            minimum_confidence=classification_job.ai_provider.minimum_confidence,
@@ -303,16 +343,26 @@ def check_ai_species_classification_finished(
            iou_threshold=iou_threshold,
        )
        logger.info(msg)
        collection.ai_pipeline_status = UserRemoteTaskStatus.SUCCESS

    elif status in [UserRemoteTaskStatus.REJECTED, UserRemoteTaskStatus.FAILURE]:
        collection.ai_pipeline_status = status
        logger.error(
            f"Classification job {classification_job_pk} failed or rejected. Unable to complete "
            "collection processing."
        )
        return

    else:
        check_ai_species_classification_finished.apply_async(
            (classification_job_pk, classification_project_pk, collection_pk, user_pk),
            (
                classification_job_pk,
                classification_project_pk,
                classification_ids,
                collection_pk,
                user_pk,
            ),
            countdown=7,
        )

    # update collection status
    collection.save(update_fields=["ai_pipeline_status"])