124 lines
7.7 KiB
Python
124 lines
7.7 KiB
Python
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
from datetime import timedelta
|
||
|
|
import json
|
||
|
|
|
||
|
|
import pytest
|
||
|
|
from django.utils import timezone
|
||
|
|
|
||
|
|
from control_plane.model_studio.backends import FakeTrainingBackend
|
||
|
|
from control_plane.model_studio.models import (
|
||
|
|
CheckpointType, CheckpointValidityStatus, Dataset, DatasetVersion, DatasetValidationStatus,
|
||
|
|
EvaluationSuite, EvaluationSuiteVersion, ExperimentStatus, ModelCheckpoint, ModelPromotionPolicy,
|
||
|
|
PromotionDecision, TrainingProject, TrainingProjectStatus,
|
||
|
|
)
|
||
|
|
from control_plane.model_studio.services import ModelStudioService
|
||
|
|
from control_plane.projects.models import Project
|
||
|
|
|
||
|
|
|
||
|
|
def studio(outcomes=None):
|
||
|
|
return ModelStudioService(backend=FakeTrainingBackend(outcomes))
|
||
|
|
|
||
|
|
|
||
|
|
def ready_project():
|
||
|
|
project = Project.objects.create(name="Guard", project_type="MODEL", goal="test")
|
||
|
|
training_project = TrainingProject.objects.create(project=project, name="Guard", slug=f"guard-{project.id.hex[:8]}", goal="test", status=TrainingProjectStatus.READY)
|
||
|
|
dataset = Dataset.objects.create(training_project=training_project, name="train")
|
||
|
|
version = DatasetVersion.objects.create(dataset=dataset, version="v1", manifest_reference="fake://train", content_hash="a" * 64, record_count=2, validation_status=DatasetValidationStatus.VALID)
|
||
|
|
checkpoint = ModelCheckpoint.objects.create(training_project=training_project, name="champion", checkpoint_type=CheckpointType.CHAMPION, reference="fake://champion", content_hash="b" * 64, validity_status=CheckpointValidityStatus.VALID, load_verified=True)
|
||
|
|
suite = EvaluationSuite.objects.create(training_project=training_project, name="suite")
|
||
|
|
suite_version = EvaluationSuiteVersion.objects.create(suite=suite, version="v1", reference="fake://suite", content_hash="c" * 64, integrity_status=DatasetValidationStatus.VALID, immutable=True)
|
||
|
|
training_project.current_champion = checkpoint
|
||
|
|
training_project.save(update_fields=["current_champion", "updated_at"])
|
||
|
|
service = studio()
|
||
|
|
baseline = service.evaluate(training_project, checkpoint, suite_version, metrics={"primary": 0.5, "critical": 0.5}, synthetic=True)
|
||
|
|
training_project.baseline_evaluation = baseline
|
||
|
|
training_project.save(update_fields=["baseline_evaluation", "updated_at"])
|
||
|
|
policy = ModelPromotionPolicy.objects.create(training_project=training_project, name="test", version="v1", criteria={"primary_metric": "primary", "minimum_delta": 0.1, "max_regression": {"critical": 0.02}})
|
||
|
|
return service, training_project, version, suite_version, policy
|
||
|
|
|
||
|
|
|
||
|
|
def contract(**overrides):
|
||
|
|
payload = {"hypothesis": "Targeted replay improves primary metric.", "reasoning": "A baseline failure cluster supports this.", "intervention": {"learning_rate": 0.0001}, "controls": {"same_champion": True}, "expected_result": "primary +0.1", "primary_success_metric": "primary", "success_threshold": {"minimum_delta": 0.1}, "regression_constraints": {"critical": 0.02}, "rejection_condition": "No meaningful improvement", "ambiguity_policy": "replicate", "maximum_runtime_seconds": 120, "compute_budget": {"gpu_seconds": 120}, "expected_improvement": 0.2, "expected_information_gain": 0.8, "estimated_compute_cost": 1, "novelty": 1, "confidence": 0.8, "estimated_runtime_seconds": 60}
|
||
|
|
return {**payload, **overrides}
|
||
|
|
|
||
|
|
|
||
|
|
def test_recipe_and_completed_dataset_version_are_immutable():
|
||
|
|
service, training_project, dataset_version, _, _ = ready_project()
|
||
|
|
experiment = service.propose_experiment(training_project, contract())
|
||
|
|
run = service.run_experiment(experiment, {"dataset_references": ["fake://train"]})
|
||
|
|
assert run.recipe.immutable is True
|
||
|
|
dataset_version.refresh_from_db()
|
||
|
|
assert dataset_version.immutable is True
|
||
|
|
run.recipe.configuration = {"changed": True}
|
||
|
|
with pytest.raises(ValueError, match="TrainingRecipe is immutable"):
|
||
|
|
run.recipe.save()
|
||
|
|
dataset_version.record_count = 3
|
||
|
|
with pytest.raises(ValueError, match="DatasetVersion is immutable"):
|
||
|
|
dataset_version.save()
|
||
|
|
|
||
|
|
|
||
|
|
def test_incomplete_scientific_contract_and_duplicate_experiment_are_rejected():
|
||
|
|
service, training_project, _, _, _ = ready_project()
|
||
|
|
with pytest.raises(ValueError, match="Incomplete scientific contract"):
|
||
|
|
service.propose_experiment(training_project, {"hypothesis": "thin"})
|
||
|
|
service.propose_experiment(training_project, contract())
|
||
|
|
with pytest.raises(ValueError, match="Duplicate experiment"):
|
||
|
|
service.propose_experiment(training_project, contract())
|
||
|
|
|
||
|
|
|
||
|
|
def test_failed_execution_is_not_hypothesis_rejection():
|
||
|
|
service, training_project, _, _, _ = ready_project()
|
||
|
|
service.backend = FakeTrainingBackend([{"status": "OOM", "failure_category": "OOM", "failure_details": "simulated"}])
|
||
|
|
experiment = service.propose_experiment(training_project, contract())
|
||
|
|
run = service.run_experiment(experiment, {})
|
||
|
|
experiment.refresh_from_db()
|
||
|
|
assert run.status == "OOM"
|
||
|
|
assert experiment.status == ExperimentStatus.FAILED
|
||
|
|
assert experiment.conclusion == "EXECUTION_FAILED"
|
||
|
|
|
||
|
|
|
||
|
|
def test_challenger_promotion_requires_same_evaluation_version_and_objective_gate():
|
||
|
|
service, training_project, _, suite_version, policy = ready_project()
|
||
|
|
service.backend = FakeTrainingBackend([{"status": "SUCCEEDED", "checkpoint_reference": "fake://challenger"}])
|
||
|
|
experiment = service.propose_experiment(training_project, contract())
|
||
|
|
run = service.run_experiment(experiment, {})
|
||
|
|
evaluation = service.evaluate(training_project, run.output_checkpoint, suite_version, metrics={"primary": 0.7, "critical": 0.5}, synthetic=True)
|
||
|
|
decision = service.decide_promotion(experiment, evaluation, policy)
|
||
|
|
training_project.refresh_from_db()
|
||
|
|
assert decision.decision == PromotionDecision.PROMOTE
|
||
|
|
assert training_project.current_champion_id == run.output_checkpoint_id
|
||
|
|
|
||
|
|
|
||
|
|
def test_regression_prevents_promotion_and_deadline_reserve_blocks_start():
|
||
|
|
service, training_project, _, suite_version, policy = ready_project()
|
||
|
|
service.backend = FakeTrainingBackend([{"status": "SUCCEEDED", "checkpoint_reference": "fake://challenger"}])
|
||
|
|
experiment = service.propose_experiment(training_project, contract())
|
||
|
|
run = service.run_experiment(experiment, {})
|
||
|
|
evaluation = service.evaluate(training_project, run.output_checkpoint, suite_version, metrics={"primary": 0.8, "critical": 0.4}, synthetic=True)
|
||
|
|
decision = service.decide_promotion(experiment, evaluation, policy)
|
||
|
|
assert decision.decision == PromotionDecision.REJECT
|
||
|
|
program = service.create_program(training_project, policy, wall_seconds=60, evaluation_reserve=30)
|
||
|
|
program.deadline = timezone.now() + timedelta(seconds=60)
|
||
|
|
program.save(update_fields=["deadline", "updated_at"])
|
||
|
|
allowed, reason = service.can_start(program, experiment, estimated_evaluation_seconds=30)
|
||
|
|
assert allowed is False
|
||
|
|
assert reason == "FINAL_EVALUATION_RESERVE"
|
||
|
|
|
||
|
|
|
||
|
|
def test_dataset_curation_respects_explicit_no_evmbench_claim(tmp_path):
|
||
|
|
service, training_project, _, _, _ = ready_project()
|
||
|
|
manifest = tmp_path / "verified_pilot.json"
|
||
|
|
manifest.write_text(json.dumps([{"input": "pragma solidity ^0.8.0;", "output": "{}"}]), encoding="utf-8")
|
||
|
|
summary = tmp_path / "verified_pilot_summary.json"
|
||
|
|
summary.write_text(json.dumps({"evmbench_source_included": False, "weak_label_data_included": False}), encoding="utf-8")
|
||
|
|
dataset = Dataset.objects.create(training_project=training_project, name="verified-pilot")
|
||
|
|
version = DatasetVersion.objects.create(dataset=dataset, version="pilot", manifest_reference=str(manifest), content_hash="d" * 64)
|
||
|
|
|
||
|
|
report = service.curate_guard_datasets(training_project)
|
||
|
|
|
||
|
|
version.refresh_from_db()
|
||
|
|
assert version.validation_status == DatasetValidationStatus.WARNING
|
||
|
|
assert version.contamination_status == DatasetValidationStatus.WARNING
|
||
|
|
assert report["valid"] == 0
|