airfRANS-model-exploration/tests/test_remote_run.py

321 lines
13 KiB
Python

from __future__ import annotations
from contextlib import redirect_stdout
from io import StringIO
import json
import tempfile
import shutil
import unittest
from unittest.mock import Mock, patch
from pathlib import Path
from airfrans_frontier.runtime import remove_pythonpath_entries
remove_pythonpath_entries()
import torch
from airfrans_frontier.remote.artifacts import verify_artifacts
from airfrans_frontier.remote.cleanup import reconcile_cleanup
from airfrans_frontier.remote.cli import _classify_artifacts, _stage_resume_checkpoint, _terminal_artifact_names, main as remote_main
from airfrans_frontier.remote.config import load_remote_run_config
from airfrans_frontier.remote.skypilot import render_skypilot_yaml
from airfrans_frontier.remote.vast import VastOffer, choose_offer
class RemoteRunConfigTests(unittest.TestCase):
def test_loads_remote_smoke_config(self) -> None:
config = load_remote_run_config("configs/remote_smoke.toml")
self.assertEqual(config.provider.kind, "vastai")
self.assertEqual(config.provider.gpu.name, "RTX 4090")
self.assertEqual(config.job.artifact_dir.as_posix(), "artifacts/current_run")
self.assertIn("checkpoint_latest.pt", config.artifacts.required)
def test_loads_remote_hf_smoke_config_for_cheap_upload(self) -> None:
config = load_remote_run_config("configs/remote_hf_smoke.toml")
self.assertEqual(config.provider.kind, "vastai")
self.assertEqual(config.provider.gpu.name, "RTX 3060")
self.assertLessEqual(config.provider.max_price_per_hour or 999.0, 0.08)
self.assertEqual(config.artifacts.mode, "object_store_upload")
self.assertIn("remote-run hf-smoke", config.job.command)
self.assertIn("hf_upload_manifest.json", config.artifacts.required)
def test_loads_remote_wandb_smoke_config_for_observability(self) -> None:
config = load_remote_run_config("configs/remote_wandb_smoke.toml")
self.assertEqual(config.provider.kind, "vastai")
self.assertEqual(config.provider.gpu.name, "RTX 3060 Ti")
self.assertLessEqual(config.provider.max_price_per_hour or 999.0, 0.09)
self.assertEqual(config.artifacts.mode, "rsync")
self.assertIn("remote-run wandb-smoke", config.job.command)
self.assertIn("wandb_smoke_manifest.json", config.artifacts.required)
class VastSelectionTests(unittest.TestCase):
def test_selection_filters_bad_hosts_and_drops_suspiciously_cheap_tail(self) -> None:
config = load_remote_run_config("configs/remote_smoke.toml")
offers = [
offer(1, price=0.10, host=1),
offer(2, price=0.20, host=2),
offer(3, price=0.30, host=3),
offer(4, price=0.40, host=4),
offer(5, price=0.50, host=59017),
offer(6, price=0.25, host=6, geo="CN"),
offer(7, price=0.26, host=7, verification="deverified"),
]
result = choose_offer(offers, config, query={"test": True})
# Four reachable RTX 4090 offers remain; drop_cheap_frac=0.30 drops floor(1.2)=1 cheapest.
self.assertEqual(result.selected_offer_id, 2)
self.assertEqual(result.candidate_count, 7)
self.assertEqual(result.survivor_count, 3)
def test_rendered_yaml_injects_selected_offer(self) -> None:
config = load_remote_run_config("configs/remote_smoke.toml")
result = choose_offer([offer(123, price=0.30, host=22), offer(124, price=0.40, host=23)], config, query={})
yaml = render_skypilot_yaml(config, result, run_id="airfrans-test")
self.assertIn("selected_offer_id: 123", yaml)
self.assertNotIn("sky launch", yaml)
self.assertIn("remote-run smoke-train", yaml)
self.assertIn("configs/aggressive_smoke.toml", yaml)
self.assertIn("df -h .", yaml)
self.assertIn("airfrans_disk_requested_gb=128", yaml)
self.assertIn("AIRFRANS_STARTUP_TIMELINE: artifacts/current_run/startup_timeline.jsonl", yaml)
self.assertIn("airfrans_timeline 'setup' 'started'", yaml)
self.assertIn("airfrans_timeline 'data_validation' 'started'", yaml)
self.assertIn("airfrans_timeline 'training_command' 'started'", yaml)
def test_terminal_collection_includes_startup_timeline(self) -> None:
self.assertIn("startup_timeline.jsonl", _terminal_artifact_names(()))
def test_rendered_yaml_can_pass_resume_checkpoint(self) -> None:
config = load_remote_run_config("configs/remote_smoke.toml")
result = choose_offer([offer(123, price=0.30, host=22), offer(124, price=0.40, host=23)], config, query={})
yaml = render_skypilot_yaml(
config,
result,
run_id="airfrans-test",
resume_checkpoint=".airfrans_resume/airfrans-test/checkpoint_latest.pt",
)
self.assertIn("AIRFRANS_RESUME_CHECKPOINT: .airfrans_resume/airfrans-test/checkpoint_latest.pt", yaml)
class VastInstanceCliTests(unittest.TestCase):
def test_vast_instances_reports_api_ground_truth(self) -> None:
stdout = StringIO()
with patch.dict("os.environ", {"VAST_API_KEY": "token"}), patch(
"airfrans_frontier.remote.cli.list_instances",
return_value=[{"id": 123, "actual_status": "running", "gpu_name": "RTX 4090"}],
), redirect_stdout(stdout):
code = remote_main(["vast-instances"])
self.assertEqual(code, 0)
payload = json.loads(stdout.getvalue())
self.assertEqual(payload["instance_count"], 1)
self.assertEqual(payload["instances"][0]["id"], 123)
self.assertEqual(payload["instances"][0]["actual_status"], "running")
def test_cleanup_reconcile_uses_current_sky_json_flag(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
vast_path = Path(tmp) / "vast.json"
vast_path.write_text(json.dumps({"instances": []}) + "\n")
stdout = StringIO()
with patch(
"airfrans_frontier.remote.cli.subprocess.run",
return_value=Mock(stdout=json.dumps({"clusters": []})),
) as run, redirect_stdout(stdout):
code = remote_main(["cleanup-reconcile", "--vast-instances-json", str(vast_path)])
self.assertEqual(code, 0)
run.assert_called_once()
self.assertEqual(run.call_args.args[0], ["sky", "status", "-o", "json"])
def test_cleanup_reconcile_treats_sky_node_names_as_live_instances(self) -> None:
destroy_instance = Mock()
report = reconcile_cleanup(
sky_state=[
{
"name": "full_70gb_model_zoo_20260726_film_fourier_inr-20260726T094528Z",
"node_names": "45895903",
"status": "UP",
}
],
vast_instances=[
{
"id": 45895903,
"label": "full-70gb-model-zoo-20260726-film-fourier-inr-20260726t094528z-166d0dab-head",
"actual_status": "running",
"gpu_name": "RTX 4090",
}
],
destroy_orphans=True,
destroy_instance=destroy_instance,
)
self.assertEqual(report["unexpected_live_count"], 0)
self.assertEqual(report["instances"][0]["cleanup_result"], "not_needed")
destroy_instance.assert_not_called()
def test_cleanup_reconcile_ignores_declared_external_vast_labels(self) -> None:
destroy_instance = Mock()
report = reconcile_cleanup(
sky_state=[],
vast_instances=[
{
"id": 45900657,
"label": "mvp-1785062905827-4",
"actual_status": "running",
"gpu_name": "Titan RTX",
},
{
"id": 45901000,
"label": "full-70gb-model-zoo-20260726-film-fourier-inr-canary-head",
"actual_status": "running",
"gpu_name": "RTX 4090",
},
],
ignored_label_prefixes=("mvp-",),
destroy_orphans=True,
destroy_instance=destroy_instance,
)
self.assertEqual(report["ignored_external_count"], 1)
self.assertEqual(report["unexpected_live_count"], 1)
self.assertEqual(report["instances"][0]["cleanup_result"], "ignored")
self.assertEqual(report["instances"][0]["label"], "mvp-1785062905827-4")
destroy_instance.assert_called_once_with(45901000)
def test_cleanup_reconcile_matches_hyphenated_vast_labels_to_sky_names(self) -> None:
destroy_instance = Mock()
report = reconcile_cleanup(
sky_state=[
{
"name": "full_70gb_model_zoo_20260726_raster_fno_unet-20260726T095257Z",
"status": "INIT",
}
],
vast_instances=[
{
"id": 45896431,
"label": "full-70gb-model-zoo-20260726-raster-fno-unet-20260726t095257z-166d0dab-head",
"actual_status": "loading",
"gpu_name": "RTX 4090",
}
],
destroy_orphans=True,
destroy_instance=destroy_instance,
)
self.assertEqual(report["unexpected_live_count"], 0)
self.assertEqual(
report["instances"][0]["associated_run_id"],
"full_70gb_model_zoo_20260726_raster_fno_unet-20260726T095257Z",
)
destroy_instance.assert_not_called()
class ArtifactVerificationTests(unittest.TestCase):
def test_verify_artifacts_requires_contract_files_and_writes_manifest(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
root = Path(tmp)
_write_contract_artifacts(root, success=True)
manifest = verify_artifacts(root)
self.assertGreaterEqual(manifest["file_count"], 8)
self.assertTrue((root / "artifact_manifest.json").is_file())
self.assertTrue((root / "checksums.txt").is_file())
self.assertTrue((root / "verification_report.json").is_file())
report = json.loads((root / "verification_report.json").read_text())
self.assertTrue(report["ok"])
def test_verify_artifacts_accepts_failure_report_terminal_state(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
root = Path(tmp)
_write_contract_artifacts(root, success=False)
(root / "failure_report.json").write_text(json.dumps({"error_type": "NonFiniteLoss"}) + "\n")
manifest = verify_artifacts(root)
self.assertGreaterEqual(manifest["file_count"], 7)
def test_classifies_remote_success_before_checkpoint_collection(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
root = Path(tmp)
(root / "final_metrics.json").write_text(json.dumps({"loss": 1.0}) + "\n")
self.assertEqual(_classify_artifacts(root), "success")
def test_classifies_and_stages_restart_checkpoint(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
root = Path(tmp)
self.assertEqual(_classify_artifacts(root), "incomplete")
_write_contract_artifacts(root, success=False)
self.assertEqual(_classify_artifacts(root), "restartable")
staged = _stage_resume_checkpoint(root, "test-run")
self.assertIsNotNone(staged)
assert staged is not None
self.assertTrue(staged.is_file())
self.assertEqual(staged.as_posix(), ".airfrans_resume/test-run/checkpoint_latest.pt")
shutil.rmtree(".airfrans_resume")
def _write_contract_artifacts(root: Path, *, success: bool) -> None:
(root / "config.toml").write_text("[run]\nname = 'test'\n")
(root / "metrics.jsonl").write_text(json.dumps({"step": 0}) + "\n")
(root / "latest_metrics.json").write_text(json.dumps({"step": 0}) + "\n")
(root / "heartbeat.json").write_text(json.dumps({"phase": "training"}) + "\n")
checkpoint = {
"schema_version": 1,
"step": 0,
"model_state_dict": {},
"optimizer_state_dict": {},
"normalization": {},
}
torch.save(checkpoint, root / "checkpoint_latest.pt")
torch.save(checkpoint, root / "checkpoint_best.pt")
if success:
torch.save(checkpoint, root / "checkpoint_final.pt")
(root / "final_metrics.json").write_text(json.dumps({"loss": 1.0}) + "\n")
def offer(
offer_id: int,
*,
price: float,
host: int,
geo: str = "US",
verification: str = "verified",
) -> VastOffer:
return VastOffer(
id=offer_id,
gpu_name="RTX 4090",
dph_total=price,
gpu_ram=24_000,
disk_space=256.0,
geolocation=geo,
inet_down_cost_per_tb=0.0,
inet_up_cost_per_tb=0.0,
host_id=host,
verification=verification,
reliability2=0.99,
cuda_max_good=12.8,
direct_port_count=1,
inet_down=500.0,
inet_up=100.0,
verified=True,
)
if __name__ == "__main__":
unittest.main()