diff --git a/src/airfrans_frontier/training/loop.py b/src/airfrans_frontier/training/loop.py index 7288812..c565ea4 100644 --- a/src/airfrans_frontier/training/loop.py +++ b/src/airfrans_frontier/training/loop.py @@ -298,6 +298,8 @@ def train(config: TrainingConfig, *, resume_path: str | Path | None = None) -> T feature_names=bundle.feature_names, target_names=bundle.target_names, ) + raw_feature_names = _raw_coordinate_feature_names(config, bundle.feature_names) + normalization_payload = _normalization_payload(stats, raw_feature_names=raw_feature_names) writer.write_split_manifest(bundle.split.to_dict()) writer.write_json( "data_manifest.json", @@ -316,6 +318,7 @@ def train(config: TrainingConfig, *, resume_path: str | Path | None = None) -> T "points_per_case": config.data.points_per_case, "feature_names": list(bundle.feature_names), "target_names": list(bundle.target_names), + "input_normalization_policy": normalization_payload["input_policy"], "cases": [ { "case_id": sample.case_id, @@ -326,13 +329,13 @@ def train(config: TrainingConfig, *, resume_path: str | Path | None = None) -> T ], }, ) - writer.write_normalization(stats.to_dict()) + writer.write_normalization(normalization_payload) - train_features = normalize_features(bundle.train.features, stats) + train_features = normalize_features(bundle.train.features, stats, raw_feature_names=raw_feature_names) train_targets = normalize_targets(bundle.train.targets, stats) - val_features = normalize_features(bundle.val.features, stats) if bundle.val is not None else None + val_features = normalize_features(bundle.val.features, stats, raw_feature_names=raw_feature_names) if bundle.val is not None else None val_targets = normalize_targets(bundle.val.targets, stats) if bundle.val is not None else None - test_features = normalize_features(bundle.test.features, stats) if bundle.test is not None else None + test_features = normalize_features(bundle.test.features, stats, raw_feature_names=raw_feature_names) if bundle.test is not None else None test_targets = normalize_targets(bundle.test.targets, stats) if bundle.test is not None else None except Exception as exc: _write_terminal_failure_bundle( @@ -1010,13 +1013,17 @@ def _train_streaming_data( streaming.prepare() stats = streaming.load_or_compute_normalization() bundle = streaming.schema_bundle() + raw_feature_names = _raw_coordinate_feature_names(config, bundle.feature_names) + normalization_payload = _normalization_payload(stats, raw_feature_names=raw_feature_names) fast_streaming_start = ( config.data.streaming_normalization_cases is not None and config.data.streaming_normalization_cases < len(bundle.split.train_ids) ) writer.write_split_manifest(bundle.split.to_dict()) - writer.write_json("data_manifest.json", streaming.data_manifest()) - writer.write_normalization(stats.to_dict()) + data_manifest = streaming.data_manifest() + data_manifest["input_normalization_policy"] = normalization_payload["input_policy"] + writer.write_json("data_manifest.json", data_manifest) + writer.write_normalization(normalization_payload) model = _build_model(config, bundle, output_dim=bundle.train.targets.shape[1]).to(device) optimizer = torch.optim.AdamW( @@ -1622,6 +1629,22 @@ def _train_streaming_data( raise +def _raw_coordinate_feature_names(config: TrainingConfig, feature_names: tuple[str, ...]) -> tuple[str, ...]: + available = set(feature_names) + return tuple(name for name in config.model.coordinate_features if name in available) + + +def _normalization_payload(stats: NormalizationStats, *, raw_feature_names: tuple[str, ...]) -> dict[str, Any]: + payload = stats.to_dict() + payload["input_policy"] = { + "target_normalization": "standardize_all_targets", + "feature_normalization": "standardize_non_coordinate_features", + "raw_feature_names": list(raw_feature_names), + "coordinate_features_raw": True, + } + return payload + + def _autocast_context(config: TrainingConfig, device: torch.device): if config.precision.dtype == "float32" or device.type != "cuda": return torch.autocast(device_type=device.type, enabled=False) @@ -2245,6 +2268,7 @@ def _checkpoint_payload( ) -> dict[str, Any]: include_optimizer = config.checkpoint.include_optimizer_state include_rng = config.checkpoint.include_rng_state + normalization_payload = _normalization_payload(stats, raw_feature_names=_raw_coordinate_feature_names(config, bundle.feature_names)) return { "schema_version": CHECKPOINT_SCHEMA_VERSION, "run_id": os.environ.get("AIRFRANS_REMOTE_RUN_ID", config.run.name), @@ -2261,7 +2285,8 @@ def _checkpoint_payload( "scheduler_state_dict": None, "config": config.config_text, "config_hash": _config_hash(config), - "normalization": stats.to_dict(), + "normalization": normalization_payload, + "normalization_policy": normalization_payload["input_policy"], "target_names": bundle.target_names, "feature_names": bundle.feature_names, "rng_state": random.getstate() if include_rng else None, @@ -2307,8 +2332,11 @@ def _validate_resume_checkpoint( raise ValueError("checkpoint missing normalization") if "optimizer_state_dict" not in checkpoint: raise ValueError("checkpoint missing optimizer state") - if checkpoint.get("normalization") != stats.to_dict(): - raise ValueError("checkpoint normalization does not match dataset") + expected_normalization = _normalization_payload(stats, raw_feature_names=_raw_coordinate_feature_names(config, bundle.feature_names)) + if checkpoint.get("normalization") != expected_normalization: + raise ValueError("checkpoint normalization does not match dataset or input policy") + if checkpoint.get("normalization_policy") != expected_normalization["input_policy"]: + raise ValueError("checkpoint normalization policy does not match config") def _restore_rng_state(checkpoint: dict[str, Any], rng: np.random.Generator) -> None: diff --git a/src/airfrans_frontier/training/normalize.py b/src/airfrans_frontier/training/normalize.py index 3528b68..4c4bc01 100644 --- a/src/airfrans_frontier/training/normalize.py +++ b/src/airfrans_frontier/training/normalize.py @@ -2,6 +2,7 @@ from __future__ import annotations import json from dataclasses import dataclass +from collections.abc import Sequence from pathlib import Path import numpy as np @@ -74,11 +75,25 @@ def compute_normalization_stats( ) -def normalize_features(features: FloatArray, stats: NormalizationStats) -> FloatArray: +def normalize_features( + features: FloatArray, + stats: NormalizationStats, + *, + raw_feature_names: Sequence[str] = (), +) -> FloatArray: _validate_matrix(features, "features") if features.shape[1] != stats.feature_mean.shape[0]: raise ValueError("Feature width does not match normalization stats") - return np.ascontiguousarray((features - stats.feature_mean) / stats.feature_std, dtype=np.float32) + normalized = (features - stats.feature_mean) / stats.feature_std + if raw_feature_names: + name_to_index = {name: index for index, name in enumerate(stats.feature_names)} + missing = [name for name in raw_feature_names if name not in name_to_index] + if missing: + raise ValueError(f"Raw feature names are missing from normalization stats: {missing}") + for name in raw_feature_names: + index = name_to_index[name] + normalized[:, index] = features[:, index] + return np.ascontiguousarray(normalized, dtype=np.float32) def normalize_targets(targets: FloatArray, stats: NormalizationStats) -> FloatArray: diff --git a/src/airfrans_frontier/training/streaming_data.py b/src/airfrans_frontier/training/streaming_data.py index b8e1c59..83d7967 100644 --- a/src/airfrans_frontier/training/streaming_data.py +++ b/src/airfrans_frontier/training/streaming_data.py @@ -26,7 +26,7 @@ from airfrans_frontier.raw.process import process_raw_case_to_npz from airfrans_frontier.training.config import DataConfig, TrainingConfig from airfrans_frontier.training.data import CaseSplit, DatasetBundle, SimulationSample, SplitArrays, create_case_split, load_simulation_npz from airfrans_frontier.training.hf_upload import _retry_after_seconds -from airfrans_frontier.training.normalize import NormalizationStats, load_normalization_stats +from airfrans_frontier.training.normalize import NormalizationStats, load_normalization_stats, normalize_features FloatArray = NDArray[np.float32] IntArray = NDArray[np.int64] @@ -747,6 +747,7 @@ class StreamingTrainingData: self.split: CaseSplit | None = None self.feature_names: tuple[str, ...] | None = None self.target_names: tuple[str, ...] | None = None + self.raw_feature_names: tuple[str, ...] = () self.stats: NormalizationStats | None = None self._split_case_ids: dict[str, tuple[str, ...]] = {} self._sampling_specs: dict[str, dict[str, SamplingSpec]] = {"train": {}, "val": {}, "test": {}} @@ -793,6 +794,8 @@ class StreamingTrainingData: sample = self.cache.ensure_case(first_case) self.feature_names = sample.feature_names self.target_names = sample.target_names + available = set(sample.feature_names) + self.raw_feature_names = tuple(name for name in self.config.model.coordinate_features if name in available) self._upload_queue.enqueue(sample.source_path) self.cache.release_case(first_case, consumed=False) @@ -989,7 +992,7 @@ class StreamingTrainingData: source_indices = _source_indices_for_local(spec, local_indices) selected_features = sample.features[source_indices] selected_targets = sample.targets[source_indices] - features = ((selected_features - self.stats.feature_mean) / self.stats.feature_std).astype(np.float32, copy=False) + features = normalize_features(selected_features, self.stats, raw_feature_names=self.raw_feature_names) targets = ((selected_targets - self.stats.target_mean) / self.stats.target_std).astype(np.float32, copy=False) return np.ascontiguousarray(features, dtype=np.float32), np.ascontiguousarray(targets, dtype=np.float32) @@ -1045,7 +1048,7 @@ class StreamingTrainingData: source_indices = _source_indices_for_local(spec, local_indices) selected_features = sample.features[source_indices] selected_targets = sample.targets[source_indices] - features[mask] = ((selected_features - self.stats.feature_mean) / self.stats.feature_std).astype(np.float32, copy=False) + features[mask] = normalize_features(selected_features, self.stats, raw_feature_names=self.raw_feature_names) targets[mask] = ((selected_targets - self.stats.target_mean) / self.stats.target_std).astype(np.float32, copy=False) self._upload_queue.enqueue(sample.source_path) self.cache.release_case(case_id) @@ -1071,7 +1074,7 @@ class StreamingTrainingData: for start in range(0, spec.count, batch_size): stop = min(start + batch_size, spec.count) source_indices = _source_indices_for_local(spec, np.arange(start, stop, dtype=np.int64)) - features = ((sample.features[source_indices] - self.stats.feature_mean) / self.stats.feature_std).astype(np.float32, copy=False) + features = normalize_features(sample.features[source_indices], self.stats, raw_feature_names=self.raw_feature_names) targets = ((sample.targets[source_indices] - self.stats.target_mean) / self.stats.target_std).astype(np.float32, copy=False) yield np.ascontiguousarray(features, dtype=np.float32), np.ascontiguousarray(targets, dtype=np.float32) self._upload_queue.enqueue(sample.source_path) diff --git a/tests/test_training_data.py b/tests/test_training_data.py index 47145e7..82f2379 100644 --- a/tests/test_training_data.py +++ b/tests/test_training_data.py @@ -11,7 +11,7 @@ remove_pythonpath_entries() import numpy as np from airfrans_frontier.training.data import build_dataset_bundle, create_case_split, load_processed_dataset, load_simulation_npz -from airfrans_frontier.training.normalize import compute_normalization_stats, normalize_targets +from airfrans_frontier.training.normalize import compute_normalization_stats, normalize_features, normalize_targets def write_case(path: Path, offset: float = 0.0) -> None: @@ -110,6 +110,28 @@ class TrainingDataTests(unittest.TestCase): self.assertAlmostEqual(float(stats.target_std[0]), 2.0) self.assertAlmostEqual(float(normalized_validation[0, 0]), 494.0) + def test_feature_normalization_preserves_raw_coordinate_columns(self) -> None: + features = np.array( + [ + [10.0, 100.0, 1.0], + [20.0, 300.0, 3.0], + ], + dtype=np.float32, + ) + targets = np.array([[1.0], [3.0]], dtype=np.float32) + stats = compute_normalization_stats( + features, + targets, + feature_names=("x", "aoa", "sdf"), + target_names=("pressure",), + ) + + normalized = normalize_features(features, stats, raw_feature_names=("x", "sdf")) + + np.testing.assert_allclose(normalized[:, 0], features[:, 0]) + np.testing.assert_allclose(normalized[:, 2], features[:, 2]) + np.testing.assert_allclose(normalized[:, 1], np.array([-1.0, 1.0], dtype=np.float32)) + if __name__ == "__main__": unittest.main() diff --git a/tests/test_training_loop.py b/tests/test_training_loop.py index 4dc10f5..f094451 100644 --- a/tests/test_training_loop.py +++ b/tests/test_training_loop.py @@ -309,8 +309,12 @@ class TrainingLoopTests(unittest.TestCase): "torch_rng_state", "batch_rng_state", "scheduler_state_dict", + "normalization_policy", ): self.assertIn(key, checkpoint) + self.assertEqual(checkpoint["normalization_policy"]["raw_feature_names"], ["x", "y", "sdf"]) + normalization_manifest = json.loads((run_dir / "normalization.json").read_text()) + self.assertEqual(normalization_manifest["input_policy"]["raw_feature_names"], ["x", "y", "sdf"]) self.assertEqual(list(run_dir.glob("*.tmp")), []) self.assertTrue((run_dir / "artifact_manifest.json").is_file()) self.assertTrue((run_dir / "checksums.txt").is_file())