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 == {}