2026-07-26 09:05:38 +00:00
|
|
|
from __future__ import annotations
|
|
|
|
|
|
|
|
|
|
import gzip
|
|
|
|
|
import json
|
2026-07-29 09:42:52 +00:00
|
|
|
import shutil
|
2026-07-26 09:05:38 +00:00
|
|
|
import os
|
|
|
|
|
import sys
|
|
|
|
|
import tempfile
|
2026-07-29 09:42:52 +00:00
|
|
|
import time
|
2026-07-26 09:05:38 +00:00
|
|
|
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,
|
2026-07-26 09:05:38 +00:00
|
|
|
) -> 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 ""}
|
2026-07-26 09:05:38 +00:00
|
|
|
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
|
|
|
|
|
|
2026-07-26 09:05:38 +00:00
|
|
|
|
|
|
|
|
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])
|
|
|
|
|
|
2026-07-26 09:05:38 +00:00
|
|
|
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")
|
2026-07-26 09:05:38 +00:00
|
|
|
|
|
|
|
|
|
|
|
|
|
if __name__ == "__main__":
|
|
|
|
|
unittest.main()
|