music3lab / tests /test_autonomous_round.py
coolpoodle's picture
code and training scripts
90884df verified
Raw
History Blame Contribute Delete
7.1 kB
"""CPU-only expected-red contract for autonomous Flow-encoder rounds.
Expected public API (specified before implementation)::
from music3lab.autonomous import (
AutonomousController, Candidate, EvaluationSuite, PromotionPolicy,
Registry,
)
This contract deliberately covers only continuous Flow-encoder waveform
reconstruction. It neither asserts nor implies Music 3 token extraction,
composition, style transfer, or native-token generation.
"""
from __future__ import annotations
from pathlib import Path
import pytest
from music3lab.autonomous import (
AutonomousController,
Candidate,
EvaluationSuite,
PromotionPolicy,
Registry,
)
USER22 = tuple(f"user22-{index:02d}" for index in range(22))
def _suite() -> EvaluationSuite:
return EvaluationSuite.freeze(
train=tuple(f"train-{index:02d}" for index in range(8)),
validation=tuple(f"validation-{index:02d}" for index in range(6)),
public=tuple(f"public-{index:02d}" for index in range(6)),
sealed=tuple(f"sealed-{index:02d}" for index in range(6)),
excluded_sources=USER22,
)
def _controller(tmp_path: Path) -> tuple[AutonomousController, Registry]:
registry = Registry(tmp_path / "registry")
registry.register_champion(
Candidate(name="base", artifact_sha256="a" * 64, kind="flow_encoder")
)
policy = PromotionPolicy(
bootstrap_samples=200,
seed=41,
protected_regression_limit=0.0,
teacher_regression_limit=0.05,
)
return AutonomousController(registry=registry, policy=policy), registry
def _metrics(*, train: float, validation: float, public: float, sealed: float, teacher: float = 1.0, technical: float = 1.0, protected: float = 1.0) -> dict[str, tuple[float, ...]]:
return {
"train": (train,) * 8,
"validation": (validation,) * 6,
"public": (public,) * 6,
"sealed": (sealed,) * 6,
"teacher": (teacher,) * 6,
"technical": (technical,) * 6,
"protected": (protected,) * 6,
}
def test_frozen_source_exclusive_suite_excludes_user22_and_is_immutable() -> None:
suite = _suite()
assert suite.splits == ("train", "validation", "public", "sealed")
assert not set(USER22) & set().union(*suite.source_ids_by_split.values())
assert len(set().union(*suite.source_ids_by_split.values())) == 26
with pytest.raises((AttributeError, TypeError, ValueError)):
suite.source_ids_by_split["train"] += ("leak",)
def test_challenger_selection_uses_validation_then_independent_public_and_sealed_metrics(tmp_path: Path) -> None:
controller, registry = _controller(tmp_path)
suite = _suite()
baseline = _metrics(train=1.0, validation=1.0, public=1.0, sealed=1.0)
winner = Candidate(name="real-latent-calibration", artifact_sha256="b" * 64, kind="flow_encoder")
weaker = Candidate(name="weaker-validation", artifact_sha256="c" * 64, kind="flow_encoder")
report = controller.run_round(
suite=suite,
baseline=baseline,
challengers={
winner: _metrics(train=0.80, validation=0.60, public=0.70, sealed=0.70),
weaker: _metrics(train=0.10, validation=0.70, public=0.01, sealed=0.01),
},
)
assert report.validation_winner == winner.artifact_sha256
assert report.decision == "PROMOTE"
assert registry.champion().artifact_sha256 == winner.artifact_sha256
assert report.selection_split == "validation"
assert report.independent_splits == ("public", "sealed")
def test_zero_latent_challenger_is_automatically_rejected(tmp_path: Path) -> None:
controller, registry = _controller(tmp_path)
baseline_hash = registry.champion().artifact_sha256
zero = Candidate(
name="zero-latent",
artifact_sha256="d" * 64,
kind="flow_encoder",
metadata={"latent_strategy": "zero"},
)
report = controller.run_round(
suite=_suite(),
baseline=_metrics(train=1.0, validation=1.0, public=1.0, sealed=1.0),
challengers={zero: _metrics(train=1.2, validation=1.4, public=1.5, sealed=1.5)},
)
assert report.decision == "REJECT"
assert "zero-latent" in report.rejected[zero.artifact_sha256].reasons
assert registry.champion().artifact_sha256 == baseline_hash
def test_train_lookup_overfit_rolls_back_without_champion_pointer_or_hash_change(tmp_path: Path) -> None:
controller, registry = _controller(tmp_path)
baseline = registry.champion()
lookup = Candidate(name="train-lookup", artifact_sha256="e" * 64, kind="flow_encoder")
report = controller.run_round(
suite=_suite(),
baseline=_metrics(train=1.0, validation=1.0, public=1.0, sealed=1.0),
challengers={lookup: _metrics(train=0.0, validation=0.60, public=1.20, sealed=1.30)},
)
assert report.decision == "ROLLBACK"
assert report.rejected[lookup.artifact_sha256].rollback is True
assert registry.champion().name == baseline.name
assert registry.champion().artifact_sha256 == baseline.artifact_sha256
def test_promotion_requires_positive_bootstrap_ci_and_all_non_regression_gates(tmp_path: Path) -> None:
controller, registry = _controller(tmp_path)
candidate = Candidate(name="candidate", artifact_sha256="f" * 64, kind="flow_encoder")
baseline = _metrics(train=1.0, validation=1.0, public=1.0, sealed=1.0)
rejected = controller.run_round(
suite=_suite(), baseline=baseline,
challengers={candidate: _metrics(train=0.7, validation=0.7, public=0.7, sealed=0.7, teacher=1.2)},
)
assert rejected.decision == "REJECT"
assert "teacher" in rejected.rejected[candidate.artifact_sha256].reasons
accepted = controller.run_round(
suite=_suite(), baseline=baseline,
challengers={candidate: _metrics(train=0.7, validation=0.7, public=0.7, sealed=0.7, teacher=1.01, technical=0.99, protected=1.0)},
)
assert accepted.decision == "PROMOTE"
assert all(interval[0] > 0 for interval in accepted.bootstrap_ci.values())
assert registry.champion().artifact_sha256 == candidate.artifact_sha256
def test_registry_contains_complete_base_champion_challenger_and_rejected_audit_rows(tmp_path: Path) -> None:
controller, registry = _controller(tmp_path)
good = Candidate(name="good", artifact_sha256="1" * 64, kind="flow_encoder")
bad = Candidate(name="bad", artifact_sha256="2" * 64, kind="flow_encoder")
report = controller.run_round(
suite=_suite(), baseline=_metrics(train=1.0, validation=1.0, public=1.0, sealed=1.0),
challengers={
good: _metrics(train=0.7, validation=0.7, public=0.7, sealed=0.7),
bad: _metrics(train=0.1, validation=0.8, public=1.2, sealed=1.2),
},
)
rows = registry.audit_rows(round_id=report.round_id)
assert {row.role for row in rows} == {"base", "champion", "challenger", "rejected"}
assert {row.artifact_sha256 for row in rows} >= {"a" * 64, good.artifact_sha256, bad.artifact_sha256}
assert all(row.suite_sha256 == _suite().sha256 for row in rows)