Skip to content
Open
31 changes: 31 additions & 0 deletions ami/jobs/migrations/0024_train_classifier_job_type.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,31 @@
# Generated by Django 4.2.10 on 2026-10-08 15:53

from django.db import migrations, models


class Migration(migrations.Migration):
dependencies = [
("jobs", "0023_alter_job_job_type_key"),
]

operations = [
migrations.AlterField(
model_name="job",
name="job_type_key",
field=models.CharField(
choices=[
("ml", "ML pipeline"),
("populate_captures_collection", "Populate capture set"),
("data_storage_sync", "Data storage sync"),
("regroup_events", "Regroup sessions"),
("unknown", "Unknown"),
("data_export", "Data Export"),
("post_processing", "Post Processing"),
("train_classifier", "Train classifier"),
],
default="unknown",
max_length=255,
verbose_name="Job Type",
),
),
]
433 changes: 431 additions & 2 deletions ami/jobs/models.py

Large diffs are not rendered by default.

37 changes: 36 additions & 1 deletion ami/jobs/serializers.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
)
from ami.main.models import Deployment, Project, SourceImage, SourceImageCollection
from ami.ml.models import Pipeline
from ami.ml.schemas import PipelineProcessingTask, PipelineTaskResult, ProcessingServiceClientInfo
from ami.ml.schemas import PipelineProcessingTask, PipelineTaskResult, ProcessingServiceClientInfo, TrainingResult
from ami.ml.serializers import PipelineNestedSerializer

from .models import (
Expand Down Expand Up @@ -292,6 +292,41 @@ class MLJobResultsRequestSerializer(serializers.Serializer):
client_info = SchemaField(schema=ProcessingServiceClientInfo, required=False, default=None)


class TrainingResultRequestSerializer(serializers.Serializer):
"""POST /jobs/{id}/training-result/ — the body a processing service posts when a run finishes.

The counterpart of MLJobResultsRequestSerializer for a retraining job: one result for
one job, rather than a list of per-item results.

``result`` is validated, since it is what the new algorithm version is built from.
``dataset`` is the metadata the service echoes back from the training set file and is
taken as it comes: a service that echoes an older shape should still have its result
recorded, and the cost is a version that cannot say which occurrence set it learned
from. It is parsed, leniently, where it is read.
"""

result = SchemaField(schema=TrainingResult)
dataset = serializers.JSONField(required=False, allow_null=True, default=None)
dataset_url = serializers.CharField(required=False, allow_null=True, default=None)
job_id = serializers.IntegerField(required=False, allow_null=True, default=None)
algorithm_key = serializers.CharField(required=False, allow_null=True, default=None)


class TrainingResultResponseSerializer(serializers.Serializer):
"""POST /jobs/{id}/training-result/ — acknowledgment returned to the processing service."""

status = serializers.CharField()
job_id = serializers.IntegerField()
algorithm = serializers.CharField(allow_null=True, help_text="Key of the version registered, if one was.")


class TrainingProgressRequestSerializer(serializers.Serializer):
"""POST /jobs/{id}/training-progress/ — how far through its epochs a run has got."""

epoch = serializers.IntegerField(min_value=0)
total_epochs = serializers.IntegerField(required=False, allow_null=True, default=None, min_value=1)


class MLJobResultsResponseSerializer(serializers.Serializer):
"""POST /jobs/{id}/result/ — acknowledgment returned to the processing service.

Expand Down
159 changes: 159 additions & 0 deletions ami/jobs/views.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,13 +7,16 @@
from django.core.cache import cache
from django.db.models import Q
from django.db.models.query import QuerySet
from django.http import Http404
from django.utils import timezone
from django_filters import rest_framework as filters
from drf_spectacular.utils import extend_schema, extend_schema_view
from rest_framework import serializers
from rest_framework.decorators import action
from rest_framework.exceptions import NotFound, PermissionDenied, ValidationError
from rest_framework.filters import BaseFilterBackend
from rest_framework.parsers import MultiPartParser
from rest_framework.permissions import AllowAny
from rest_framework.response import Response

from ami.base.filters import RelatedIdFilter
Expand All @@ -28,6 +31,9 @@
MLJobResultsResponseSerializer,
MLJobTasksRequestSerializer,
MLJobTasksResponseSerializer,
TrainingProgressRequestSerializer,
TrainingResultRequestSerializer,
TrainingResultResponseSerializer,
)
from ami.jobs.tasks import (
HEARTBEAT_THROTTLE_SECONDS,
Expand Down Expand Up @@ -591,3 +597,156 @@ def result(self, request, pk=None):
},
status=503,
)

def _job_for_callback(self, pk, request) -> Job:
"""
The job a processing service is reporting about, if its token proves it may.

get_object() applies project visibility, which an unauthenticated service fails,
so the job is looked up directly and the signed token is what authorises the call.
"""
from ami.ml.training import verify_callback_token

job = Job.objects.filter(pk=pk).first()
if not job:
raise Http404("Job not found.")

token = request.headers.get("Authorization", "").removeprefix("Token ").strip()
if not verify_callback_token(token, job):
raise PermissionDenied("Invalid or expired training callback token.")
return job

@extend_schema(
request=TrainingResultRequestSerializer,
responses={200: TrainingResultResponseSerializer},
)
@action(
detail=True,
methods=["post"],
url_path="training-result",
name="training-result",
# A processing service has no Antenna account. It proves itself with the signed
# token Antenna issued when it dispatched the job, checked below.
permission_classes=[AllowAny],
authentication_classes=[],
)
def training_result(self, request, pk=None):
"""
Receive the outcome of a retraining run from a processing service.

Training outlasts the request that started it, so this is the one way a result
comes back. A service that also answers its caller inline is acknowledged and that
body ignored, so there is a single shape to read.
"""
from ami.jobs.models import TrainClassifierJob

job = self._job_for_callback(pk, request)

if job.job_type_key != TrainClassifierJob.key:
raise ValidationError(f"Job #{job.pk} is not a training job.")

if job.status in JobState.final_states():
# A retry of a callback that already landed. The first answer stands.
logger.info("Ignoring a training result for job %s, which already finished", job.pk)
return Response({"status": "already recorded", "job_id": job.pk, "algorithm": None})

serializer = TrainingResultRequestSerializer(data=request.data)
serializer.is_valid(raise_exception=True)

new_version = TrainClassifierJob.record_result(
job=job,
result=serializer.validated_data["result"],
dataset=serializer.validated_data["dataset"],
dataset_url=serializer.validated_data["dataset_url"],
)
logger.info("Recorded a training result for job %s", job.pk)
return Response(
{
"status": "recorded",
"job_id": job.pk,
"algorithm": new_version.key if new_version else None,
}
)

@extend_schema(exclude=True)
@action(
detail=True,
methods=["post"],
url_path="training-progress",
name="training-progress",
# Same token as the result callback: a processing service has no Antenna account.
permission_classes=[AllowAny],
authentication_classes=[],
)
def training_progress(self, request, pk=None):
"""
Receive how far through its epochs a retraining run is.

Training is the long stage of the job and the service says nothing while it fits, so
this is what moves the progress bar in between.
"""
from ami.jobs.models import TrainClassifierJob

job = self._job_for_callback(pk, request)

if job.job_type_key != TrainClassifierJob.key:
raise ValidationError(f"Job #{job.pk} is not a training job.")

if job.status in JobState.final_states():
# A ping overtaken by the result. The job is done; its stages say so.
return Response({"status": "finished"})

serializer = TrainingProgressRequestSerializer(data=request.data)
serializer.is_valid(raise_exception=True)

TrainClassifierJob.record_progress(
job=job,
epoch=serializer.validated_data["epoch"],
total_epochs=serializer.validated_data["total_epochs"],
)
return Response({"status": "recorded"})

@extend_schema(exclude=True)
@action(
detail=True,
methods=["post"],
url_path="training-head",
name="training-head",
# Same token as the result callback: a processing service has no Antenna account.
permission_classes=[AllowAny],
authentication_classes=[],
parser_classes=[MultiPartParser],
)
def training_head(self, request, pk=None):
"""
Receive the head a retraining run produced, so Antenna keeps a copy of the weights.

Without this the head exists only on the service's disk, under a cache directory,
and Antenna records that a version exists without being able to say where it is.
"""
from ami.jobs.models import TrainClassifierJob
from ami.ml.training import HeadTooLarge, store_head

job = self._job_for_callback(pk, request)

if job.job_type_key != TrainClassifierJob.key:
raise ValidationError(f"Job #{job.pk} is not a training job.")

if job.status in JobState.final_states():
# The token stays valid for 24 hours and the head is stored at a path fixed by
# the job, so without this the weights a registered version points at could be
# replaced for a day after the run ended. A service uploads before it reports
# its result, so a run that is still going is unaffected.
raise ValidationError(f"Job #{job.pk} has already finished; its head can no longer be replaced.")

if not request.FILES:
raise ValidationError("No head files were uploaded.")

algorithm_key = (job.params or {}).get("algorithm_key") or "head"
try:
stored = store_head(algorithm_key=algorithm_key, job_id=job.pk, files=request.FILES)
except HeadTooLarge as e:
raise ValidationError(str(e))

job.logger.info(f"Stored the retrained head: {', '.join(item['path'] for item in stored.values())}")
return Response({"files": stored})
58 changes: 58 additions & 0 deletions ami/ml/schemas.py
Original file line number Diff line number Diff line change
Expand Up @@ -473,6 +473,64 @@ class AsyncPipelineRegistrationRequest(pydantic.BaseModel):
pipelines: list[PipelineConfigResponse] = []


class TrainingRequest(pydantic.BaseModel):
"""
What Antenna asks a processing service to do, mirroring the service's own TrainRequest.

Declared rather than assembled as a dict so the two sides of the contract are written
down in one place, the way the pipeline request and response already are.
"""

dataset_url: str = pydantic.Field(description="Where the service downloads the training set from.")
algorithm_key: str = pydantic.Field(description="The head to retrain. Its current weights are the baseline.")
job_id: int
name: str

# The settings the training set was built under, plus the ones the service fits with.
min_per_species: int
min_improvement: float
head_type: str
epochs: int
learning_rate: float
weight_decay: float

# Where the service reports back. The result callback is the only way a run that
# outlasts its request can finish, so it is always sent; the other two are optional for
# the service and ignored by one that does not support them.
callback_url: str
callback_token: str
head_upload_url: str | None = None
progress_url: str | None = None


class TrainingResult(pydantic.BaseModel):
"""
What a processing service reports when a run finishes, mirroring its TrainResponse.

Extra fields are kept: a newer service may report more than this one knows about, and
dropping it silently would lose the only record of what a run did.
"""

promote: bool = False
reason: str = ""
warnings: list[str] = pydantic.Field(default_factory=list)
labels: list[str] = pydantic.Field(default_factory=list, description="The new head's classes, in order.")
rows: dict = pydantic.Field(default_factory=dict, description="How many rows were used, kept and dropped.")
counts: dict = pydantic.Field(default_factory=dict)
dropped_species: list[str] = pydantic.Field(default_factory=list)
classes_restored_from_current_head: int = 0
candidate_metrics: dict = pydantic.Field(default_factory=dict, description="The new head, on the held-out rows.")
# Null when there was no current head to score, which a service reports rather than
# leaving out, so the field has to accept it.
incumbent_metrics: dict | None = pydantic.Field(default=None, description="The current head, on the same rows.")
saved: dict[str, str] | None = None
head_url: str | None = pydantic.Field(default=None, description="Where the uploaded head landed, if it was sent.")
trained_at: str | None = None

class Config:
extra = "allow"


class AlgorithmTrainingInfo(pydantic.BaseModel):
"""
What actually happened when this version was trained. Written by the service, read-only.
Expand Down
Loading