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,
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:

View file

@ -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:

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.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)

View file

@ -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()

View file

@ -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())