fix: normalization does not apply to coordinate parameters

This commit is contained in:
Zachery Aaron Shores-Chmielewski 2026-07-29 14:14:25 +04:00
parent d2b102cafb
commit 2065313f44
5 changed files with 88 additions and 16 deletions

View file

@ -298,6 +298,8 @@ def train(config: TrainingConfig, *, resume_path: str | Path | None = None) -> T
feature_names=bundle.feature_names, feature_names=bundle.feature_names,
target_names=bundle.target_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_split_manifest(bundle.split.to_dict())
writer.write_json( writer.write_json(
"data_manifest.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, "points_per_case": config.data.points_per_case,
"feature_names": list(bundle.feature_names), "feature_names": list(bundle.feature_names),
"target_names": list(bundle.target_names), "target_names": list(bundle.target_names),
"input_normalization_policy": normalization_payload["input_policy"],
"cases": [ "cases": [
{ {
"case_id": sample.case_id, "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) 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 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 test_targets = normalize_targets(bundle.test.targets, stats) if bundle.test is not None else None
except Exception as exc: except Exception as exc:
_write_terminal_failure_bundle( _write_terminal_failure_bundle(
@ -1010,13 +1013,17 @@ def _train_streaming_data(
streaming.prepare() streaming.prepare()
stats = streaming.load_or_compute_normalization() stats = streaming.load_or_compute_normalization()
bundle = streaming.schema_bundle() 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 = ( fast_streaming_start = (
config.data.streaming_normalization_cases is not None config.data.streaming_normalization_cases is not None
and config.data.streaming_normalization_cases < len(bundle.split.train_ids) and config.data.streaming_normalization_cases < len(bundle.split.train_ids)
) )
writer.write_split_manifest(bundle.split.to_dict()) writer.write_split_manifest(bundle.split.to_dict())
writer.write_json("data_manifest.json", streaming.data_manifest()) data_manifest = streaming.data_manifest()
writer.write_normalization(stats.to_dict()) 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) model = _build_model(config, bundle, output_dim=bundle.train.targets.shape[1]).to(device)
optimizer = torch.optim.AdamW( optimizer = torch.optim.AdamW(
@ -1622,6 +1629,22 @@ def _train_streaming_data(
raise 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): def _autocast_context(config: TrainingConfig, device: torch.device):
if config.precision.dtype == "float32" or device.type != "cuda": if config.precision.dtype == "float32" or device.type != "cuda":
return torch.autocast(device_type=device.type, enabled=False) return torch.autocast(device_type=device.type, enabled=False)
@ -2245,6 +2268,7 @@ def _checkpoint_payload(
) -> dict[str, Any]: ) -> dict[str, Any]:
include_optimizer = config.checkpoint.include_optimizer_state include_optimizer = config.checkpoint.include_optimizer_state
include_rng = config.checkpoint.include_rng_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 { return {
"schema_version": CHECKPOINT_SCHEMA_VERSION, "schema_version": CHECKPOINT_SCHEMA_VERSION,
"run_id": os.environ.get("AIRFRANS_REMOTE_RUN_ID", config.run.name), "run_id": os.environ.get("AIRFRANS_REMOTE_RUN_ID", config.run.name),
@ -2261,7 +2285,8 @@ def _checkpoint_payload(
"scheduler_state_dict": None, "scheduler_state_dict": None,
"config": config.config_text, "config": config.config_text,
"config_hash": _config_hash(config), "config_hash": _config_hash(config),
"normalization": stats.to_dict(), "normalization": normalization_payload,
"normalization_policy": normalization_payload["input_policy"],
"target_names": bundle.target_names, "target_names": bundle.target_names,
"feature_names": bundle.feature_names, "feature_names": bundle.feature_names,
"rng_state": random.getstate() if include_rng else None, "rng_state": random.getstate() if include_rng else None,
@ -2307,8 +2332,11 @@ def _validate_resume_checkpoint(
raise ValueError("checkpoint missing normalization") raise ValueError("checkpoint missing normalization")
if "optimizer_state_dict" not in checkpoint: if "optimizer_state_dict" not in checkpoint:
raise ValueError("checkpoint missing optimizer state") raise ValueError("checkpoint missing optimizer state")
if checkpoint.get("normalization") != stats.to_dict(): expected_normalization = _normalization_payload(stats, raw_feature_names=_raw_coordinate_feature_names(config, bundle.feature_names))
raise ValueError("checkpoint normalization does not match dataset") 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: def _restore_rng_state(checkpoint: dict[str, Any], rng: np.random.Generator) -> None:

View file

@ -2,6 +2,7 @@ from __future__ import annotations
import json import json
from dataclasses import dataclass from dataclasses import dataclass
from collections.abc import Sequence
from pathlib import Path from pathlib import Path
import numpy as np 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") _validate_matrix(features, "features")
if features.shape[1] != stats.feature_mean.shape[0]: if features.shape[1] != stats.feature_mean.shape[0]:
raise ValueError("Feature width does not match normalization stats") 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: def normalize_targets(targets: FloatArray, stats: NormalizationStats) -> FloatArray:

View file

@ -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.config import DataConfig, TrainingConfig
from airfrans_frontier.training.data import CaseSplit, DatasetBundle, SimulationSample, SplitArrays, create_case_split, load_simulation_npz 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.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] FloatArray = NDArray[np.float32]
IntArray = NDArray[np.int64] IntArray = NDArray[np.int64]
@ -747,6 +747,7 @@ class StreamingTrainingData:
self.split: CaseSplit | None = None self.split: CaseSplit | None = None
self.feature_names: tuple[str, ...] | None = None self.feature_names: tuple[str, ...] | None = None
self.target_names: tuple[str, ...] | None = None self.target_names: tuple[str, ...] | None = None
self.raw_feature_names: tuple[str, ...] = ()
self.stats: NormalizationStats | None = None self.stats: NormalizationStats | None = None
self._split_case_ids: dict[str, tuple[str, ...]] = {} self._split_case_ids: dict[str, tuple[str, ...]] = {}
self._sampling_specs: dict[str, dict[str, SamplingSpec]] = {"train": {}, "val": {}, "test": {}} self._sampling_specs: dict[str, dict[str, SamplingSpec]] = {"train": {}, "val": {}, "test": {}}
@ -793,6 +794,8 @@ class StreamingTrainingData:
sample = self.cache.ensure_case(first_case) sample = self.cache.ensure_case(first_case)
self.feature_names = sample.feature_names self.feature_names = sample.feature_names
self.target_names = sample.target_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._upload_queue.enqueue(sample.source_path)
self.cache.release_case(first_case, consumed=False) self.cache.release_case(first_case, consumed=False)
@ -989,7 +992,7 @@ class StreamingTrainingData:
source_indices = _source_indices_for_local(spec, local_indices) source_indices = _source_indices_for_local(spec, local_indices)
selected_features = sample.features[source_indices] selected_features = sample.features[source_indices]
selected_targets = sample.targets[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) 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) 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) source_indices = _source_indices_for_local(spec, local_indices)
selected_features = sample.features[source_indices] selected_features = sample.features[source_indices]
selected_targets = sample.targets[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) 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._upload_queue.enqueue(sample.source_path)
self.cache.release_case(case_id) self.cache.release_case(case_id)
@ -1071,7 +1074,7 @@ class StreamingTrainingData:
for start in range(0, spec.count, batch_size): for start in range(0, spec.count, batch_size):
stop = min(start + batch_size, spec.count) stop = min(start + batch_size, spec.count)
source_indices = _source_indices_for_local(spec, np.arange(start, stop, dtype=np.int64)) 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) 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) yield np.ascontiguousarray(features, dtype=np.float32), np.ascontiguousarray(targets, dtype=np.float32)
self._upload_queue.enqueue(sample.source_path) self._upload_queue.enqueue(sample.source_path)

View file

@ -11,7 +11,7 @@ remove_pythonpath_entries()
import numpy as np 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.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: 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(stats.target_std[0]), 2.0)
self.assertAlmostEqual(float(normalized_validation[0, 0]), 494.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__": if __name__ == "__main__":
unittest.main() unittest.main()

View file

@ -309,8 +309,12 @@ class TrainingLoopTests(unittest.TestCase):
"torch_rng_state", "torch_rng_state",
"batch_rng_state", "batch_rng_state",
"scheduler_state_dict", "scheduler_state_dict",
"normalization_policy",
): ):
self.assertIn(key, checkpoint) 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.assertEqual(list(run_dir.glob("*.tmp")), [])
self.assertTrue((run_dir / "artifact_manifest.json").is_file()) self.assertTrue((run_dir / "artifact_manifest.json").is_file())
self.assertTrue((run_dir / "checksums.txt").is_file()) self.assertTrue((run_dir / "checksums.txt").is_file())