diff --git a/ami/jobs/migrations/0024_train_classifier_job_type.py b/ami/jobs/migrations/0024_train_classifier_job_type.py new file mode 100644 index 000000000..314738095 --- /dev/null +++ b/ami/jobs/migrations/0024_train_classifier_job_type.py @@ -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", + ), + ), + ] diff --git a/ami/jobs/models.py b/ami/jobs/models.py index 20d03a54c..c8ec4753f 100644 --- a/ami/jobs/models.py +++ b/ami/jobs/models.py @@ -17,13 +17,17 @@ from ami.base.models import BaseModel from ami.base.schemas import ConfigurableStage, ConfigurableStageParam from ami.jobs.tasks import cleanup_async_job_if_needed, run_job -from ami.main.models import Deployment, Project, SourceImage, SourceImageCollection +from ami.main.models import Deployment, OccurrenceSet, Project, SourceImage, SourceImageCollection from ami.ml.models import Pipeline from ami.ml.post_processing.registry import get_postprocessing_task from ami.utils.schemas import OrderedEnum logger = logging.getLogger(__name__) +if typing.TYPE_CHECKING: + from ami.ml.models import Algorithm + from ami.ml.schemas import TrainingResult + class JobDispatchMode(models.TextChoices): """ @@ -997,6 +1001,429 @@ def run(cls, job: "Job"): raise ValueError(f"Unknown job type '{job.job_type()}'") +class TrainClassifierJob(JobType): + """ + Retrain a classifier head from the species people have verified in this project. + + Antenna prepares the training set and hands the processing service a URL to it, rather + than posting the rows: a project with a few hundred thousand verified labels runs to + hundreds of megabytes, which is fragile to send in one request and has to start over if + the connection drops. + + The job does not wait for training to finish. It dispatches and leaves the job STARTED; + the service reports back through the job's result endpoint. Waiting inline would hit the + sub-task timeout on anything but a toy dataset. + """ + + name = "Train classifier" + key = "train_classifier" + required_params = ("algorithm_key",) + user_creatable = True + + # This job dispatches and then waits for the service to report back, writing nothing to + # its own row in between, so the default threshold reads a healthy run as a dead one. + # Matched to the life of the callback token: once that expires Antenna refuses the + # callback anyway, so a job still waiting past it can never finish. + stalled_after_minutes = 24 * 60 + + STAGE_PREPARE = "prepare" + STAGE_DISPATCH = "dispatch" + STAGE_TRAIN = "train" + + # Reported by the service while it fits, and read back to move the training stage. + PARAM_EPOCH = "Epoch" + PARAM_TOTAL_EPOCHS = "Total epochs" + + @classmethod + def run(cls, job: "Job"): + from ami.ml.models import Algorithm + from ami.ml.models.processing_service import ProcessingService + from ami.ml.training import NotEnoughVerifiedData, build_training_dataset + + params = job.params or {} + algorithm_key = params.get("algorithm_key") + if not algorithm_key: + raise ValueError("A train_classifier job needs an 'algorithm_key' in its params.") + + # A key is shared by every version of an algorithm, and embeddings are stored per + # version row, so the newest is chosen rather than whichever row came back first. + algorithm = Algorithm.objects.filter(key=algorithm_key).order_by("-version", "-pk").first() + if not algorithm: + raise ValueError(f"No algorithm with key '{algorithm_key}'.") + if not algorithm.trainable: + raise ValueError( + f"Algorithm '{algorithm_key}' is not marked trainable by its processing service. " + "Re-register the pipelines if the service has since been updated." + ) + + # Settled before any stage exists, so a bad setting is refused outright rather than + # failing halfway through a run that already looks started. + occurrence_set = cls.target_occurrence_set(job) + config = cls.config_for(job, algorithm) + + job.progress.add_stage("Preparing training set", cls.STAGE_PREPARE) + job.progress.add_stage("Sending to processing service", cls.STAGE_DISPATCH) + job.progress.add_stage("Training", cls.STAGE_TRAIN) + # The work goes to an external service that reports back on its own, which is what + # this mode means. Set here rather than in Job.setup(), which reads it off the + # pipeline a training job does not have. + job.dispatch_mode = JobDispatchMode.ASYNC_API + job.update_status(JobState.STARTED) + job.started_at = datetime.datetime.now() + job.finished_at = None + job.progress.update_stage(cls.STAGE_PREPARE, status=JobState.STARTED, progress=0) + job.save() + + taxa_list = cls.target_taxa_list(job) + try: + dataset = build_training_dataset( + project=job.project, + algorithm=algorithm, + min_per_species=config.min_per_species, + test_fraction=config.test_fraction, + split_salt=config.split_salt, + taxa_list=taxa_list, + occurrence_set=occurrence_set, + job=job, + ) + except NotEnoughVerifiedData as e: + # A data problem, not a crash. Say so plainly rather than failing with a traceback. + job.logger.error(str(e)) + job.progress.update_stage(cls.STAGE_PREPARE, status=JobState.FAILURE, progress=0) + job.finished_at = datetime.datetime.now() + job.result = {"error": str(e)} + job.update_status(JobState.FAILURE, save=True) + return + + meta = dataset["metadata"] + job.logger.info( + f"Training set: {meta.rows} verified crops over {len(meta.classes)} species " + f"({meta.train} train / {meta.test} held out)" + ) + if meta.taxa_list: + job.logger.info( + f"Species list comes from taxa list '{meta.taxa_list.name}' " f"({len(meta.classes)} species)" + ) + without_data = meta.classes_without_verified_data or [] + if without_data: + job.logger.info( + f"{len(without_data)} species in the list have no verified crops yet. They keep " + "whatever the current head already knows." + ) + if meta.dropped_species: + job.logger.warning( + f"{len(meta.dropped_species)} verified species are not in the taxa list and were " + f"left out: {', '.join(meta.dropped_species[:5])}" + f"{' ...' if len(meta.dropped_species) > 5 else ''}" + ) + if meta.verified_detections_without_embedding: + job.logger.warning( + f"{meta.verified_detections_without_embedding} verified detection(s) have no embedding " + "from this algorithm and were left out. Re-run the pipeline over them to include them." + ) + job.progress.add_stage_param(cls.STAGE_PREPARE, "Rows", meta.rows) + job.progress.add_stage_param(cls.STAGE_PREPARE, "Species", len(meta.classes)) + job.progress.add_stage_param(cls.STAGE_PREPARE, "Dataset", dataset["url"]) + job.progress.update_stage(cls.STAGE_PREPARE, status=JobState.SUCCESS, progress=1) + job.save() + + service = ( + ProcessingService.objects.filter( + projects=job.project, + pipelines__algorithms=algorithm, + endpoint_url__isnull=False, + ) + .exclude(endpoint_url="") + .distinct() + .first() + ) + if not service: + # Pull-mode workers register with a null endpoint_url, so there is nothing to + # send a training request to. Supporting them means routing this through the + # task queue instead. + raise ValueError( + f"No push-mode processing service in this project serves '{algorithm_key}'. " + "Training cannot be dispatched to a pull-mode worker." + ) + + job.progress.update_stage(cls.STAGE_DISPATCH, status=JobState.STARTED, progress=0) + job.save() + cls.dispatch(job=job, service=service, algorithm=algorithm, dataset=dataset) + + @classmethod + def config_for(cls, job: "Job", algorithm): + """ + The training settings for this run, checked before anything uses them. + + The algorithm's training_config holds the defaults its processing service + published; a job may override any of them for one run. Those overrides arrive in + job.params, which any member who can create a job sets, so they are validated here + instead of being read raw at each point of use. + """ + import pydantic + + from ami.ml.schemas import AlgorithmTrainingConfig + + try: + return AlgorithmTrainingConfig.for_run(algorithm.training_config, job.params or {}) + except pydantic.ValidationError as e: + problems = "; ".join(f"{'.'.join(str(part) for part in err['loc'])} {err['msg']}" for err in e.errors()) + raise ValueError(f"This job's training settings are out of range: {problems}") + + @classmethod + def target_occurrence_set(cls, job: "Job") -> OccurrenceSet | None: + """ + The set of occurrences this run learns from, or None for every verified one. + + Naming a set is worth doing: a set is fixed once created, so the run can be repeated + and says for itself what it learned from. It is not required, because the common case + is retraining on everything verified so far, and making someone save a set first to + do that is a step with no decision in it. + + The dataset file records which set was used, or that there was none, so a run without + one is still legible afterwards. + """ + occurrence_set_id = (job.params or {}).get("occurrence_set_id") + if not occurrence_set_id: + return None + # Scoped to the project, since params are set by any member who can create a job. + occurrence_set = OccurrenceSet.objects.for_project(job.project).filter(pk=occurrence_set_id).first() + if not occurrence_set: + raise ValueError(f"No occurrence set with id {occurrence_set_id} in this project.") + return occurrence_set + + @classmethod + def target_taxa_list(cls, job: "Job"): + """ + Which taxa list sets the head's species. + + A job may name one; otherwise the project's default is used. Without either, the + class list falls back to whatever has been verified. + """ + from ami.main.models import TaxaList + + taxa_list_id = (job.params or {}).get("taxa_list_id") + if taxa_list_id: + # Scoped for the same reason as the evaluation set above. + taxa_list = TaxaList.objects.for_project(job.project).filter(pk=taxa_list_id).first() + if not taxa_list: + raise ValueError(f"No taxa list with id {taxa_list_id} in this project.") + return taxa_list + return job.project.default_taxa_list + + @classmethod + def dispatch(cls, job: "Job", service, algorithm, dataset: dict) -> None: + """Hand the service the dataset URL and leave the job running until it reports back.""" + from ami.ml.training import send_training_request + + # Written before the request so the stage reads "0 of 300" while the service is + # still starting, rather than an empty pair of parameters. + config = cls.config_for(job, algorithm) + job.progress.add_stage_param(cls.STAGE_TRAIN, cls.PARAM_EPOCH, 0) + job.progress.add_stage_param(cls.STAGE_TRAIN, cls.PARAM_TOTAL_EPOCHS, config.epochs) + + send_training_request(job=job, service=service, algorithm=algorithm, dataset=dataset) + + # A fast run reports its result before the request it came from returns, so the + # job can already be finished here. Writing this copy's progress would put the + # training stage back to "started" on a job that is done. + job.refresh_from_db() + if job.status in JobState.final_states(): + job.logger.info(f"{service.name} already reported its result.") + return + + job.progress.update_stage(cls.STAGE_DISPATCH, status=JobState.SUCCESS, progress=1) + job.progress.update_stage(cls.STAGE_TRAIN, status=JobState.STARTED, progress=0) + job.save(update_fields=["progress", "updated_at"]) + # Leaving the job STARTED is the point: a real training set outlasts the request + # that started it, and the result arrives at the job's callback. + job.logger.info(f"Training request accepted by {service.name}. Waiting for it to report back.") + + @classmethod + def record_progress(cls, job: "Job", epoch: int, total_epochs: int | None = None) -> None: + """ + Move the training stage to the epoch a service says it has reached. + + Training is the long stage and the service is silent while it runs, so without this + the stage sits at nought for as long as the run takes. The epoch count is the only + progress the service has to report. + """ + # As in record_result: a callback can arrive on a job whose stages were never set + # up, for instance after a restart. + if not any(stage.key == cls.STAGE_TRAIN for stage in job.progress.stages): + job.progress.add_stage("Training", cls.STAGE_TRAIN) + + try: + reached = job.progress.get_stage_param(cls.STAGE_TRAIN, job.progress.make_key(cls.PARAM_EPOCH)).value + except ValueError: + reached = 0 + # Pings can arrive late or out of order, and a stale one must not drag a run + # backwards past the epoch it actually reached. + if epoch <= (reached or 0): + return + + total = total_epochs + if not total: + try: + total = job.progress.get_stage_param( + cls.STAGE_TRAIN, job.progress.make_key(cls.PARAM_TOTAL_EPOCHS) + ).value + except ValueError: + total = None + + job.progress.add_or_update_stage_param(cls.STAGE_TRAIN, cls.PARAM_EPOCH, epoch) + if total: + job.progress.add_or_update_stage_param(cls.STAGE_TRAIN, cls.PARAM_TOTAL_EPOCHS, total) + job.progress.update_stage( + cls.STAGE_TRAIN, + status=JobState.STARTED, + # Never a full stage: the result finishes it, not the last epoch, since the + # service still has to score the head and upload it. + progress=min(epoch / total, 0.99) if total else 0, + ) + job.save(update_fields=["progress", "updated_at"]) + + @classmethod + def record_result( + cls, + job: "Job", + result: "TrainingResult", + dataset: dict | None = None, + dataset_url: str | None = None, + ) -> "Algorithm | None": + """ + Store what the service reported, register the new version, and finish the job. + + The result arrives one way only: through the callback, parsed at the view. A + service that also answers inline is acknowledged and its body ignored, so there is + a single shape to read here. + """ + job.refresh_from_db(fields=["status", "result"]) + if job.status in JobState.final_states(): + # A retry of the callback, or a second service reporting on the same job. The + # first answer stands; recording this one would register a second version. + job.logger.info("A training result is already recorded for this job; ignoring a duplicate.") + return None + + job.result = {"result": result.dict(), "dataset": dataset, "dataset_url": dataset_url} + + # A result can arrive through the callback on a job whose stages were never set up, + # for instance after a restart. Make sure the stage exists before reporting into it. + if not any(stage.key == cls.STAGE_TRAIN for stage in job.progress.stages): + job.progress.add_stage("Training", cls.STAGE_TRAIN) + # A service that is reporting back plainly received the request, and the callback + # can arrive before the dispatching code closes that stage. + if any(stage.key == cls.STAGE_DISPATCH for stage in job.progress.stages): + job.progress.update_stage(cls.STAGE_DISPATCH, status=JobState.SUCCESS, progress=1) + + for warning in result.warnings: + job.logger.warning(warning) + + job.progress.add_stage_param(cls.STAGE_TRAIN, "New head top-1", result.candidate_metrics.get("top1")) + incumbent = result.incumbent_metrics or {} + job.progress.add_stage_param(cls.STAGE_TRAIN, "Current head top-1", incumbent.get("top1")) + job.progress.add_stage_param(cls.STAGE_TRAIN, "Better", result.promote) + job.logger.info(result.reason or "Training finished.") + + new_version = cls.register_new_version(job=job, result=result, dataset=dataset, dataset_url=dataset_url) + if new_version: + job.progress.add_stage_param(cls.STAGE_TRAIN, "New version", new_version.key) + job.logger.info(f"Registered {new_version} as version {new_version.version}") + + job.progress.update_stage(cls.STAGE_TRAIN, status=JobState.SUCCESS, progress=1) + job.finished_at = datetime.datetime.now() + job.update_status(JobState.SUCCESS, save=True) + job.save() + return new_version + + @classmethod + def register_new_version( + cls, + job: "Job", + result: "TrainingResult", + dataset: dict | None = None, + dataset_url: str | None = None, + ): + """ + Record the trained head as a new version of the algorithm it was trained from. + + A Classification points at an Algorithm row, so a version that changed in place + would make past predictions untraceable. Each retrain therefore gets its own row: + same name, next version number, and a new key, because key is unique on its own. + """ + import pydantic + + from ami.ml.models import Algorithm, AlgorithmCategoryMap + from ami.ml.schemas import AlgorithmTrainingInfo, TrainingDatasetMetadata + + # The service echoes the training set's own metadata, which is where the run's + # provenance comes from. Read leniently: a service echoing an older shape should + # still have its result recorded, and the only cost is a version that cannot say + # which set it learned from. + metadata = None + if dataset: + try: + metadata = TrainingDatasetMetadata.parse_obj(dataset) + except pydantic.ValidationError as e: + job.logger.warning(f"The service returned training set metadata Antenna could not read: {e}") + + parent_key = (job.params or {}).get("algorithm_key") + parent = Algorithm.objects.filter(key=parent_key).first() + if not parent: + job.logger.warning(f"Cannot register a new version: no algorithm with key '{parent_key}'.") + return None + + labels = result.labels + if not labels: + job.logger.warning("The service returned no class list, so no new version was registered.") + return None + + trained_at = result.trained_at or "" + stamp = str(trained_at).replace(":", "").replace("-", "").replace(".", "")[:15] or f"job{job.pk}" + version = ( + Algorithm.objects.filter(name=parent.name).order_by("-version").values_list("version", flat=True).first() + or parent.version + ) + 1 + + category_map = AlgorithmCategoryMap.objects.create( + labels=labels, + data=[{"label": label, "index": i} for i, label in enumerate(labels)], + version=stamp, + description=( + f"Retrained from {metadata.rows if metadata else len(labels)} verified crops in " + f"project '{job.project.name}' by job #{job.pk}." + ), + ) + + return Algorithm.objects.create( + name=parent.name, + key=f"{parent.key}-v{version}-{stamp}", + # Where the weights are kept. Empty when the service did not upload them, in + # which case the head exists only on that service's own disk. + uri=result.head_url or "", + version=version, + version_name=str(trained_at) or stamp, + task_type=parent.task_type, + description=parent.description, + trainable=parent.trainable, + training_config=parent.training_config, + category_map=category_map, + training_info=AlgorithmTrainingInfo( + trained_at=trained_at or None, + dataset_url=dataset_url or (metadata.url if metadata else None), + occurrence_set_id=metadata.occurrence_set.id if metadata and metadata.occurrence_set else None, + occurrence_set_name=metadata.occurrence_set.name if metadata and metadata.occurrence_set else None, + dataset_rows=result.rows.get("kept"), + dataset_classes=len(labels), + metrics=result.candidate_metrics, + previous_metrics=result.incumbent_metrics or {}, + parent_algorithm_key=parent.key, + job_id=job.pk, + warnings=result.warnings, + ), + ) + + VALID_JOB_TYPES = [ MLJob, SourceImageCollectionPopulateJob, @@ -1005,6 +1432,7 @@ def run(cls, job: "Job"): UnknownJobType, DataExportJob, PostProcessingJob, + TrainClassifierJob, ] @@ -1041,7 +1469,8 @@ class Job(BaseModel): # Redis SREM-driven progress save, so this is effectively "no progress for # N minutes". 10 is conservative; raise if legitimate long-running jobs get # reaped. - STALLED_JOBS_MAX_MINUTES = 10 + # The default deadline, owned by JobType so a job type can set its own. + STALLED_JOBS_MAX_MINUTES = JobType.stalled_after_minutes # Zombie-stream reaper: age threshold above which a NATS stream for a job # in a terminal state (or missing from Django) is considered safe to drop. # Kept well above :attr:`STALLED_JOBS_MAX_MINUTES` so newly-dispatched jobs diff --git a/ami/jobs/serializers.py b/ami/jobs/serializers.py index 7fe1e2581..ac521f39e 100644 --- a/ami/jobs/serializers.py +++ b/ami/jobs/serializers.py @@ -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 ( @@ -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. diff --git a/ami/jobs/views.py b/ami/jobs/views.py index 09b314e4b..175ea70b4 100644 --- a/ami/jobs/views.py +++ b/ami/jobs/views.py @@ -7,6 +7,7 @@ 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 @@ -14,6 +15,8 @@ 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 @@ -28,6 +31,9 @@ MLJobResultsResponseSerializer, MLJobTasksRequestSerializer, MLJobTasksResponseSerializer, + TrainingProgressRequestSerializer, + TrainingResultRequestSerializer, + TrainingResultResponseSerializer, ) from ami.jobs.tasks import ( HEARTBEAT_THROTTLE_SECONDS, @@ -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}) diff --git a/ami/ml/schemas.py b/ami/ml/schemas.py index d1d00a575..88ff797af 100644 --- a/ami/ml/schemas.py +++ b/ami/ml/schemas.py @@ -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. diff --git a/ami/ml/test_training.py b/ami/ml/test_training.py index 896e7fc7a..b75eede40 100644 --- a/ami/ml/test_training.py +++ b/ami/ml/test_training.py @@ -217,15 +217,22 @@ def test_a_set_from_another_project_is_not_counted(self): self.assertEqual(response.status_code, 404) def test_the_split_ratio_moves_the_counts(self): + """ + Compared across two ratios rather than asserted exactly. + + Which side a given occurrence falls on is a hash of its id, so with a handful of + rows a particular ratio does not produce a particular count. + """ for taxon in self.taxa: self.make_occurrence(taxon) self.make_occurrence(taxon) - everything_held_out = self._summary(test_fraction=0.99) + mostly_train = self._summary(test_fraction=0.01) + mostly_test = self._summary(test_fraction=0.99) - self.assertEqual(everything_held_out["test"], everything_held_out["rows"]) - self.assertEqual(everything_held_out["train"], 0) - self.assertEqual(everything_held_out["settings"]["test_fraction"], 0.99) + self.assertEqual(mostly_test["rows"], mostly_train["rows"]) + self.assertGreater(mostly_test["test"], mostly_train["test"]) + self.assertEqual(mostly_test["settings"]["test_fraction"], 0.99) def test_a_ratio_no_run_could_use_is_refused(self): response = self.client.get( @@ -258,3 +265,254 @@ def test_someone_who_cannot_retrain_cannot_read_the_training_set(self): response = self.client.get(self.url, {"project_id": self.project.pk, "algorithm": self.algorithm.key}) self.assertEqual(response.status_code, 403) + + +class TestTheJobsOccurrenceSet(TrainingSetFixture): + def _job(self, **params): + from ami.jobs.models import Job, TrainClassifierJob + + return Job.objects.create( + project=self.project, + name="Retrain", + job_type_key=TrainClassifierJob.key, + params={"algorithm_key": self.algorithm.key, **params}, + ) + + def test_a_job_without_a_set_learns_from_everything_verified(self): + from ami.jobs.models import TrainClassifierJob + + self.assertIsNone(TrainClassifierJob.target_occurrence_set(self._job())) + + def test_a_set_from_another_project_is_refused(self): + from ami.jobs.models import TrainClassifierJob + + other = Project.objects.create(name="Someone else", owner=self.user) + theirs = OccurrenceSet.objects.create(name="Theirs") + theirs.projects.add(other) + + with self.assertRaises(ValueError) as caught: + TrainClassifierJob.target_occurrence_set(self._job(occurrence_set_id=theirs.pk)) + + self.assertIn("in this project", str(caught.exception)) + + def test_a_set_in_this_project_is_used(self): + from ami.jobs.models import TrainClassifierJob + + mine = self.make_set("Mine", [self.make_occurrence(self.taxa[0])]) + + self.assertEqual(TrainClassifierJob.target_occurrence_set(self._job(occurrence_set_id=mine.pk)), mine) + + +class TestTheTrainingStageFollowsTheEpochs(TrainingSetFixture): + """ + Training is the long stage and the service is silent while it fits, so the epochs it + reports are the only thing that can move the stage while a run is going. + """ + + def setUp(self) -> None: + super().setUp() + from ami.jobs.models import Job, TrainClassifierJob + + self.job = Job.objects.create( + project=self.project, + name="Retrain", + job_type_key=TrainClassifierJob.key, + params={"algorithm_key": self.algorithm.key}, + ) + self.job.progress.add_stage("Training", TrainClassifierJob.STAGE_TRAIN) + self.url = reverse("api:job-training-progress", args=[self.job.pk]) + + def _stage(self): + from ami.jobs.models import Job, TrainClassifierJob + + job = Job.objects.get(pk=self.job.pk) + return job.progress.get_stage(TrainClassifierJob.STAGE_TRAIN) + + def _param(self, name: str): + from ami.jobs.models import Job, TrainClassifierJob + + job = Job.objects.get(pk=self.job.pk) + return job.progress.get_stage_param(TrainClassifierJob.STAGE_TRAIN, job.progress.make_key(name)).value + + def _post(self, **payload): + from ami.ml.training import make_callback_token + + return APIClient().post( + self.url, + payload, + format="json", + HTTP_AUTHORIZATION=f"Token {make_callback_token(self.job)}", + ) + + def test_a_ping_moves_the_stage(self): + from ami.jobs.models import TrainClassifierJob + + response = self._post(epoch=150, total_epochs=300) + + self.assertEqual(response.status_code, 200, response.data) + self.assertEqual(self._param(TrainClassifierJob.PARAM_EPOCH), 150) + self.assertEqual(self._param(TrainClassifierJob.PARAM_TOTAL_EPOCHS), 300) + self.assertEqual(self._stage().progress, 0.5) + + def test_the_stage_is_never_finished_by_an_epoch(self): + """The result finishes it: the service still has to score the head and upload it.""" + self._post(epoch=300, total_epochs=300) + + self.assertEqual(self._stage().progress, 0.99) + + def test_a_ping_that_arrives_late_is_ignored(self): + self._post(epoch=200, total_epochs=300) + self._post(epoch=100, total_epochs=300) + + from ami.jobs.models import TrainClassifierJob + + self.assertEqual(self._param(TrainClassifierJob.PARAM_EPOCH), 200) + + def test_a_ping_without_the_token_is_refused(self): + response = APIClient().post(self.url, {"epoch": 10}, format="json") + + self.assertEqual(response.status_code, 403) + + def test_a_ping_after_the_job_finished_changes_nothing(self): + from ami.jobs.models import JobState, TrainClassifierJob + + self.job.status = JobState.SUCCESS + self.job.save() + + response = self._post(epoch=10, total_epochs=300) + + self.assertEqual(response.status_code, 200) + with self.assertRaises(ValueError): + self._param(TrainClassifierJob.PARAM_EPOCH) + + +class TestTheResultComingBackFromAService(TrainingSetFixture): + """ + The callback is the one way a result arrives, so its shape is the contract. + + A service posts the result together with the training set's own metadata, echoed back + from the file Antenna wrote. That echo is where a registered version learns which + occurrence set it came from. + """ + + def setUp(self) -> None: + super().setUp() + from ami.jobs.models import Job, TrainClassifierJob + + self.occurrence_set = self.make_set("Mine", [self.make_occurrence(self.taxa[0])]) + self.job = Job.objects.create( + project=self.project, + name="Retrain", + job_type_key=TrainClassifierJob.key, + params={"algorithm_key": self.algorithm.key, "occurrence_set_id": self.occurrence_set.pk}, + ) + self.url = reverse("api:job-training-result", args=[self.job.pk]) + + def _dataset_echo(self, **overrides) -> dict: + """The metadata a service reads out of the training set and sends back.""" + from ami.ml.schemas import ( + NamedReference, + TrainedAlgorithmReference, + TrainingDatasetMetadata, + TrainingDatasetSettings, + ) + + metadata = TrainingDatasetMetadata( + url="/media/training/set.npz", + project=NamedReference(id=self.project.pk, name=self.project.name), + algorithm=TrainedAlgorithmReference(key=self.algorithm.key, name=self.algorithm.name, version=1), + dimensions=VECTOR_LENGTH, + dtype="float32", + classes=[t.name for t in self.taxa], + rows=6, + train=5, + test=1, + occurrence_set=NamedReference(id=self.occurrence_set.pk, name=self.occurrence_set.name), + settings=TrainingDatasetSettings(min_per_species=2, split_salt="antenna-head-v1", test_fraction=0.2), + ) + return {**metadata.dict(), **overrides} + + def _post(self, payload: dict): + from ami.ml.training import make_callback_token + + return APIClient().post( + self.url, + payload, + format="json", + HTTP_AUTHORIZATION=f"Token {make_callback_token(self.job)}", + ) + + def _result(self, **overrides) -> dict: + return { + "promote": False, + "reason": "The new head did not beat the current one.", + "warnings": ["Only 1 held-out row(s)."], + "labels": [t.name for t in self.taxa], + "rows": {"kept": 6, "dropped": 0}, + "candidate_metrics": {"top1": 0.8}, + "incumbent_metrics": {"top1": 0.9}, + "trained_at": "2026-10-08T23:00:00", + "head_url": "/media/algorithms/head.npz", + **overrides, + } + + def test_a_result_registers_a_version_that_says_what_it_learned_from(self): + from ami.jobs.models import Job + from ami.ml.models import Algorithm + + response = self._post({"result": self._result(), "dataset": self._dataset_echo()}) + + self.assertEqual(response.status_code, 200, response.data) + version = Algorithm.objects.get(key=response.data["algorithm"]) + self.assertEqual(version.training_info.occurrence_set_id, self.occurrence_set.pk) + self.assertEqual(version.training_info.occurrence_set_name, "Mine") + self.assertEqual(version.training_info.dataset_rows, 6) + self.assertEqual(version.training_info.metrics, {"top1": 0.8}) + self.assertEqual(version.training_info.previous_metrics, {"top1": 0.9}) + self.assertEqual(Job.objects.get(pk=self.job.pk).status, "SUCCESS") + + def test_an_echo_antenna_cannot_read_still_records_the_result(self): + """ + A service echoing an older shape loses the provenance, not the run. + + Refusing the whole callback would throw away a head that was trained successfully. + """ + from ami.ml.models import Algorithm + + response = self._post({"result": self._result(), "dataset": {"unexpected": "shape"}}) + + self.assertEqual(response.status_code, 200, response.data) + version = Algorithm.objects.get(key=response.data["algorithm"]) + self.assertIsNone(version.training_info.occurrence_set_id) + + def test_a_run_with_no_current_head_to_compare_against_is_recorded(self): + """ + A service reports a null incumbent rather than leaving the field out. + + Refusing that would throw away the result of a first-ever retrain, which is the + one run guaranteed to have nothing to compare against. + """ + from ami.ml.models import Algorithm + + response = self._post({"result": self._result(incumbent_metrics=None), "dataset": self._dataset_echo()}) + + self.assertEqual(response.status_code, 200, response.data) + version = Algorithm.objects.get(key=response.data["algorithm"]) + self.assertEqual(version.training_info.previous_metrics, {}) + + def test_a_result_antenna_cannot_read_is_refused(self): + """The result is what the new version is built from, so a broken one is a 400.""" + response = self._post({"result": {"labels": "not a list"}, "dataset": self._dataset_echo()}) + + self.assertEqual(response.status_code, 400) + + def test_a_second_callback_does_not_register_a_second_version(self): + from ami.ml.models import Algorithm + + self._post({"result": self._result(), "dataset": self._dataset_echo()}) + before = Algorithm.objects.count() + again = self._post({"result": self._result(), "dataset": self._dataset_echo()}) + + self.assertEqual(again.status_code, 200) + self.assertEqual(again.data["status"], "already recorded") + self.assertEqual(Algorithm.objects.count(), before) diff --git a/ami/ml/training/__init__.py b/ami/ml/training/__init__.py index 6770457ad..91bab508e 100644 --- a/ami/ml/training/__init__.py +++ b/ami/ml/training/__init__.py @@ -1,9 +1,9 @@ """ Retraining a classifier head from the species people have verified. -``dataset`` decides what the head learns from and writes it to storage. The names -re-exported here are the ones callers outside this package use; anything else is internal -to it. +``dataset`` decides what the head learns from and writes it to storage; ``service`` hands a +processing service the URL and keeps the head that comes back. The names re-exported here +are the ones callers outside this package use; anything else is internal to it. """ from ami.ml.training.dataset import ( @@ -22,20 +22,48 @@ verified_occurrence_ids, verified_training_rows, ) +from ami.ml.training.service import ( + CALLBACK_MAX_AGE_SECONDS, + DISPATCH_TIMEOUT_SECONDS, + MAX_HEAD_BYTES, + HeadTooLarge, + absolute_media_url, + callback_url_for, + head_path, + head_upload_url_for, + make_callback_token, + progress_url_for, + send_training_request, + store_head, + verify_callback_token, +) __all__ = [ + "CALLBACK_MAX_AGE_SECONDS", "DEFAULT_SPLIT_SALT", "DEFAULT_TEST_FRACTION", + "DISPATCH_TIMEOUT_SECONDS", + "MAX_HEAD_BYTES", "SPLITS", "SPLIT_TEST", "SPLIT_TRAIN", + "HeadTooLarge", "NotEnoughVerifiedData", + "absolute_media_url", "build_training_dataset", + "callback_url_for", "count_missing_embeddings", + "head_path", + "head_upload_url_for", "label_counts", + "make_callback_token", + "progress_url_for", "row_as_dict", + "send_training_request", "species_with_enough_examples", "split_for", + "store_head", "verified_occurrence_ids", "verified_training_rows", + "verify_callback_token", ] diff --git a/ami/ml/training/service.py b/ami/ml/training/service.py new file mode 100644 index 000000000..f684708f6 --- /dev/null +++ b/ami/ml/training/service.py @@ -0,0 +1,200 @@ +""" +Hand a training request to a processing service, and keep what it sends back. + +Outward: a URL to the dataset rather than the rows (see ``dataset``), plus a signed token +the service reports back with. + +Inward: the service uploads the finished head through Antenna rather than to storage. It +holds no storage credentials, and a presigned upload would only work on S3, not against the +local filesystem. A head is one small matrix, so passing it through costs little. +""" + +import logging +import typing +from typing import TYPE_CHECKING +from urllib.parse import urljoin + +from django.conf import settings +from django.core import signing +from django.core.files.storage import default_storage +from django.urls import reverse +from django.utils.text import slugify + +from ami.ml.schemas import AlgorithmTrainingConfig, TrainingRequest +from ami.utils.requests import create_session, extract_error_message_from_response + +if TYPE_CHECKING: + from ami.jobs.models import Job + from ami.ml.models import Algorithm, ProcessingService + +logger = logging.getLogger(__name__) + + +# How long to wait for the service to acknowledge the request. This is not how long +# training takes: a service that answers within this window returns its result inline, +# and one that does not is expected to report back through the job's result endpoint. +DISPATCH_TIMEOUT_SECONDS = 600 + + +# A processing service has no Antenna account, so the callback is authorised by a signed +# token instead. Nothing is stored: the signature carries the job id and Django's secret +# key proves Antenna issued it. +CALLBACK_SALT = "ami.ml.training.callback" +CALLBACK_MAX_AGE_SECONDS = 60 * 60 * 24 + + +def make_callback_token(job: "Job") -> str: + """A token only Antenna could have produced, tied to this one job.""" + return signing.dumps({"job_id": job.pk}, salt=CALLBACK_SALT) + + +def verify_callback_token(token: str, job: "Job") -> bool: + """True when the token is Antenna's, unexpired, and for this job.""" + if not token: + return False + try: + payload = signing.loads(token, salt=CALLBACK_SALT, max_age=CALLBACK_MAX_AGE_SECONDS) + except signing.BadSignature: + return False + return payload.get("job_id") == job.pk + + +def _callback_base(job: "Job") -> str: + """ + Where the service can reach this Antenna. + + Deliberately not taken from the job's params. Params are free JSON that any member who + can create a job may set, and this base decides where the signed callback token is + sent -- a token that also unlocks the head upload. A per-deployment override belongs + in settings, not in a request body. + """ + base = getattr(settings, "EXTERNAL_BASE_URL", "") + if not base: + raise ValueError( + "No base URL is configured, so the processing service has no way to report back. " "Set EXTERNAL_BASE_URL." + ) + return base + + +def callback_url_for(job: "Job") -> str: + """Where the service should post its result when training finishes.""" + # reverse(), so the path follows the router rather than a copy of it here. + path = reverse("api:job-training-result", args=[job.pk]) + return urljoin(_callback_base(job).rstrip("/") + "/", path.lstrip("/")) + + +def head_upload_url_for(job: "Job") -> str: + """Where the service should upload the head it produced, so Antenna keeps a copy.""" + path = reverse("api:job-training-head", args=[job.pk]) + return urljoin(_callback_base(job).rstrip("/") + "/", path.lstrip("/")) + + +def progress_url_for(job: "Job") -> str: + """Where the service should report how far through the epochs it is.""" + path = reverse("api:job-training-progress", args=[job.pk]) + return urljoin(_callback_base(job).rstrip("/") + "/", path.lstrip("/")) + + +def absolute_media_url(url: str, base_url: str | None = None) -> str: + """ + Turn a stored file's URL into one a processing service can fetch. + + In production MEDIA_URL is already an absolute S3 URL and this is a no-op. Locally it is + the relative /media/ path, so it needs a base in front. + """ + if url.startswith("http://") or url.startswith("https://"): + return url + base = base_url or getattr(settings, "EXTERNAL_BASE_URL", "") + if not base: + raise ValueError( + "The training set is stored at a relative URL and no base URL is configured, so the " + "processing service has no way to download it. Set EXTERNAL_BASE_URL." + ) + return urljoin(base.rstrip("/") + "/", url.lstrip("/")) + + +def send_training_request(job: "Job", service: "ProcessingService", algorithm: "Algorithm", dataset: dict) -> None: + """ + Ask a processing service to retrain a head. + + Returns once the service has accepted the work. The result is not read from the + response: it arrives at the job's result callback, which is the only way a run that + outlasts its request can report. A service that answers inline as well is simply + acknowledged. Raises if the service refused the request. + """ + endpoint = urljoin(service.endpoint_url.rstrip("/") + "/", "train") + # The same merge the job validated before it built the dataset, so the service is told + # the settings the training set was actually made under. Reading job.params again here + # would let the two drift, and would send values nothing had checked. + config = AlgorithmTrainingConfig.for_run(algorithm.training_config, job.params or {}) + request = TrainingRequest( + dataset_url=absolute_media_url(dataset["url"]), + algorithm_key=algorithm.key, + job_id=job.pk, + name=f"{algorithm.key}-job-{job.pk}", + min_per_species=config.min_per_species, + # The fitting settings the service published, so an admin can tune them in Antenna + # without redeploying the service. + min_improvement=config.min_improvement, + head_type=config.head_type, + epochs=config.epochs, + learning_rate=config.learning_rate, + weight_decay=config.weight_decay, + # Always sent: a service that finishes after the request times out reports back + # here instead, which is the only way a real training set can work. + callback_url=callback_url_for(job), + callback_token=make_callback_token(job), + # The head itself comes back here. A service that does not support the upload + # ignores this, and the version is registered without a stored copy. + head_upload_url=head_upload_url_for(job), + # Without this the training stage simply stays at nought until the result lands. + progress_url=progress_url_for(job), + ) + + job.logger.info(f"Sending training request to {endpoint} for {algorithm.key}") + session = create_session() + response = session.post(endpoint, json=request.dict(), timeout=DISPATCH_TIMEOUT_SECONDS) + + if not response.ok: + message = extract_error_message_from_response(response) + raise ValueError(f"The processing service refused the training request: {message}") + + +HEAD_DIRECTORY = "algorithms" + +# A head is one Linear layer: 72 KB over 17 species, about 3 MB over 749. The cap is well +# clear of that and only exists so a wrong or hostile upload cannot fill the bucket. +MAX_HEAD_BYTES = 64 * 1024 * 1024 + + +class HeadTooLarge(Exception): + """The upload is larger than any head this could reasonably be.""" + + +def head_path(algorithm_key: str, job_id: int, filename: str) -> str: + """Where one job's head is kept. Deterministic, so a re-run replaces its own file.""" + return f"{HEAD_DIRECTORY}/{slugify(algorithm_key)}-job-{job_id}/{filename}" + + +def store_head(algorithm_key: str, job_id: int, files: dict) -> dict[str, typing.Any]: + """ + Save the uploaded head files and return where they landed. + + Returns the storage path and URL of each file, keyed by the name it was uploaded + under, so the caller can record the head's location against the algorithm version. + """ + stored = {} + for name, uploaded in files.items(): + if uploaded.size > MAX_HEAD_BYTES: + raise HeadTooLarge(f"'{name}' is {uploaded.size} bytes; the limit is {MAX_HEAD_BYTES}.") + + path = head_path(algorithm_key, job_id, uploaded.name) + if default_storage.exists(path): + # A re-run of the same job replaces its head instead of piling up copies, + # matching how the training set is written. + default_storage.delete(path) + saved_path = default_storage.save(path, uploaded) + stored[name] = {"path": saved_path, "url": default_storage.url(saved_path)} + logger.info(f"Stored '{name}' for job {job_id} at {saved_path}") + + return stored