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