airfRANS-model-exploration/src/airfrans_frontier/training/config.py

208 lines
6.1 KiB
Python
Raw Normal View History

2026-07-21 08:32:30 +00:00
from __future__ import annotations
import tomllib
from dataclasses import dataclass
from pathlib import Path
from typing import Any
@dataclass(frozen=True)
class RunConfig:
name: str
seed: int
artifact_dir: Path
@dataclass(frozen=True)
class DataConfig:
root: Path
train_cases: int
val_cases: int
test_cases: int
points_per_case: int
batch_size: int
@dataclass(frozen=True)
class ModelConfig:
type: str
hidden_width: int
depth: int
activation: str
@dataclass(frozen=True)
class OptimConfig:
lr: float
weight_decay: float
steps: int
log_interval: int | None = None
@dataclass(frozen=True)
class DeviceConfig:
type: str
allow_cpu_fallback: bool
benchmark_kernels: bool
@dataclass(frozen=True)
class LossConfig:
type: str
@dataclass(frozen=True)
class TrainingConfig:
path: Path
config_text: str
run: RunConfig
data: DataConfig
model: ModelConfig
optim: OptimConfig
device: DeviceConfig
loss: LossConfig
_REQUIRED_SECTIONS = ("run", "data", "model", "optim", "device", "loss")
def load_training_config(path: str | Path) -> TrainingConfig:
config_path = Path(path).expanduser()
if not config_path.exists():
raise FileNotFoundError(f"Training config not found: {config_path}")
if not config_path.is_file():
raise ValueError(f"Training config is not a file: {config_path}")
text = config_path.read_text()
try:
raw = tomllib.loads(text)
except tomllib.TOMLDecodeError as exc:
raise ValueError(f"Invalid TOML config {config_path}: {exc}") from exc
for section_name in _REQUIRED_SECTIONS:
if section_name not in raw or not isinstance(raw[section_name], dict):
raise ValueError(f"Training config missing [{section_name}] section")
run_raw = raw["run"]
data_raw = raw["data"]
model_raw = raw["model"]
optim_raw = raw["optim"]
device_raw = raw["device"]
loss_raw = raw["loss"]
run = RunConfig(
name=_string(run_raw, "name"),
seed=_integer(run_raw, "seed", minimum=0),
artifact_dir=_path(run_raw, "artifact_dir"),
)
data = DataConfig(
root=_path(data_raw, "root"),
train_cases=_integer(data_raw, "train_cases", minimum=1),
val_cases=_integer(data_raw, "val_cases", minimum=0),
test_cases=_integer(data_raw, "test_cases", minimum=0),
points_per_case=_integer(data_raw, "points_per_case", minimum=1),
batch_size=_integer(data_raw, "batch_size", minimum=1),
)
model = ModelConfig(
type=_choice(_string(model_raw, "type"), {"mlp"}, "model.type"),
hidden_width=_integer(model_raw, "hidden_width", minimum=1),
depth=_integer(model_raw, "depth", minimum=1),
activation=_choice(_string(model_raw, "activation").lower(), {"gelu", "relu", "silu", "tanh"}, "model.activation"),
)
optim = OptimConfig(
lr=_number(optim_raw, "lr", minimum=0.0, exclusive_minimum=True),
weight_decay=_number(optim_raw, "weight_decay", minimum=0.0),
steps=_integer(optim_raw, "steps", minimum=1),
log_interval=_optional_integer(optim_raw, "log_interval", minimum=1),
)
device = DeviceConfig(
type=_choice(_string(device_raw, "type").lower(), {"cuda", "cpu", "auto"}, "device.type"),
allow_cpu_fallback=_boolean(device_raw, "allow_cpu_fallback"),
benchmark_kernels=_boolean(device_raw, "benchmark_kernels"),
)
loss = LossConfig(type=_choice(_string(loss_raw, "type"), {"normalized_mse"}, "loss.type"))
requested_cases = data.train_cases + data.val_cases + data.test_cases
if requested_cases <= 0:
raise ValueError("Training config must request at least one case")
return TrainingConfig(
path=config_path,
config_text=text,
run=run,
data=data,
model=model,
optim=optim,
device=device,
loss=loss,
)
def _path(section: dict[str, Any], key: str) -> Path:
value = _string(section, key)
path = Path(value).expanduser()
if path.is_absolute():
return path
return Path.cwd() / path
def _string(section: dict[str, Any], key: str) -> str:
value = _required(section, key)
if not isinstance(value, str) or not value:
raise ValueError(f"Expected non-empty string for {key}")
return value
def _integer(section: dict[str, Any], key: str, *, minimum: int | None = None) -> int:
value = _required(section, key)
if isinstance(value, bool) or not isinstance(value, int):
raise ValueError(f"Expected integer for {key}")
if minimum is not None and value < minimum:
raise ValueError(f"Expected {key} >= {minimum}")
return value
def _optional_integer(section: dict[str, Any], key: str, *, minimum: int | None = None) -> int | None:
if key not in section:
return None
return _integer(section, key, minimum=minimum)
def _number(
section: dict[str, Any],
key: str,
*,
minimum: float | None = None,
exclusive_minimum: bool = False,
) -> float:
value = _required(section, key)
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise ValueError(f"Expected number for {key}")
result = float(value)
if minimum is not None:
if exclusive_minimum and result <= minimum:
raise ValueError(f"Expected {key} > {minimum}")
if not exclusive_minimum and result < minimum:
raise ValueError(f"Expected {key} >= {minimum}")
return result
def _boolean(section: dict[str, Any], key: str) -> bool:
value = _required(section, key)
if not isinstance(value, bool):
raise ValueError(f"Expected boolean for {key}")
return value
def _choice(value: str, allowed: set[str], key: str) -> str:
if value not in allowed:
allowed_text = ", ".join(sorted(allowed))
raise ValueError(f"Expected {key} to be one of: {allowed_text}")
return value
def _required(section: dict[str, Any], key: str) -> Any:
if key not in section:
raise ValueError(f"Training config missing key: {key}")
return section[key]