AI_Development_Automation_Manager / tests /test_experiment_tracker.py
SyntheticMDProductions's picture
Update ADAM safety, UI, and model workflows (#1)
c61c435
Raw
History Blame Contribute Delete
3.86 kB
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from pathlib import Path
import sqlite3
from adam.experiment_tracker import ExperimentStore
from adam.models import ExecutionPlan, Job, JobStatus, PlanStep, SystemSnapshot
def test_experiment_store_records_training_job_and_clone_request(tmp_path: Path) -> None:
dataset = tmp_path / "dataset"
dataset.mkdir()
for index in range(3):
(dataset / f"{index}.png").write_bytes(b"image")
output = tmp_path / "output"
output.mkdir()
checkpoint = output / "model.safetensors"
checkpoint.write_bytes(b"weights")
now = datetime.now(timezone.utc)
job = Job(
id="ABC123",
plan=ExecutionPlan(
request="train",
summary="Train",
steps=[
PlanStep(
"ddpm_trainer",
"Train DDPM",
"Train",
{
"dataset_dir": str(dataset),
"model_name": "Demo Model",
"epochs": 12,
"output_dir": str(output),
"resolution": 128,
"batch_size": 2,
"learning_rate": 0.0001,
"preview_seed": 44,
},
)
],
),
status=JobStatus.FINISHED,
started_at=(now - timedelta(minutes=5)).isoformat(),
ended_at=now.isoformat(),
output_folder=str(output),
logs=["loss: 0.25"],
)
store = ExperimentStore(tmp_path)
run = store.record_job(job, SystemSnapshot(gpu_name="Test GPU", vram_used_gb=4, vram_total_gb=8))
assert run is not None
assert store.list_runs()[0].dataset_item_count == 3
assert store.list_runs()[0].final_loss == 0.25
request = store.clone_request("EXP-ABC123")
assert "Demo Model Clone" in request
assert "dataset_dir" not in request
def test_experiment_store_updates_notes_and_compare(tmp_path: Path) -> None:
store = ExperimentStore(tmp_path)
for suffix in ("A", "B"):
job = Job(
id=f"JOB{suffix}",
plan=ExecutionPlan(
request="train",
summary="Train",
steps=[
PlanStep(
"flow_trainer",
"Train",
"Train",
{"dataset_dir": str(tmp_path), "model_name": suffix, "epochs": 5, "output_dir": str(tmp_path)},
)
],
),
status=JobStatus.FINISHED,
ended_at=datetime.now(timezone.utc).isoformat(),
)
store.record_job(job)
store.update_notes("EXP-JOBA", "best so far", 89)
assert store.get("EXP-JOBA").notes == "best so far" # type: ignore[union-attr]
comparison = store.compare(["EXP-JOBA", "EXP-JOBB"])
assert any(row["field"] == "quality_score" and row["EXP-JOBA"] == 89 for row in comparison)
def test_experiment_store_migrates_older_sqlite_schema(tmp_path: Path) -> None:
path = tmp_path / "data" / "experiments.sqlite3"
path.parent.mkdir()
with sqlite3.connect(path) as db:
db.execute(
"CREATE TABLE experiments (id TEXT PRIMARY KEY, job_id TEXT UNIQUE NOT NULL, timestamp TEXT NOT NULL, model_architecture TEXT NOT NULL, model_name TEXT NOT NULL)"
)
db.execute(
"INSERT INTO experiments (id, job_id, timestamp, model_architecture, model_name) VALUES ('EXP-OLD', 'OLD', '2026-01-01T00:00:00+00:00', 'ddpm', 'Old Run')"
)
store = ExperimentStore(tmp_path)
run = store.get("EXP-OLD")
assert run is not None
assert run.dataset_path == ""
assert run.checkpoint_paths == []
assert run.settings == {}