from __future__ import annotations import json import tempfile import unittest from pathlib import Path from unittest.mock import Mock, patch from airfrans_frontier.sweep import ( BudgetLedger, PoolConfig, claim_next_job, collect_job_status, complete_job_attempt, generate_jobs, generate_minimal_scaling_jobs, read_jobs, rescue_templates, run_node, utilization_summary, can_expand_pool, ) from airfrans_frontier.training.config import load_training_config class SweepFoundationTests(unittest.TestCase): def test_generate_jobs_writes_deterministic_manifests_and_coordinate_encoding_configs(self) -> None: with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) jobs = generate_jobs( output_dir=root, data_root="artifacts/data_cache/airfrans_processed/processed/full", bands=("100m",), families=("film_fourier_inr",), encodings=("nerf_multires", "raw"), group="test_sweep", ) persisted = read_jobs(root / "jobs.jsonl") self.assertEqual([job["job_id"] for job in persisted], [job["job_id"] for job in jobs]) self.assertTrue((root / "pool.toml").is_file()) self.assertTrue((root / "budget_ledger.json").is_file()) config = load_training_config(root / "configs" / "100m_film_fourier_inr_nerf_multires.toml") self.assertEqual(config.coordinate_encoding.type, "nerf_multires") self.assertEqual(config.coordinate_encoding.features, ("x", "y", "sdf")) self.assertEqual(config.model_metadata.reported_family, "film_fourier_inr") self.assertFalse(config.model_metadata.is_proxy) self.assertTrue((root / "rescue_templates.json").is_file()) self.assertEqual(config.model.encoding_levels, 16) self.assertEqual(config.model.fourier_scales, ()) self.assertEqual(config.model.hidden_width, 2048) self.assertEqual(config.model.depth, 16) def test_generate_minimal_scaling_jobs_are_all_points_and_non_raw_for_main_coordinate_models(self) -> None: with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) jobs = generate_minimal_scaling_jobs(output_dir=root, include_proxies=False) self.assertEqual(len(jobs), 23) self.assertEqual([job["job_id"] for job in jobs[:3]], ["A_data_all_points_smoke", "A_checkpoint_700m_footprint_smoke", "A_plateau_interrupt_smoke"]) serious_raw = [ job["job_id"] for job in jobs if job["phase"] in {"B_family_viability", "C_scaling_ladder"} and job["coordinate_encoding"] == "raw" and job["model_family"] != "siren_conditioned_inr" ] self.assertEqual(serious_raw, []) self.assertTrue(all(job["all_points_per_case"] for job in jobs)) c10 = load_training_config(root / "configs" / "C_10m_film_fourier_inr_nerf_multires_scale.toml") c100 = load_training_config(root / "configs" / "C_100m_film_fourier_inr_nerf_multires_scale.toml") config = load_training_config(root / "configs" / "C_700m_film_fourier_inr_nerf_multires_scale.toml") self.assertTrue(config.data.all_points_per_case) self.assertEqual(config.data.source, "huggingface_streaming") self.assertEqual(config.data.hf_repo_id, "zacheryasc/airfrans-processed") self.assertEqual(config.data.hf_path_prefix, "processed/full") self.assertEqual(config.data.train_cases, 900) self.assertEqual(config.data.val_cases, 50) self.assertEqual(config.data.test_cases, 50) self.assertEqual(config.data.streaming_normalization_cases, 4) self.assertEqual(config.data.streaming_queue_max_cases, 8) self.assertEqual(c10.run.seed, c100.run.seed) self.assertEqual(c100.run.seed, config.run.seed) self.assertEqual(config.checkpoint.policy, "full") self.assertTrue(config.checkpoint.include_optimizer_state) self.assertEqual(config.stability.dead_curve_patience_evals, 8) def test_generate_jobs_only_crosses_encoding_axis_for_compatible_families(self) -> None: with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) generate_jobs( output_dir=root, data_root="data/full", bands=("100m",), families=("film_fourier_inr", "siren_conditioned_inr"), encodings=("raw", "random_fourier"), ) job_ids = [job["job_id"] for job in read_jobs(root / "jobs.jsonl")] proxy_config = load_training_config(root / "configs" / "100m_raster_fno_unet_raw.toml") if (root / "configs" / "100m_raster_fno_unet_raw.toml").is_file() else None if proxy_config is not None: self.assertEqual(proxy_config.model_metadata.reported_family, "raster_fno_unet_proxy") self.assertIn("100m_film_fourier_inr_random_fourier", job_ids) self.assertIn("100m_siren_conditioned_inr_raw", job_ids) self.assertNotIn("100m_siren_conditioned_inr_random_fourier", job_ids) def test_collect_reconstructs_success_from_required_job_artifacts(self) -> None: with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) run_root = root / "runs" jobs = generate_jobs( output_dir=root / "sweep", data_root="data/full", artifact_dir=str(run_root), bands=("100m",), families=("mlp",), encodings=("raw",), ) run_dir = run_root / "20260727T000000Z_100m_mlp_raw" run_dir.mkdir(parents=True) for name in ( "config.toml", "job_manifest.json", "metrics.jsonl", "latest_metrics.json", "final_metrics.json", "run_manifest.json", "environment_manifest.json", "utilization.jsonl", ): (run_dir / name).write_text("{}\n") status = collect_job_status(jobs_path=root / "sweep" / "jobs.jsonl") self.assertEqual(status["counts"], {"succeeded": 1}) self.assertEqual(status["jobs"][0]["job_id"], jobs[0]["job_id"]) self.assertEqual(status["jobs"][0]["missing_artifacts"], []) def test_claim_complete_and_stale_recovery_are_durable(self) -> None: with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) jobs = generate_jobs( output_dir=root / "sweep", data_root="data/full", artifact_dir=str(root / "runs"), bands=("100m",), families=("mlp",), encodings=("raw",), ) jobs_path = root / "sweep" / "jobs.jsonl" claimed = claim_next_job(jobs_path=jobs_path, node_name="node-a", now=100.0) assert claimed is not None self.assertEqual(claimed["status"], "running") self.assertEqual(claimed["attempts"], 1) self.assertEqual(claimed["node"], "node-a") self.assertIsNone(claim_next_job(jobs_path=jobs_path, node_name="node-b", now=101.0)) completed = complete_job_attempt( jobs_path=jobs_path, attempt_id=claimed["last_attempt_id"], returncode=7, stdout="out", stderr="err", now=102.0, ) self.assertEqual(completed["status"], "failed") attempts = [json.loads(line) for line in (root / "sweep" / "attempts.jsonl").read_text().splitlines()] self.assertEqual([item["event"] for item in attempts], ["started", "failed"]) persisted = read_jobs(jobs_path) persisted[0].update({"status": "running", "lease_updated_at": 10.0, "last_attempt_id": "stale"}) from airfrans_frontier.sweep import write_jobs write_jobs(jobs_path, persisted) reclaimed = claim_next_job(jobs_path=jobs_path, node_name="node-c", stale_after_seconds=5.0, now=20.0) assert reclaimed is not None self.assertEqual(reclaimed["job_id"], jobs[0]["job_id"]) self.assertEqual(reclaimed["node"], "node-c") events = [json.loads(line)["event"] for line in (root / "sweep" / "attempts.jsonl").read_text().splitlines()] self.assertIn("stale_recovered", events) def test_run_node_drains_jobs_with_patched_training_process(self) -> None: with tempfile.TemporaryDirectory() as tmp: root = Path(tmp) generate_jobs( output_dir=root / "sweep", data_root="data/full", artifact_dir=str(root / "runs"), bands=("100m",), families=("mlp",), encodings=("raw",), ) completed = Mock(returncode=0, stdout="ok", stderr="") with patch("airfrans_frontier.sweep.subprocess.run", Mock(return_value=completed)): summary = run_node(jobs_path=root / "sweep" / "jobs.jsonl", node_name="node-a", max_jobs=1) self.assertEqual(summary["claimed"], 1) self.assertEqual(summary["succeeded"], 1) self.assertEqual(read_jobs(root / "sweep" / "jobs.jsonl")[0]["status"], "succeeded") def test_rescue_templates_cover_spec_axes_deterministically(self) -> None: templates = rescue_templates(("siren_conditioned_inr", "film_fourier_inr", "meshgraphnet_or_point_transformer_local")) keys = [(str(template["scope"]), str(template["template_id"])) for template in templates] ids = [template["template_id"] for template in templates] self.assertEqual(keys, sorted(keys)) self.assertIn("optimizer_lr_3e-4", ids) self.assertIn("siren_omega0_10", ids) self.assertIn("film_condition_width_2048", ids) self.assertIn("local_neighbors_32", ids) def test_expansion_gate_requires_utilization_backlog_stability_and_budget(self) -> None: pool = PoolConfig(nodes=(), budget_usd=25.0) good_utilization = {"gpu_util_median": 91.0, "gpu_idle_fraction": 0.04} allowed, reasons = can_expand_pool( pool=pool, utilization=good_utilization, backlog=3, recent_failures=0, ledger=BudgetLedger(budget_usd=25.0, spent_usd=5.0), next_pool_cost_usd=10.0, ) self.assertTrue(allowed) self.assertEqual(reasons, ()) blocked, reasons = can_expand_pool( pool=pool, utilization={"gpu_util_median": 70.0, "gpu_idle_fraction": 0.20}, backlog=0, recent_failures=1, ledger=BudgetLedger(budget_usd=25.0, spent_usd=24.0), next_pool_cost_usd=2.0, ) self.assertFalse(blocked) self.assertIn("median GPU utilization below gate", reasons) self.assertIn("GPU idle fraction above gate", reasons) self.assertIn("no job backlog", reasons) self.assertIn("recent failures are not isolated", reasons) self.assertIn("remaining budget does not support larger pool", reasons) def test_utilization_summary_reports_gate_metrics(self) -> None: with tempfile.TemporaryDirectory() as tmp: path = Path(tmp) / "utilization.jsonl" path.write_text( "".join( json.dumps(sample) + "\n" for sample in ( {"gpu_util_percent": 90, "memory_used_mb": 1000}, {"gpu_util_percent": 95, "memory_used_mb": 1500}, {"gpu_util_percent": 0, "memory_used_mb": 1200}, ) ) ) summary = utilization_summary(path) self.assertEqual(summary["gpu_util_median"], 90.0) self.assertAlmostEqual(summary["gpu_idle_fraction"], 1 / 3) self.assertEqual(summary["memory_used_peak"], 1500.0) if __name__ == "__main__": unittest.main()