fix: normalization does not apply to coordinate parameters
This commit is contained in:
parent
d2b102cafb
commit
2065313f44
5 changed files with 88 additions and 16 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
Loading…
Reference in a new issue