96 lines
3.3 KiB
Python
96 lines
3.3 KiB
Python
from __future__ import annotations
|
|
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
from airfrans_frontier.runtime import remove_pythonpath_entries
|
|
|
|
remove_pythonpath_entries()
|
|
|
|
import numpy as np
|
|
|
|
from airfrans_frontier.training.data import create_case_split, load_processed_dataset, load_simulation_npz
|
|
from airfrans_frontier.training.normalize import compute_normalization_stats, normalize_targets
|
|
|
|
|
|
def write_case(path: Path, offset: float = 0.0) -> None:
|
|
features = np.array(
|
|
[
|
|
[offset + 0.0, 1.0],
|
|
[offset + 1.0, 2.0],
|
|
[offset + 2.0, 3.0],
|
|
],
|
|
dtype=np.float32,
|
|
)
|
|
targets = np.array(
|
|
[
|
|
[offset + 10.0, -1.0],
|
|
[offset + 11.0, 0.0],
|
|
[offset + 12.0, 1.0],
|
|
],
|
|
dtype=np.float32,
|
|
)
|
|
np.savez(
|
|
path,
|
|
features=features,
|
|
targets=targets,
|
|
feature_names=np.array(["x", "y"]),
|
|
target_names=np.array(["pressure", "velocity"]),
|
|
)
|
|
|
|
|
|
class TrainingDataTests(unittest.TestCase):
|
|
def test_dataset_loader_rejects_malformed_npz(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
path = Path(tmp) / "bad.npz"
|
|
np.savez(path, features=np.array([1.0, 2.0], dtype=np.float32), targets=np.ones((2, 1)))
|
|
|
|
with self.assertRaisesRegex(ValueError, "features to be a 2D array"):
|
|
load_simulation_npz(path)
|
|
|
|
def test_dataset_loader_loads_common_schema(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
root = Path(tmp)
|
|
write_case(root / "case_a.npz", offset=0.0)
|
|
write_case(root / "case_b.npz", offset=1.0)
|
|
|
|
samples = load_processed_dataset(root)
|
|
|
|
self.assertEqual([sample.case_id for sample in samples], ["case_a", "case_b"])
|
|
self.assertEqual(samples[0].feature_names, ("x", "y"))
|
|
self.assertEqual(samples[0].target_names, ("pressure", "velocity"))
|
|
|
|
def test_case_split_is_deterministic_and_case_level(self) -> None:
|
|
case_ids = [f"case_{index}" for index in range(10)]
|
|
|
|
first = create_case_split(case_ids, train_cases=6, val_cases=2, test_cases=2, seed=7)
|
|
second = create_case_split(case_ids, train_cases=6, val_cases=2, test_cases=2, seed=7)
|
|
|
|
self.assertEqual(first, second)
|
|
self.assertEqual(len(set(first.train_ids) & set(first.val_ids)), 0)
|
|
self.assertEqual(len(set(first.train_ids) & set(first.test_ids)), 0)
|
|
self.assertEqual(len(first.train_ids), 6)
|
|
self.assertEqual(len(first.val_ids), 2)
|
|
self.assertEqual(len(first.test_ids), 2)
|
|
|
|
def test_normalization_uses_train_split_only(self) -> None:
|
|
train_features = np.array([[0.0], [2.0]], dtype=np.float32)
|
|
train_targets = np.array([[10.0], [14.0]], dtype=np.float32)
|
|
validation_targets = np.array([[1000.0]], dtype=np.float32)
|
|
|
|
stats = compute_normalization_stats(
|
|
train_features,
|
|
train_targets,
|
|
feature_names=("x",),
|
|
target_names=("pressure",),
|
|
)
|
|
normalized_validation = normalize_targets(validation_targets, stats)
|
|
|
|
self.assertAlmostEqual(float(stats.target_mean[0]), 12.0)
|
|
self.assertAlmostEqual(float(stats.target_std[0]), 2.0)
|
|
self.assertAlmostEqual(float(normalized_validation[0, 0]), 494.0)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|