from __future__ import annotations import gzip import json import shutil import os import sys import tempfile import time import types import unittest import zipfile from pathlib import Path from unittest.mock import patch from airfrans_frontier.runtime import remove_pythonpath_entries remove_pythonpath_entries() import numpy as np from airfrans_frontier.raw.public import process_of_dataset_url_streaming from airfrans_frontier.training.config import load_training_config from airfrans_frontier.training.data import build_dataset_bundle, load_processed_dataset from airfrans_frontier.training.loop import train from airfrans_frontier.training.normalize import compute_normalization_stats from airfrans_frontier.training.streaming_data import StreamingEventRecorder, StreamingTrainingData def write_minimal_airfrans_archive(archive: Path, case_names: list[str]) -> None: with zipfile.ZipFile(archive, "w") as zf: for index, case_name in enumerate(case_names): base = f"OF_dataset/{case_name}" u_value = 1.0 + 0.1 * index p_value = 0.5 + 0.2 * index nut_value = 0.01 + 0.001 * index zf.writestr(f"{base}/constant/transportProperties", "nu 1e-5;\n") zf.writestr( f"{base}/constant/polyMesh/boundary", "\naerofoil\n{\n type wall;\n nFaces 1;\n startFace 0;\n}\nfarfield\n{\n type patch;\n nFaces 3;\n startFace 1;\n}\n", ) zf.writestr(f"{base}/constant/polyMesh/points.gz", gzip.compress(b"4\n(\n(0 0 0)\n(1 0 0)\n(1 1 0)\n(0 1 0)\n)\n")) zf.writestr(f"{base}/constant/polyMesh/faces.gz", gzip.compress(b"4\n(\n2(0 1)\n2(1 2)\n2(2 3)\n2(3 0)\n)\n")) zf.writestr(f"{base}/constant/polyMesh/owner.gz", gzip.compress(b"4\n(\n0\n0\n0\n0\n)\n")) zf.writestr(f"{base}/constant/polyMesh/neighbour.gz", gzip.compress(b"0\n(\n)\n")) zf.writestr(f"{base}/1/U.gz", gzip.compress(f"1\n(\n({u_value} 0 0)\n)\n".encode())) zf.writestr(f"{base}/1/p.gz", gzip.compress(f"1\n(\n{p_value}\n)\n".encode())) zf.writestr(f"{base}/1/nut.gz", gzip.compress(f"1\n(\n{nut_value}\n)\n".encode())) def write_malformed_airfrans_archive(archive: Path, case_name: str) -> None: with zipfile.ZipFile(archive, "w") as zf: zf.writestr(f"OF_dataset/{case_name}/constant/transportProperties", "nu 1e-5;\n") def write_streaming_config( path: Path, *, archive: Path, cache_dir: Path, artifact_dir: Path, train_cases: int = 2, val_cases: int = 1, test_cases: int = 1, steps: int = 2, log_interval: int = 1, batch_size: int = 2, high_water_bytes: int = 32 * 1024 * 1024, low_water_bytes: int = 16 * 1024 * 1024, upload_processed: bool = False, upload_batch_size: int = 1, normalization_cases: int | None = None, ) -> None: path.write_text( f""" [run] name = "streaming_test" seed = 7 artifact_dir = "{artifact_dir}" [data] root = "{cache_dir}" source = "public_zip_streaming" public_source_url = "{archive}" cache_dir = "{cache_dir}" streaming_scratch_dir = "{cache_dir / '_raw'}" train_cases = {train_cases} val_cases = {val_cases} test_cases = {test_cases} points_per_case = 999999999 batch_size = {batch_size} streaming_cache_max_bytes = {max(high_water_bytes, high_water_bytes + 1)} streaming_cache_high_water_bytes = {high_water_bytes} streaming_cache_low_water_bytes = {low_water_bytes} streaming_queue_max_cases = 1 streaming_upload_processed = {str(upload_processed).lower()} streaming_upload_batch_size = {upload_batch_size} {f"streaming_normalization_cases = {normalization_cases}" if normalization_cases is not None else ""} hf_repo_id = "owner/airfrans-processed" hf_repo_type = "dataset" hf_path_prefix = "processed/full" [model] type = "mlp" hidden_width = 16 depth = 2 activation = "gelu" [optim] lr = 0.01 weight_decay = 0.0 steps = {steps} log_interval = {log_interval} [device] type = "cpu" allow_cpu_fallback = false benchmark_kernels = false [loss] type = "normalized_mse" [checkpoint] interval_seconds = 0 """.strip() + "\n" ) def write_huggingface_streaming_config( path: Path, *, cache_dir: Path, artifact_dir: Path, train_cases: int = 4, val_cases: int = 1, test_cases: int = 1, steps: int = 1, normalization_cases: int = 1, ) -> None: path.write_text( f""" [run] name = "hf_streaming_test" seed = 7 artifact_dir = "{artifact_dir}" [data] root = "{cache_dir / 'processed' / 'full'}" source = "huggingface_streaming" hf_repo_id = "owner/airfrans-processed" hf_repo_type = "dataset" hf_path_prefix = "processed/full" cache_dir = "{cache_dir}" train_cases = {train_cases} val_cases = {val_cases} test_cases = {test_cases} all_points_per_case = true batch_size = 2 streaming_queue_max_cases = 2 streaming_normalization_cases = {normalization_cases} [model] type = "mlp" hidden_width = 16 depth = 2 activation = "gelu" [optim] lr = 0.01 weight_decay = 0.0 steps = {steps} log_interval = 1 [device] type = "cpu" allow_cpu_fallback = false benchmark_kernels = false [loss] type = "normalized_mse" [checkpoint] interval_seconds = 0 policy = "full" include_optimizer_state = true include_rng_state = true """.strip() + "\n" ) def write_processed_case(root: Path, relative_prefix: str, case_id: str, offset: float) -> None: target = root / relative_prefix / f"{case_id}.npz" target.parent.mkdir(parents=True, exist_ok=True) features = np.asarray( [ [offset + 0.0, 0.0, 1.0, 0.1], [offset + 1.0, 1.0, 0.5, 0.2], [offset + 2.0, 0.5, 0.25, 0.3], [offset + 3.0, 0.25, 0.125, 0.4], ], dtype=np.float32, ) targets = np.asarray( [ [offset + 0.0, 0.1, 0.2], [offset + 0.2, 0.3, 0.4], [offset + 0.4, 0.5, 0.6], [offset + 0.6, 0.7, 0.8], ], dtype=np.float32, ) np.savez( target, features=features, targets=targets, feature_names=np.asarray(["x", "y", "sdf", "alpha"], dtype="U"), target_names=np.asarray(["u", "v", "p"], dtype="U"), ) def fake_huggingface_module(source_root: Path, calls: list[str]) -> types.ModuleType: module = types.ModuleType("huggingface_hub") class FakeHfApi: def __init__(self, token: str | None = None) -> None: self.token = token def list_repo_files(self, *, repo_id: str, repo_type: str) -> list[str]: return sorted(str(path.relative_to(source_root)) for path in source_root.rglob("*") if path.is_file()) def hf_hub_download(*, repo_id: str, filename: str, repo_type: str, local_dir: str, token: str | None = None) -> str: calls.append(filename) source = source_root / filename destination = Path(local_dir) / filename time.sleep(0.01) destination.parent.mkdir(parents=True, exist_ok=True) shutil.copy2(source, destination) return str(destination) module.HfApi = FakeHfApi module.hf_hub_download = hf_hub_download return module def read_events(run_dir: Path) -> list[dict[str, object]]: return [json.loads(line) for line in (run_dir / "streaming_events.jsonl").read_text().splitlines() if line.strip()] class FullDataBackpressureStreamingTests(unittest.TestCase): def test_streaming_training_smoke_writes_artifacts_without_eager_concatenation(self) -> None: with tempfile.TemporaryDirectory() as tmp: tmp_path = Path(tmp) archive = tmp_path / "OF_dataset.zip" case_names = [f"airFoil2D_SST_{10 + index}.0_5.0_0012" for index in range(5)] write_minimal_airfrans_archive(archive, case_names) config_path = tmp_path / "streaming.toml" artifact_dir = tmp_path / "artifacts" write_streaming_config(config_path, archive=archive, cache_dir=tmp_path / "cache", artifact_dir=artifact_dir) config = load_training_config(config_path) with patch("airfrans_frontier.training.loop.load_processed_dataset", side_effect=AssertionError("eager load called")), patch( "airfrans_frontier.training.loop.build_dataset_bundle", side_effect=AssertionError("eager concat called") ): result = train(config) self.assertTrue(np.isfinite(result.final_metrics["train_loss"])) self.assertEqual(result.final_metrics["data_mode"], "public_zip_streaming") for name in ( "metrics.jsonl", "checkpoint_latest.pt", "checkpoint_best.pt", "checkpoint_final.pt", "final_metrics.json", "split_manifest.json", "data_manifest.json", "normalization.json", "streaming_events.jsonl", "streaming_state.json", "streaming_summary.json", "processed_upload_manifest.json", "artifact_manifest.json", "checksums.txt", "verification_report.json", ): self.assertTrue((result.run_dir / name).is_file(), name) events = read_events(result.run_dir) event_names = {event["event"] for event in events} self.assertIn("dataset_enumeration_start", event_names) self.assertIn("dataset_enumeration_end", event_names) self.assertIn("split_selection", event_names) self.assertIn("normalization_start", event_names) self.assertIn("normalization_end", event_names) self.assertIn("first_batch_ready", event_names) self.assertIn("first_gpu_batch_consumed", event_names) self.assertIn("first_metric", event_names) self.assertIn("first_checkpoint_written", event_names) selected_cases = set(json.loads((result.run_dir / "data_manifest.json").read_text())["cases"][index]["case_id"] for index in range(4)) processed_cases = {str(event["case_id"]) for event in events if event["event"] == "processing_end"} self.assertLessEqual(processed_cases, selected_cases) self.assertFalse(any((tmp_path / "cache" / "_raw").glob("airFoil2D_*"))) def test_streaming_fast_start_reaches_first_gpu_batch_before_all_train_cases_are_processed(self) -> None: with tempfile.TemporaryDirectory() as tmp: tmp_path = Path(tmp) archive = tmp_path / "OF_dataset.zip" case_names = [f"airFoil2D_SST_{10 + index}.0_5.0_0012" for index in range(6)] write_minimal_airfrans_archive(archive, case_names) config_path = tmp_path / "streaming.toml" artifact_dir = tmp_path / "artifacts" write_streaming_config( config_path, archive=archive, cache_dir=tmp_path / "cache", artifact_dir=artifact_dir, train_cases=4, val_cases=1, test_cases=1, steps=1, normalization_cases=1, ) result = train(load_training_config(config_path)) events = read_events(result.run_dir) first_gpu = next(index for index, event in enumerate(events) if event["event"] == "first_gpu_batch_consumed") processed_before_gpu = { str(event["case_id"]) for event in events[:first_gpu] if event["event"] == "processing_end" } self.assertLess(len(processed_before_gpu), 4) normalization_end = next(event for event in events if event["event"] == "normalization_end") self.assertEqual(normalization_end["normalization_cases"], 1) def test_huggingface_streaming_fast_start_downloads_remaining_selected_cases(self) -> None: with tempfile.TemporaryDirectory() as tmp: tmp_path = Path(tmp) hf_source = tmp_path / "hf_source" case_names = [f"case_{index:04d}" for index in range(6)] for index, case_name in enumerate(case_names): write_processed_case(hf_source, "processed/full", case_name, float(index)) config_path = tmp_path / "hf_streaming.toml" cache_dir = tmp_path / "cache" artifact_dir = tmp_path / "artifacts" write_huggingface_streaming_config( config_path, cache_dir=cache_dir, artifact_dir=artifact_dir, train_cases=4, val_cases=1, test_cases=1, steps=1, normalization_cases=1, ) calls: list[str] = [] fake_module = fake_huggingface_module(hf_source, calls) with patch.dict(sys.modules, {"huggingface_hub": fake_module}): result = train(load_training_config(config_path)) self.assertEqual(result.final_metrics["data_mode"], "huggingface_streaming") events = read_events(result.run_dir) event_names = {event["event"] for event in events} self.assertIn("background_acquisition_start", event_names) self.assertIn("hf_case_download_end", event_names) self.assertIn("background_acquisition_complete", event_names) first_gpu = next(index for index, event in enumerate(events) if event["event"] == "first_gpu_batch_consumed") processed_before_gpu = { str(event["case_id"]) for event in events[:first_gpu] if event["event"] == "processing_end" } self.assertLess(len(processed_before_gpu), 4) manifest = json.loads((result.run_dir / "data_manifest.json").read_text()) self.assertEqual(manifest["source"], "huggingface_streaming") cached_cases = sorted(path.stem for path in (cache_dir / "processed" / "full").glob("*.npz")) self.assertEqual(cached_cases, case_names) self.assertEqual(sorted(set(calls)), [f"processed/full/{case_name}.npz" for case_name in case_names]) def test_backpressure_pauses_resumes_and_bounds_cache_with_inflight_slack(self) -> None: with tempfile.TemporaryDirectory() as tmp: tmp_path = Path(tmp) archive = tmp_path / "OF_dataset.zip" case_names = [f"airFoil2D_SST_{10 + index}.0_5.0_0012" for index in range(4)] write_minimal_airfrans_archive(archive, case_names) config_path = tmp_path / "streaming.toml" high_water = 256 write_streaming_config( config_path, archive=archive, cache_dir=tmp_path / "cache", artifact_dir=tmp_path / "artifacts", high_water_bytes=high_water, low_water_bytes=128, steps=2, ) result = train(load_training_config(config_path)) summary = json.loads((result.run_dir / "streaming_summary.json").read_text()) self.assertGreater(summary["cache_high_water_events"], 0) self.assertGreater(summary["cache_low_water_events"], 0) self.assertGreater(summary["producer_pause_events"], 0) self.assertGreater(summary["producer_resume_events"], 0) self.assertGreater(summary["evicted_units"], 0) self.assertLessEqual(summary["processed_cache_high_water_bytes"], high_water + summary["max_processed_unit_bytes"]) event_names = {event["event"] for event in read_events(result.run_dir)} self.assertIn("producer_paused", event_names) self.assertIn("producer_resumed", event_names) self.assertIn("cleanup_eviction", event_names) def test_streaming_normalization_matches_eager_train_split_statistics(self) -> None: with tempfile.TemporaryDirectory() as tmp: tmp_path = Path(tmp) archive = tmp_path / "OF_dataset.zip" case_names = [f"airFoil2D_SST_{10 + index}.0_5.0_0012" for index in range(4)] write_minimal_airfrans_archive(archive, case_names) config_path = tmp_path / "streaming.toml" write_streaming_config(config_path, archive=archive, cache_dir=tmp_path / "cache", artifact_dir=tmp_path / "artifacts", steps=1) config = load_training_config(config_path) run_dir = tmp_path / "run" recorder = StreamingEventRecorder(run_dir) streaming = StreamingTrainingData.from_config(config, run_dir=run_dir, recorder=recorder) streaming.prepare() streaming_stats = streaming.load_or_compute_normalization() eager_root = tmp_path / "eager_processed" process_of_dataset_url_streaming(str(archive), eager_root, scratch_dir=tmp_path / "eager_raw", min_cases=4) eager_bundle = build_dataset_bundle( load_processed_dataset(eager_root), train_cases=config.data.train_cases, val_cases=config.data.val_cases, test_cases=config.data.test_cases, points_per_case=config.data.points_per_case, seed=config.run.seed, ) eager_stats = compute_normalization_stats( eager_bundle.train.features, eager_bundle.train.targets, feature_names=eager_bundle.feature_names, target_names=eager_bundle.target_names, ) np.testing.assert_allclose(streaming_stats.feature_mean, eager_stats.feature_mean, rtol=1e-6, atol=1e-6) np.testing.assert_allclose(streaming_stats.feature_std, eager_stats.feature_std, rtol=1e-6, atol=1e-6) np.testing.assert_allclose(streaming_stats.target_mean, eager_stats.target_mean, rtol=1e-6, atol=1e-6) np.testing.assert_allclose(streaming_stats.target_std, eager_stats.target_std, rtol=1e-6, atol=1e-6) def test_resume_reuses_validated_units_and_discards_partial_units(self) -> None: with tempfile.TemporaryDirectory() as tmp: tmp_path = Path(tmp) archive = tmp_path / "OF_dataset.zip" case_names = [f"airFoil2D_SST_{10 + index}.0_5.0_0012" for index in range(3)] write_minimal_airfrans_archive(archive, case_names) config_path = tmp_path / "streaming.toml" write_streaming_config( config_path, archive=archive, cache_dir=tmp_path / "cache", artifact_dir=tmp_path / "artifacts", train_cases=1, val_cases=1, test_cases=1, steps=1, ) config = load_training_config(config_path) run_dir = tmp_path / "run" first = StreamingTrainingData.from_config(config, run_dir=run_dir, recorder=StreamingEventRecorder(run_dir)) first.prepare() assert first.split is not None first_case = first.split.train_ids[0] (tmp_path / "cache" / f"{first_case}.npz.tmp.npz").write_bytes(b"partial") second = StreamingTrainingData.from_config(config, run_dir=run_dir, recorder=StreamingEventRecorder(run_dir)) second.prepare() events = read_events(run_dir) self.assertTrue(any(event["event"] == "partial_unit_discarded" and event.get("case_id") == first_case for event in events)) self.assertTrue(any(event["event"] == "resume_validated_unit_reused" and event.get("case_id") == first_case for event in events)) processing_events = [event for event in events if event["event"] == "processing_end" and event.get("case_id") == first_case] self.assertEqual(len(processing_events), 1) def test_processed_upload_rate_limit_does_not_fail_training(self) -> None: class FakeRateLimitError(RuntimeError): def __init__(self) -> None: super().__init__("429 Too Many Requests") self.response = types.SimpleNamespace(headers={"Retry-After": "600"}) class FakeCommitOperationAdd: def __init__(self, *, path_in_repo: str, path_or_fileobj: str) -> None: self.path_in_repo = path_in_repo self.path_or_fileobj = path_or_fileobj class FakeApi: def __init__(self, token: str) -> None: self.token = token def create_repo(self, *, repo_id: str, repo_type: str, private: bool, exist_ok: bool) -> None: return None def create_commit(self, **kwargs): raise FakeRateLimitError() fake_module = types.SimpleNamespace(HfApi=FakeApi, CommitOperationAdd=FakeCommitOperationAdd) with tempfile.TemporaryDirectory() as tmp, patch.dict(sys.modules, {"huggingface_hub": fake_module}), patch.dict(os.environ, {"HF_TOKEN": "token"}): tmp_path = Path(tmp) archive = tmp_path / "OF_dataset.zip" case_names = [f"airFoil2D_SST_{10 + index}.0_5.0_0012" for index in range(4)] write_minimal_airfrans_archive(archive, case_names) config_path = tmp_path / "streaming.toml" write_streaming_config( config_path, archive=archive, cache_dir=tmp_path / "cache", artifact_dir=tmp_path / "artifacts", upload_processed=True, upload_batch_size=1, steps=1, ) result = train(load_training_config(config_path)) self.assertTrue(np.isfinite(result.final_metrics["train_loss"])) manifest = json.loads((result.run_dir / "processed_upload_manifest.json").read_text()) self.assertTrue(manifest["enabled"]) self.assertGreater(manifest["queue_depth"], 0) self.assertGreater(manifest["rate_limit_until"], 0) self.assertEqual(manifest["rate_limit_retry_after_seconds"], 600.0) events = {event["event"] for event in read_events(result.run_dir)} self.assertIn("processed_data_upload_rate_limited", events) self.assertIn("processed_data_upload_suppressed", events) run_manifest = json.loads((result.run_dir / "run_manifest.json").read_text()) self.assertEqual(run_manifest["phase"], "completed") def test_streaming_failure_writes_diagnostic_artifacts(self) -> None: with tempfile.TemporaryDirectory() as tmp: tmp_path = Path(tmp) archive = tmp_path / "bad.zip" case_name = "airFoil2D_SST_10.0_5.0_0012" write_malformed_airfrans_archive(archive, case_name) config_path = tmp_path / "streaming.toml" artifact_dir = tmp_path / "artifacts" write_streaming_config( config_path, archive=archive, cache_dir=tmp_path / "cache", artifact_dir=artifact_dir, train_cases=1, val_cases=0, test_cases=0, steps=1, ) with self.assertRaises(Exception): train(load_training_config(config_path)) run_dir = next(path for path in artifact_dir.iterdir() if path.is_dir()) for name in ("failure_report.json", "metrics.jsonl", "streaming_events.jsonl", "streaming_state.json", "streaming_summary.json", "verification_report.json"): self.assertTrue((run_dir / name).is_file(), name) report = json.loads((run_dir / "failure_report.json").read_text()) self.assertEqual(report["phase"], "streaming_training") events = {event["event"] for event in read_events(run_dir)} self.assertIn("processing_failure", events) verification = json.loads((run_dir / "verification_report.json").read_text()) self.assertTrue(verification["ok"]) self.assertEqual(verification["checks"]["terminal_artifact"], "failure_report.json") if __name__ == "__main__": unittest.main()