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,
|
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:
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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())
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue