airfRANS-model-exploration/tests/test_streaming_data.py

572 lines
24 KiB
Python
Raw Normal View History

from __future__ import annotations
import gzip
import json
2026-07-29 09:42:52 +00:00
import shutil
import os
import sys
import tempfile
2026-07-29 09:42:52 +00:00
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,
2026-07-27 08:48:37 +00:00
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}
2026-07-27 08:48:37 +00:00
{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"
)
2026-07-29 09:42:52 +00:00
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_*")))
2026-07-27 08:48:37 +00:00
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)
2026-07-29 09:42:52 +00:00
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())
2026-07-28 18:04:50 +00:00
self.assertTrue(verification["ok"])
self.assertEqual(verification["checks"]["terminal_artifact"], "failure_report.json")
if __name__ == "__main__":
unittest.main()