Back to report index

linebitmapimagerasterizer (ML) b3707fd: AI3D-379 Pydantic config models via iolabs-common ConfigModel

Miroslav Simko <ms@iolabs.ch> 2026-09-02T08:47:59+02:00

Commit #81 · 27 snippets

 CLAUDE.md                    |   6 +-
 pyproject.toml               |  12 ++-
 scripts/train.py             |  49 ++++++---
 src/train/config.py          | 248 +++++++++++++++++++++++--------------------
 test/test_ml_harness.py      |   2 +-
 test/test_train_overrides.py |  85 ++++++++++++---
 6 files changed, 257 insertions(+), 145 deletions(-)
Importance #1: src/train/config.py @@ -1,95 +1,99 @@
1"""YAML-backed configuration for the training harness.1"""YAML-backed configuration for the training harness.
22
3Plain dataclasses + yaml, no config framework. Unknown keys raise, so config3Pydantic models on `iolabs.common.config_loader.ConfigModel` + yaml. Unknown keys
4typos fail fast instead of silently training with defaults.4raise, so config typos fail fast instead of silently training with defaults, and
5values are coerced by the shared fleet matrix.
6
7Adding a config key = adding one field with its default to the model below.
8Instances are frozen: derive a changed config with ``model_copy(update=...)``.
5"""9"""
6from dataclasses import dataclass, field, fields10import logging
11from collections.abc import Mapping
7from pathlib import Path12from pathlib import Path
8from typing import Any, TypeVar13from typing import Any, Literal
914
15import pydantic
10import yaml16import yaml
17from iolabs.common import config_loader
18
19logger = logging.getLogger(__name__)
1120
12T = TypeVar("T")21_STROKE_KINDS = frozenset({"solid", "dashed"})
22_PAIR_LIST_FIELDS = ("pairs", "val_pairs", "test_pairs")
23_PATH_FIELD_SUFFIXES = ("path", "paths", "dir", "dirs", "root", "roots")
1324
1425
15def _build(cls: type[T], data: dict[str, Any] | None, where: str) -> T:26class ConfigError(config_loader.ConfigError):
16 data = dict(data or {})27 """Raised when a harness config holds unknown keys or invalid values."""
17 known = {f.name for f in fields(cls)}
18 unknown = sorted(set(data) - known)
19 if unknown:
20 raise KeyError(f"unknown key(s) {unknown} in config section {where!r}; "
21 f"known keys: {sorted(known)}")
22 return cls(**data)
2328
2429
25@dataclass30class PairSpec(config_loader.ConfigModel):
26class PairSpec:
27 """One images-dir / masks-dir pair (see src.dataset.index_tile_pairs)."""31 """One images-dir / masks-dir pair (see src.dataset.index_tile_pairs)."""
28 images: str32 images: str
29 masks: str33 masks: str
3034
3135
32@dataclass36class DataConfig(config_loader.ConfigModel):
33class DataConfig:37 """Tile corpus, split, crop sampling and label rasterization knobs."""
34 pairs: list = field(default_factory=list)38 pairs: list[PairSpec] = []
35 # Optional explicit, pre-split directories (e.g. the symlink folders under39 # Optional explicit, pre-split directories (e.g. the symlink folders under
36 # data/02_processed/<ds>/{train,val,test} built by scripts/build_processed_40 # data/02_processed/<ds>/{train,val,test} built by scripts/build_processed_
37 # splits.py). When val_pairs is set, `pairs` is used in full as the training41 # splits.py). When val_pairs is set, `pairs` is used in full as the training
38 # set and is NOT re-split — val_fraction is ignored — so a geographic split42 # set and is NOT re-split — val_fraction is ignored — so a geographic split
39 # materialised on disk is honoured verbatim. test_pairs feeds test_dataloader.43 # materialised on disk is honoured verbatim. test_pairs feeds test_dataloader.
40 val_pairs: list = field(default_factory=list)44 val_pairs: list[PairSpec] = []
41 test_pairs: list = field(default_factory=list)45 test_pairs: list[PairSpec] = []
42 crop_size: int = 51246 crop_size: int = pydantic.Field(default=512, gt=0)
43 batch_size: int = 1647 batch_size: int = pydantic.Field(default=16, gt=0)
44 num_workers: int = 448 num_workers: int = pydantic.Field(default=4, ge=0)
45 val_fraction: float = 0.1549 val_fraction: float = pydantic.Field(default=0.15, ge=0.0, le=1.0)
46 crops_per_tile: int = 450 crops_per_tile: int = pydantic.Field(default=4, gt=0)
47 pos_crop_prob: float = 0.751 pos_crop_prob: float = pydantic.Field(default=0.7, ge=0.0, le=1.0)
48 min_valid_fraction: float = 0.1052 min_valid_fraction: float = pydantic.Field(default=0.10, ge=0.0, le=1.0)
49 augment: bool = True53 augment: bool = True
50 label_source: str = "rendered" # rendered *_lines.png | vector *_vectors.json | review54 # rendered *_lines.png | vector *_vectors.json | reviewer-confirmed tiles
51 label_stroke_px: int | float | dict[str, float] = 4 # scalar or {solid, dashed}55 label_source: Literal["rendered", "vector", "review"] = "rendered"
52 review_statuses: list = field(default_factory=lambda: ["ok"]) # label_source: review56 # scalar or {solid, dashed}; int stays int so reports render "5", not "5.0"
5357 label_stroke_px: int | float | dict[str, int | float] = 4
5458 review_statuses: list[str] = ["ok"] # label_source: review
55_STROKE_KINDS = frozenset({"solid", "dashed"})
5659
5760 @pydantic.field_validator("label_stroke_px", mode="before")
58def _parse_label_stroke_px(value: Any) -> int | float | dict[str, float]:61 @classmethod
59 """Scalar > 0, or exactly ``{solid, dashed}`` with positive numeric values."""62 def _check_label_stroke_px(cls, value: Any) -> Any:
60 if isinstance(value, (int, float)) and not isinstance(value, bool):63 """Scalar > 0, or exactly ``{solid, dashed}`` with positive numeric values."""
61 if value <= 0:64 if isinstance(value, (int, float)) and not isinstance(value, bool):
62 raise ValueError(f"data.label_stroke_px must be > 0, got {value}")65 if value <= 0:
63 return value66 raise ValueError(f"data.label_stroke_px must be > 0, got {value}")
64 if isinstance(value, dict):67 return value
65 keys = set(value)68 if isinstance(value, Mapping):
66 if keys != _STROKE_KINDS:69 keys = set(value)
67 raise ValueError(70 if keys != _STROKE_KINDS:
68 f"data.label_stroke_px mapping must have exactly the keys "
69 f"{sorted(_STROKE_KINDS)}, got {sorted(keys)}")
70 out: dict[str, float] = {}
71 for kind, width in value.items():
72 if (not isinstance(width, (int, float)) or isinstance(width, bool)
73 or width <= 0):
74 raise ValueError(71 raise ValueError(
75 f"data.label_stroke_px[{kind!r}] must be a positive number, "72 f"data.label_stroke_px mapping must have exactly the keys "
76 f"got {width!r}")73 f"{sorted(_STROKE_KINDS)}, got {sorted(keys)}")
77 out[kind] = width74 for kind, width in value.items():
78 return out75 if (not isinstance(width, (int, float)) or isinstance(width, bool)
79 raise ValueError(76 or width <= 0):
80 f"data.label_stroke_px must be a positive number or a "77 raise ValueError(
81 f"{{solid, dashed}} mapping, got {type(value).__name__}: {value!r}")78 f"data.label_stroke_px[{kind!r}] must be a positive number, "
79 f"got {width!r}")
80 return dict(value)
81 raise ValueError(
82 f"data.label_stroke_px must be a positive number or a "
83 f"{{solid, dashed}} mapping, got {type(value).__name__}: {value!r}")
8284
8385
84def _reroot_path(value: Any, root: Path) -> Any:86def _reroot_path(value: Any, root: Path) -> Any:
87 """Return *value* as an absolute path string, joined onto *root* if relative."""
85 path = Path(value)88 path = Path(value)
86 if path.is_absolute():89 if path.is_absolute():
87 return str(path)90 return str(path)
88 return str(root / path)91 return str(root / path)
8992
9093
91def _reroot_path_value(value: Any, root: Path) -> Any:94def _reroot_path_value(value: Any, root: Path) -> Any:
95 """Re-root every path-like leaf of a scalar/list/tuple/dict value."""
92 if isinstance(value, (str, Path)):96 if isinstance(value, (str, Path)):
93 return _reroot_path(value, root)97 return _reroot_path(value, root)
94 if isinstance(value, list):98 if isinstance(value, list):
95 return [_reroot_path_value(item, root) for item in value]99 return [_reroot_path_value(item, root) for item in value]
Importance #2: src/train/config.py @@ -100,85 +104,103 @@
100 return value104 return value
101105
102106
103def reroot_data_paths(data_cfg: DataConfig, root: str | Path) -> DataConfig:107def reroot_data_paths(data_cfg: DataConfig, root: str | Path) -> DataConfig:
104 """Re-root relative dataset paths in a DataConfig onto root."""108 """Re-root relative dataset paths in a DataConfig onto root.
109
110 Config models are frozen, so this returns an updated copy instead of
111 mutating ``data_cfg`` in place.
112
113 Args:
114 data_cfg: The data section to re-root; never mutated.
115 root: Directory relative paths are joined onto. Absolute paths are kept.
116
117 Returns:
118 A copy of ``data_cfg`` with the pair lists and every path-like field
119 (name ending in path/paths/dir/dirs/root/roots) made absolute.
120 """
105 root = Path(root)121 root = Path(root)
106 for pair_list_name in ("pairs", "val_pairs", "test_pairs"):122 updates: dict[str, Any] = {}
107 for pair in getattr(data_cfg, pair_list_name):123 for name in _PAIR_LIST_FIELDS:
108 pair.images = _reroot_path(pair.images, root)124 specs = getattr(data_cfg, name)
109 pair.masks = _reroot_path(pair.masks, root)125 if specs:
110126 updates[name] = [
111 path_field_suffixes = ("path", "paths", "dir", "dirs", "root", "roots")127 spec.model_copy(update={
112 for field_info in fields(data_cfg):128 "images": _reroot_path(spec.images, root),
113 name = field_info.name129 "masks": _reroot_path(spec.masks, root)})
114 if name in {"pairs", "val_pairs", "test_pairs"}:130 for spec in specs]
131 for name in type(data_cfg).model_fields:
132 if name in _PAIR_LIST_FIELDS or not name.endswith(_PATH_FIELD_SUFFIXES):
115 continue133 continue
116 if name.endswith(path_field_suffixes):134 updates[name] = _reroot_path_value(getattr(data_cfg, name), root)
117 setattr(data_cfg, name, _reroot_path_value(getattr(data_cfg, name), root))135 return data_cfg.model_copy(update=updates)
118 return data_cfg
119136
120137
121@dataclass138class ModelConfig(config_loader.ConfigModel):
122class ModelConfig:139 """Segmentation-models-pytorch architecture/encoder selection."""
123 name: str = "unet"140 name: str = "unet"
124 encoder_name: str = "resnet18"141 encoder_name: str = "resnet18"
125 encoder_weights: str | None = "imagenet" # None = train from scratch142 encoder_weights: str | None = "imagenet" # None = train from scratch
126 in_channels: int = 1143 in_channels: int = pydantic.Field(default=1, gt=0)
127 num_classes: int = 3144 num_classes: int = pydantic.Field(default=3, gt=0)
128 extra: dict = field(default_factory=dict) # passed through to the model factory145 extra: dict[str, Any] = {} # passed through to the model factory
129146
130147
131@dataclass148class LossConfig(config_loader.ConfigModel):
132class LossConfig:149 """Loss selection by registry name plus factory keyword arguments."""
133 name: str = "dice_focal"150 name: str = "dice_focal"
134 args: dict = field(default_factory=dict)151 args: dict[str, Any] = {}
135152
136153
137@dataclass154class TrainerConfig(config_loader.ConfigModel):
138class TrainerConfig:155 """Lightning trainer, logger, and callback knobs."""
139 max_epochs: int = -1 # -1 = no cap; early stopping ends training instead156 max_epochs: int = pydantic.Field(default=-1, ge=-1) # -1 = no cap
140 lr: float = 3.0e-4157 lr: float = pydantic.Field(default=3.0e-4, gt=0)
141 weight_decay: float = 1.0e-4158 weight_decay: float = pydantic.Field(default=1.0e-4, ge=0)
142 precision: str = "auto" # auto -> 16-mixed on CUDA, 32-true on CPU159 precision: str = "auto" # auto -> 16-mixed on CUDA, 32-true on CPU
143 accumulate_grad_batches: int = 1 # match effective batch across experiments160 accumulate_grad_batches: int = pydantic.Field(default=1, ge=1)
144 accelerator: str = "auto"161 accelerator: str = "auto"
145 devices: int | str = 1162 devices: int | str = 1
146 viz_every_n_epochs: int = 2163 viz_every_n_epochs: int = pydantic.Field(default=2, ge=0)
147 viz_samples: int = 4164 viz_samples: int = pydantic.Field(default=4, ge=0)
148 monitor: str = "val/f1_mean_fg"165 monitor: str = "val/f1_mean_fg"
149 monitor_mode: str = "max"166 monitor_mode: Literal["max", "min"] = "max"
150 early_stop_monitor: str = "val/loss" # stop when this stops improving167 early_stop_monitor: str = "val/loss" # stop when this stops improving
151 early_stop_mode: str = "min"168 early_stop_mode: Literal["min", "max"] = "min"
152 early_stop_patience: int = 4 # epochs without improvement before stopping; 0 disables169 # epochs without improvement before stopping; 0 disables
170 early_stop_patience: int = pydantic.Field(default=4, ge=0)
153 log_dir: str = "runs"171 log_dir: str = "runs"
154 log_every_n_steps: int = 10172 log_every_n_steps: int = pydantic.Field(default=10, ge=1)
155173
156174
157@dataclass175class HarnessConfig(config_loader.ConfigModel):
158class HarnessConfig:176 """Top-level training config: one YAML file, one instance."""
159 experiment: str = "experiment"177 experiment: str = "experiment"
160 seed: int = 1337178 seed: int = 1337
161 data: DataConfig = field(default_factory=DataConfig)179 data: DataConfig = DataConfig()
162 model: ModelConfig = field(default_factory=ModelConfig)180 model: ModelConfig = ModelConfig()
163 loss: LossConfig = field(default_factory=LossConfig)181 loss: LossConfig = LossConfig()
164 train: TrainerConfig = field(default_factory=TrainerConfig)182 train: TrainerConfig = TrainerConfig()
165183
166 @classmethod184 @classmethod
167 def from_yaml(cls, path: str | Path) -> "HarnessConfig":185 def from_yaml(cls, path: str | Path) -> "HarnessConfig":
168 raw = yaml.safe_load(Path(path).read_text()) or {}186 """Loads and validates a harness YAML config.
169 data = _build(DataConfig, raw.pop("data", {}), "data")187
170 data.pairs = [_build(PairSpec, p, "data.pairs[]") for p in data.pairs]188 Args:
171 data.val_pairs = [_build(PairSpec, p, "data.val_pairs[]") for p in data.val_pairs]189 path: Path of the YAML file, read as UTF-8.
172 data.test_pairs = [_build(PairSpec, p, "data.test_pairs[]") for p in data.test_pairs]190
173 data.label_stroke_px = _parse_label_stroke_px(data.label_stroke_px)191 Returns:
174 cfg = cls(192 The validated, frozen config.
175 experiment=raw.pop("experiment", cls.experiment),193
176 seed=raw.pop("seed", cls.seed),194 Raises:
177 data=data,195 FileNotFoundError: If ``path`` does not exist.
178 model=_build(ModelConfig, raw.pop("model", {}), "model"),196 ConfigError: If the document is not a mapping, holds an unknown key,
179 loss=_build(LossConfig, raw.pop("loss", {}), "loss"),197 or holds a value invalid for its field. Derives from
180 train=_build(TrainerConfig, raw.pop("train", {}), "train"),198 ``ValueError``.
181 )199 """
182 if raw:200 raw = yaml.safe_load(Path(path).read_text(encoding="utf-8")) or {}
183 raise KeyError(f"unknown top-level config key(s) {sorted(raw)} in {path}")201 if not isinstance(raw, Mapping):
184 return cfg202 raise ConfigError(
203 f"config {str(path)!r} must contain a top-level mapping, "
204 f"got {type(raw).__name__}")
205 return config_loader.validate_config(
206 cls, raw, context=str(path), error_cls=ConfigError)
Importance #3: scripts/train.py @@ -22,8 +22,9 @@
22segformer, ...) and model.encoder_name any smp encoder. encoder_weights: imagenet22segformer, ...) and model.encoder_name any smp encoder. encoder_weights: imagenet
23downloads pretrained encoder weights on first use.23downloads pretrained encoder weights on first use.
24"""24"""
25import argparse25import argparse
26from typing import Any
2627
27import lightning.pytorch as pl28import lightning.pytorch as pl
28import torch29import torch
29from lightning.pytorch.callbacks import (30from lightning.pytorch.callbacks import (
Importance #4: scripts/train.py @@ -53,29 +54,51 @@
53 parser.add_argument("--cpu", action="store_true", help="force CPU training")54 parser.add_argument("--cpu", action="store_true", help="force CPU training")
54 return parser.parse_args()55 return parser.parse_args()
5556
5657
58def apply_cli_overrides(cfg: config.HarnessConfig,
59 args: argparse.Namespace) -> config.HarnessConfig:
60 """Returns a copy of ``cfg`` with the CLI overrides applied.
61
62 Config models are frozen, so overrides are applied by copying each touched
63 section instead of assigning to it.
64
65 Args:
66 cfg: The config parsed from the YAML file.
67 args: Parsed CLI arguments; ``None``/empty values override nothing.
68
69 Returns:
70 ``cfg`` itself when no override was given, otherwise an updated copy.
71 """
72 sections: dict[str, dict[str, Any]] = {"train": {}, "data": {}, "model": {}}
73 if args.log_dir:
74 sections["train"]["log_dir"] = args.log_dir
75 if args.num_workers is not None:
76 sections["data"]["num_workers"] = args.num_workers
77 if args.max_epochs is not None:
78 sections["train"]["max_epochs"] = args.max_epochs
79 if args.batch_size is not None:
80 sections["data"]["batch_size"] = args.batch_size
81 if args.model:
82 sections["model"]["name"] = args.model
83 if args.encoder:
84 sections["model"]["encoder_name"] = args.encoder
85 updates = {name: getattr(cfg, name).model_copy(update=values)
86 for name, values in sections.items() if values}
87 return cfg.model_copy(update=updates) if updates else cfg
88
89
57def main() -> None:90def main() -> None:
58 args = parse_args()91 args = parse_args()
59 cfg = config.HarnessConfig.from_yaml(args.config)92 cfg = config.HarnessConfig.from_yaml(args.config)
60 if args.data_root:93 if args.data_root:
61 config.reroot_data_paths(cfg.data, args.data_root)94 cfg = cfg.model_copy(update={
62 if args.log_dir:95 "data": config.reroot_data_paths(cfg.data, args.data_root)})
63 cfg.train.log_dir = args.log_dir
64 if args.num_workers is not None:
65 cfg.data.num_workers = args.num_workers
66 if cfg.model.num_classes != len(tiles.CLASS_NAMES):96 if cfg.model.num_classes != len(tiles.CLASS_NAMES):
67 raise ValueError(97 raise ValueError(
68 f"model.num_classes={cfg.model.num_classes} but the harness tracks "98 f"model.num_classes={cfg.model.num_classes} but the harness tracks "
69 f"{len(tiles.CLASS_NAMES)} classes {tiles.CLASS_NAMES} — metrics would mis-bin")99 f"{len(tiles.CLASS_NAMES)} classes {tiles.CLASS_NAMES} — metrics would mis-bin")
70 if args.max_epochs is not None:100 cfg = apply_cli_overrides(cfg, args)
71 cfg.train.max_epochs = args.max_epochs
72 if args.batch_size is not None:
73 cfg.data.batch_size = args.batch_size
74 if args.model:
75 cfg.model.name = args.model
76 if args.encoder:
77 cfg.model.encoder_name = args.encoder
78101
79 pl.seed_everything(cfg.seed, workers=True)102 pl.seed_everything(cfg.seed, workers=True)
80103
81 dm = datamodule.TilesDataModule(cfg.data)104 dm = datamodule.TilesDataModule(cfg.data)
Importance #5: test/test_ml_harness.py @@ -185,9 +185,9 @@
185 cfg = HarnessConfig.from_yaml(good)185 cfg = HarnessConfig.from_yaml(good)
186 assert cfg.experiment == "t" and cfg.model.name == "unet"186 assert cfg.experiment == "t" and cfg.model.name == "unet"
187 bad = tmp_path / "bad.yaml"187 bad = tmp_path / "bad.yaml"
188 bad.write_text("model:\n encoder: oops\n")188 bad.write_text("model:\n encoder: oops\n")
189 with pytest.raises(KeyError, match="encoder"):189 with pytest.raises(ValueError, match="encoder"):
190 HarnessConfig.from_yaml(bad)190 HarnessConfig.from_yaml(bad)
191191
192192
193def test_shipped_baseline_config_parses() -> None:193def test_shipped_baseline_config_parses() -> None:
Importance #6: test/test_train_overrides.py @@ -1,17 +1,45 @@
1"""Unit tests for train.py configuration overrides."""1"""Unit tests for train.py configuration overrides."""
2from dataclasses import dataclass, field, fields2import argparse
3import importlib.util
4import sys
3from pathlib import Path5from pathlib import Path
46
5from src.train.config import DataConfig, PairSpec, reroot_data_paths7import pytest
8
9from src.train.config import DataConfig, HarnessConfig, PairSpec, reroot_data_paths
10
11_REPO_ROOT = Path(__file__).resolve().parents[1]
12_CLI_FLAGS = ("log_dir", "num_workers", "max_epochs", "batch_size", "model", "encoder")
13
14
15def _train_script():
16 """Imports scripts/train.py as a module; skips when the ml extra is missing."""
17 pytest.importorskip("torch")
18 pytest.importorskip("lightning")
19 pytest.importorskip("segmentation_models_pytorch")
20 pytest.importorskip("albumentations")
21 if "train_script" in sys.modules:
22 return sys.modules["train_script"]
23 spec = importlib.util.spec_from_file_location(
24 "train_script", _REPO_ROOT / "scripts" / "train.py")
25 module = importlib.util.module_from_spec(spec)
26 sys.modules["train_script"] = module
27 spec.loader.exec_module(module)
28 return module
29
30
31def _args(**overrides: object) -> argparse.Namespace:
32 """Builds a parsed-CLI namespace where unset flags are None."""
33 return argparse.Namespace(**{name: overrides.get(name) for name in _CLI_FLAGS})
634
735
8def test_reroot_data_paths_rewrites_relative_pair_dirs(tmp_path: Path) -> None:36def test_reroot_data_paths_rewrites_relative_pair_dirs(tmp_path: Path) -> None:
9 cfg = DataConfig(pairs=[37 cfg = DataConfig(pairs=[
10 PairSpec(images="data/train/images", masks="data/train/masks"),38 PairSpec(images="data/train/images", masks="data/train/masks"),
11 ])39 ])
1240
13 reroot_data_paths(cfg, tmp_path)41 cfg = reroot_data_paths(cfg, tmp_path)
1442
15 assert cfg.pairs[0].images == str(tmp_path / "data/train/images")43 assert cfg.pairs[0].images == str(tmp_path / "data/train/images")
16 assert cfg.pairs[0].masks == str(tmp_path / "data/train/masks")44 assert cfg.pairs[0].masks == str(tmp_path / "data/train/masks")
1745
Importance #7: test/test_train_overrides.py @@ -20,9 +48,9 @@
20 cfg = DataConfig(pairs=[48 cfg = DataConfig(pairs=[
21 PairSpec(images="/mnt/input/images", masks="/mnt/input/masks"),49 PairSpec(images="/mnt/input/images", masks="/mnt/input/masks"),
22 ])50 ])
2351
24 reroot_data_paths(cfg, tmp_path)52 cfg = reroot_data_paths(cfg, tmp_path)
2553
26 assert cfg.pairs[0].images == "/mnt/input/images"54 assert cfg.pairs[0].images == "/mnt/input/images"
27 assert cfg.pairs[0].masks == "/mnt/input/masks"55 assert cfg.pairs[0].masks == "/mnt/input/masks"
2856
Importance #8: test/test_train_overrides.py @@ -32,28 +60,27 @@
32 val_pairs=[PairSpec(images="val/images", masks="val/masks")],60 val_pairs=[PairSpec(images="val/images", masks="val/masks")],
33 test_pairs=[PairSpec(images="test/images", masks="test/masks")],61 test_pairs=[PairSpec(images="test/images", masks="test/masks")],
34 )62 )
3563
36 reroot_data_paths(cfg, tmp_path)64 cfg = reroot_data_paths(cfg, tmp_path)
3765
38 assert cfg.val_pairs[0].images == str(tmp_path / "val/images")66 assert cfg.val_pairs[0].images == str(tmp_path / "val/images")
39 assert cfg.val_pairs[0].masks == str(tmp_path / "val/masks")67 assert cfg.val_pairs[0].masks == str(tmp_path / "val/masks")
40 assert cfg.test_pairs[0].images == str(tmp_path / "test/images")68 assert cfg.test_pairs[0].images == str(tmp_path / "test/images")
41 assert cfg.test_pairs[0].masks == str(tmp_path / "test/masks")69 assert cfg.test_pairs[0].masks == str(tmp_path / "test/masks")
4270
4371
44def test_reroot_data_paths_covers_extra_path_fields(tmp_path: Path) -> None:72def test_reroot_data_paths_covers_extra_path_fields(tmp_path: Path) -> None:
45 @dataclass
46 class ExtendedDataConfig(DataConfig):73 class ExtendedDataConfig(DataConfig):
47 review_sidecar_path: str = "review/dataset_review.sidecar.json"74 review_sidecar_path: str = "review/dataset_review.sidecar.json"
48 materialised_split_dirs: list = field(default_factory=lambda: [75 materialised_split_dirs: list[str] = [
49 "splits/train",76 "splits/train",
50 "/mnt/splits/val",77 "/mnt/splits/val",
51 ])78 ]
5279
53 cfg = ExtendedDataConfig()80 cfg = ExtendedDataConfig()
5481
55 reroot_data_paths(cfg, tmp_path)82 cfg = reroot_data_paths(cfg, tmp_path)
5683
57 assert cfg.review_sidecar_path == str(tmp_path / "review/dataset_review.sidecar.json")84 assert cfg.review_sidecar_path == str(tmp_path / "review/dataset_review.sidecar.json")
58 assert cfg.materialised_split_dirs == [85 assert cfg.materialised_split_dirs == [
59 str(tmp_path / "splits/train"),86 str(tmp_path / "splits/train"),
Importance #9: test/test_train_overrides.py @@ -63,11 +90,11 @@
6390
64def test_current_data_path_fields_are_covered_by_pair_lists() -> None:91def test_current_data_path_fields_are_covered_by_pair_lists() -> None:
65 pair_fields = {"pairs", "val_pairs", "test_pairs"}92 pair_fields = {"pairs", "val_pairs", "test_pairs"}
66 path_like_fields = {93 path_like_fields = {
67 field.name94 name
68 for field in fields(DataConfig)95 for name in DataConfig.model_fields
69 if field.name.endswith(("path", "paths", "dir", "dirs", "root", "roots"))96 if name.endswith(("path", "paths", "dir", "dirs", "root", "roots"))
70 }97 }
7198
72 assert path_like_fields <= pair_fields99 assert path_like_fields <= pair_fields
73100
Importance #10: test/test_train_overrides.py @@ -84,4 +111,38 @@
84 assert cfg.val_pairs[0].images == "val/images"111 assert cfg.val_pairs[0].images == "val/images"
85 assert cfg.val_pairs[0].masks == "val/masks"112 assert cfg.val_pairs[0].masks == "val/masks"
86 assert cfg.test_pairs[0].images == "/abs/test/images"113 assert cfg.test_pairs[0].images == "/abs/test/images"
87 assert cfg.test_pairs[0].masks == "/abs/test/masks"114 assert cfg.test_pairs[0].masks == "/abs/test/masks"
115
116
117def test_apply_cli_overrides_copies_every_touched_section() -> None:
118 cfg = HarnessConfig()
119 args = _args(log_dir="tb", num_workers=2, max_epochs=7, batch_size=3,
120 model="unetplusplus", encoder="resnet34")
121
122 updated = _train_script().apply_cli_overrides(cfg, args)
123
124 assert updated.train.log_dir == "tb" and updated.train.max_epochs == 7
125 assert updated.data.num_workers == 2 and updated.data.batch_size == 3
126 assert updated.model.name == "unetplusplus"
127 assert updated.model.encoder_name == "resnet34"
128 # untouched fields of the copied sections survive
129 assert updated.train.lr == cfg.train.lr
130 assert updated.data.crop_size == cfg.data.crop_size
131 assert updated.model.num_classes == cfg.model.num_classes
132 assert updated.loss == cfg.loss and updated.seed == cfg.seed
133 # the frozen input is never mutated
134 assert cfg.train.log_dir == "runs" and cfg.data.batch_size == 16
135
136
137def test_apply_cli_overrides_without_flags_returns_the_input() -> None:
138 cfg = HarnessConfig()
139
140 assert _train_script().apply_cli_overrides(cfg, _args()) is cfg
141
142
143def test_apply_cli_overrides_honours_zero_valued_flags() -> None:
144 cfg = HarnessConfig()
145
146 updated = _train_script().apply_cli_overrides(cfg, _args(num_workers=0, max_epochs=0))
147
148 assert updated.data.num_workers == 0 and updated.train.max_epochs == 0
Importance #11: pyproject.toml @@ -1,7 +1,7 @@
1[project]1[project]
2name = "iolabs-image-analyzer-line-bitmap-segmentation"2name = "iolabs-image-analyzer-line-bitmap-segmentation"
3version = "0.1.0"3version = "0.1.1"
4description = "Extraction of highway road markings (solid/dashed lane lines) from LiDAR intensity rasters"4description = "Extraction of highway road markings (solid/dashed lane lines) from LiDAR intensity rasters"
5requires-python = ">=3.11,<3.13"5requires-python = ">=3.11,<3.13"
6dependencies = [6dependencies = [
7 "numpy>=1.26",7 "numpy>=1.26",
Importance #12: pyproject.toml @@ -10,8 +10,10 @@
10 "scikit-image>=0.22",10 "scikit-image>=0.22",
11 "matplotlib>=3.7",11 "matplotlib>=3.7",
12 "Pillow>=10.0",12 "Pillow>=10.0",
13 "ezdxf>=1.1",13 "ezdxf>=1.1",
14 "pydantic>=2.7",
15 "iolabs-common>=0.9.0",
14]16]
1517
16[project.optional-dependencies]18[project.optional-dependencies]
17dev = [19dev = [
Importance #13: pyproject.toml @@ -39,19 +41,19 @@
39[tool.pytest.ini_options]41[tool.pytest.ini_options]
40testpaths = ["test"]42testpaths = ["test"]
4143
42# Inference package is published to the private Nexus index (single source of44# Inference package is published to the private Nexus index (single source of
43# truth); pull it from there, not a local sibling checkout. iolabs-common /45# truth); pull it from there, not a local sibling checkout. iolabs-common is a
44# iolabs-logstash are transitive private deps it pulls in — sources mirror the46# direct dependency (the config layer) and iolabs-logstash a transitive private
45# inference repo's own pyproject so `uv sync --extra ml` resolves them identically.47# dep of the inference package; both come from the same Nexus index.
46[[tool.uv.index]]48[[tool.uv.index]]
47name = "nexus"49name = "nexus"
48url = "https://nexus.iolabs.ch/repository/pypi-private/simple/"50url = "https://nexus.iolabs.ch/repository/pypi-private/simple/"
49authenticate = "always"51authenticate = "always"
5052
51[tool.uv.sources]53[tool.uv.sources]
52iolabs-image-analyzer-line-bitmap-inference = { index = "nexus" }54iolabs-image-analyzer-line-bitmap-inference = { index = "nexus" }
53iolabs-common = { path = "../3dai.common" }55iolabs-common = { index = "nexus" }
54iolabs-logstash = { index = "nexus" }56iolabs-logstash = { index = "nexus" }
5557
56[build-system]58[build-system]
57requires = ["hatchling"]59requires = ["hatchling"]
Importance #14: CLAUDE.md @@ -39,9 +39,13 @@
39- `src/model/` — registry; `model.name` = any smp arch, swap via config39- `src/model/` — registry; `model.name` = any smp arch, swap via config
40- `src/losses/`, `src/metrics/` — dice_focal default; also focal_tversky /40- `src/losses/`, `src/metrics/` — dice_focal default; also focal_tversky /
41 soft_cldice / ftl_cldice (handoff recipe); IoU/F1, coverage/tightness,41 soft_cldice / ftl_cldice (handoff recipe); IoU/F1, coverage/tightness,
42 predicted-vs-label mask area, val clDice (centerline topology)42 predicted-vs-label mask area, val clDice (centerline topology)
43- `src/train/` — LightningModule/DataModule, YAML config, `MaskOverlayWriter`43- `src/train/` — LightningModule/DataModule, YAML config (pydantic models on
44 `iolabs.common.config_loader.ConfigModel`: unknown keys rejected, values
45 coerced, instances frozen — **adding a config key = adding one field with
46 its default to the model in `src/train/config.py`**; `from_yaml` raises
47 `config.ConfigError`, a `ValueError`), `MaskOverlayWriter`
44 (`intensity | label | prediction` sheets → TensorBoard + `runs/.../overlays/`)48 (`intensity | label | prediction` sheets → TensorBoard + `runs/.../overlays/`)
45- `scripts/train.py --config configs/<experiment>.yaml` — train (from repo49- `scripts/train.py --config configs/<experiment>.yaml` — train (from repo
46 root, needs `--extra ml`); `tensorboard --logdir runs` to monitor. Configs:50 root, needs `--extra ml`); `tensorboard --logdir runs` to monitor. Configs:
47 `unet_baseline`, `unet_vector_labels`, `unet_confirmed_good` (reviewer-`ok`51 `unet_baseline`, `unet_vector_labels`, `unet_confirmed_good` (reviewer-`ok`
Importance #15: pyproject.toml @@ -1,7 +1,7 @@
1[project]1[project]
2name = "iolabs-image-analyzer-line-bitmap-segmentation"2name = "iolabs-image-analyzer-line-bitmap-segmentation"
3version = "0.1.0"3version = "0.1.1"
4description = "Extraction of highway road markings (solid/dashed lane lines) from LiDAR intensity rasters"4description = "Extraction of highway road markings (solid/dashed lane lines) from LiDAR intensity rasters"
5requires-python = ">=3.11,<3.13"5requires-python = ">=3.11,<3.13"
6dependencies = [6dependencies = [
7 "numpy>=1.26",7 "numpy>=1.26",
Importance #16: pyproject.toml @@ -10,8 +10,10 @@
10 "scikit-image>=0.22",10 "scikit-image>=0.22",
11 "matplotlib>=3.7",11 "matplotlib>=3.7",
12 "Pillow>=10.0",12 "Pillow>=10.0",
13 "ezdxf>=1.1",13 "ezdxf>=1.1",
14 "pydantic>=2.7",
15 "iolabs-common>=0.9.0",
14]16]
1517
16[project.optional-dependencies]18[project.optional-dependencies]
17dev = [19dev = [
Importance #17: pyproject.toml @@ -39,19 +41,19 @@
39[tool.pytest.ini_options]41[tool.pytest.ini_options]
40testpaths = ["test"]42testpaths = ["test"]
4143
42# Inference package is published to the private Nexus index (single source of44# Inference package is published to the private Nexus index (single source of
43# truth); pull it from there, not a local sibling checkout. iolabs-common /45# truth); pull it from there, not a local sibling checkout. iolabs-common is a
44# iolabs-logstash are transitive private deps it pulls in — sources mirror the46# direct dependency (the config layer) and iolabs-logstash a transitive private
45# inference repo's own pyproject so `uv sync --extra ml` resolves them identically.47# dep of the inference package; both come from the same Nexus index.
46[[tool.uv.index]]48[[tool.uv.index]]
47name = "nexus"49name = "nexus"
48url = "https://nexus.iolabs.ch/repository/pypi-private/simple/"50url = "https://nexus.iolabs.ch/repository/pypi-private/simple/"
49authenticate = "always"51authenticate = "always"
5052
51[tool.uv.sources]53[tool.uv.sources]
52iolabs-image-analyzer-line-bitmap-inference = { index = "nexus" }54iolabs-image-analyzer-line-bitmap-inference = { index = "nexus" }
53iolabs-common = { path = "../3dai.common" }55iolabs-common = { index = "nexus" }
54iolabs-logstash = { index = "nexus" }56iolabs-logstash = { index = "nexus" }
5557
56[build-system]58[build-system]
57requires = ["hatchling"]59requires = ["hatchling"]
Importance #18: scripts/train.py @@ -22,8 +22,9 @@
22segformer, ...) and model.encoder_name any smp encoder. encoder_weights: imagenet22segformer, ...) and model.encoder_name any smp encoder. encoder_weights: imagenet
23downloads pretrained encoder weights on first use.23downloads pretrained encoder weights on first use.
24"""24"""
25import argparse25import argparse
26from typing import Any
2627
27import lightning.pytorch as pl28import lightning.pytorch as pl
28import torch29import torch
29from lightning.pytorch.callbacks import (30from lightning.pytorch.callbacks import (
Importance #19: scripts/train.py @@ -53,29 +54,51 @@
53 parser.add_argument("--cpu", action="store_true", help="force CPU training")54 parser.add_argument("--cpu", action="store_true", help="force CPU training")
54 return parser.parse_args()55 return parser.parse_args()
5556
5657
58def apply_cli_overrides(cfg: config.HarnessConfig,
59 args: argparse.Namespace) -> config.HarnessConfig:
60 """Returns a copy of ``cfg`` with the CLI overrides applied.
61
62 Config models are frozen, so overrides are applied by copying each touched
63 section instead of assigning to it.
64
65 Args:
66 cfg: The config parsed from the YAML file.
67 args: Parsed CLI arguments; ``None``/empty values override nothing.
68
69 Returns:
70 ``cfg`` itself when no override was given, otherwise an updated copy.
71 """
72 sections: dict[str, dict[str, Any]] = {"train": {}, "data": {}, "model": {}}
73 if args.log_dir:
74 sections["train"]["log_dir"] = args.log_dir
75 if args.num_workers is not None:
76 sections["data"]["num_workers"] = args.num_workers
77 if args.max_epochs is not None:
78 sections["train"]["max_epochs"] = args.max_epochs
79 if args.batch_size is not None:
80 sections["data"]["batch_size"] = args.batch_size
81 if args.model:
82 sections["model"]["name"] = args.model
83 if args.encoder:
84 sections["model"]["encoder_name"] = args.encoder
85 updates = {name: getattr(cfg, name).model_copy(update=values)
86 for name, values in sections.items() if values}
87 return cfg.model_copy(update=updates) if updates else cfg
88
89
57def main() -> None:90def main() -> None:
58 args = parse_args()91 args = parse_args()
59 cfg = config.HarnessConfig.from_yaml(args.config)92 cfg = config.HarnessConfig.from_yaml(args.config)
60 if args.data_root:93 if args.data_root:
61 config.reroot_data_paths(cfg.data, args.data_root)94 cfg = cfg.model_copy(update={
62 if args.log_dir:95 "data": config.reroot_data_paths(cfg.data, args.data_root)})
63 cfg.train.log_dir = args.log_dir
64 if args.num_workers is not None:
65 cfg.data.num_workers = args.num_workers
66 if cfg.model.num_classes != len(tiles.CLASS_NAMES):96 if cfg.model.num_classes != len(tiles.CLASS_NAMES):
67 raise ValueError(97 raise ValueError(
68 f"model.num_classes={cfg.model.num_classes} but the harness tracks "98 f"model.num_classes={cfg.model.num_classes} but the harness tracks "
69 f"{len(tiles.CLASS_NAMES)} classes {tiles.CLASS_NAMES} — metrics would mis-bin")99 f"{len(tiles.CLASS_NAMES)} classes {tiles.CLASS_NAMES} — metrics would mis-bin")
70 if args.max_epochs is not None:100 cfg = apply_cli_overrides(cfg, args)
71 cfg.train.max_epochs = args.max_epochs
72 if args.batch_size is not None:
73 cfg.data.batch_size = args.batch_size
74 if args.model:
75 cfg.model.name = args.model
76 if args.encoder:
77 cfg.model.encoder_name = args.encoder
78101
79 pl.seed_everything(cfg.seed, workers=True)102 pl.seed_everything(cfg.seed, workers=True)
80103
81 dm = datamodule.TilesDataModule(cfg.data)104 dm = datamodule.TilesDataModule(cfg.data)
Importance #20: src/train/config.py @@ -1,95 +1,99 @@
1"""YAML-backed configuration for the training harness.1"""YAML-backed configuration for the training harness.
22
3Plain dataclasses + yaml, no config framework. Unknown keys raise, so config3Pydantic models on `iolabs.common.config_loader.ConfigModel` + yaml. Unknown keys
4typos fail fast instead of silently training with defaults.4raise, so config typos fail fast instead of silently training with defaults, and
5values are coerced by the shared fleet matrix.
6
7Adding a config key = adding one field with its default to the model below.
8Instances are frozen: derive a changed config with ``model_copy(update=...)``.
5"""9"""
6from dataclasses import dataclass, field, fields10import logging
11from collections.abc import Mapping
7from pathlib import Path12from pathlib import Path
8from typing import Any, TypeVar13from typing import Any, Literal
914
15import pydantic
10import yaml16import yaml
17from iolabs.common import config_loader
18
19logger = logging.getLogger(__name__)
1120
12T = TypeVar("T")21_STROKE_KINDS = frozenset({"solid", "dashed"})
22_PAIR_LIST_FIELDS = ("pairs", "val_pairs", "test_pairs")
23_PATH_FIELD_SUFFIXES = ("path", "paths", "dir", "dirs", "root", "roots")
1324
1425
15def _build(cls: type[T], data: dict[str, Any] | None, where: str) -> T:26class ConfigError(config_loader.ConfigError):
16 data = dict(data or {})27 """Raised when a harness config holds unknown keys or invalid values."""
17 known = {f.name for f in fields(cls)}
18 unknown = sorted(set(data) - known)
19 if unknown:
20 raise KeyError(f"unknown key(s) {unknown} in config section {where!r}; "
21 f"known keys: {sorted(known)}")
22 return cls(**data)
2328
2429
25@dataclass30class PairSpec(config_loader.ConfigModel):
26class PairSpec:
27 """One images-dir / masks-dir pair (see src.dataset.index_tile_pairs)."""31 """One images-dir / masks-dir pair (see src.dataset.index_tile_pairs)."""
28 images: str32 images: str
29 masks: str33 masks: str
3034
3135
32@dataclass36class DataConfig(config_loader.ConfigModel):
33class DataConfig:37 """Tile corpus, split, crop sampling and label rasterization knobs."""
34 pairs: list = field(default_factory=list)38 pairs: list[PairSpec] = []
35 # Optional explicit, pre-split directories (e.g. the symlink folders under39 # Optional explicit, pre-split directories (e.g. the symlink folders under
36 # data/02_processed/<ds>/{train,val,test} built by scripts/build_processed_40 # data/02_processed/<ds>/{train,val,test} built by scripts/build_processed_
37 # splits.py). When val_pairs is set, `pairs` is used in full as the training41 # splits.py). When val_pairs is set, `pairs` is used in full as the training
38 # set and is NOT re-split — val_fraction is ignored — so a geographic split42 # set and is NOT re-split — val_fraction is ignored — so a geographic split
39 # materialised on disk is honoured verbatim. test_pairs feeds test_dataloader.43 # materialised on disk is honoured verbatim. test_pairs feeds test_dataloader.
40 val_pairs: list = field(default_factory=list)44 val_pairs: list[PairSpec] = []
41 test_pairs: list = field(default_factory=list)45 test_pairs: list[PairSpec] = []
42 crop_size: int = 51246 crop_size: int = pydantic.Field(default=512, gt=0)
43 batch_size: int = 1647 batch_size: int = pydantic.Field(default=16, gt=0)
44 num_workers: int = 448 num_workers: int = pydantic.Field(default=4, ge=0)
45 val_fraction: float = 0.1549 val_fraction: float = pydantic.Field(default=0.15, ge=0.0, le=1.0)
46 crops_per_tile: int = 450 crops_per_tile: int = pydantic.Field(default=4, gt=0)
47 pos_crop_prob: float = 0.751 pos_crop_prob: float = pydantic.Field(default=0.7, ge=0.0, le=1.0)
48 min_valid_fraction: float = 0.1052 min_valid_fraction: float = pydantic.Field(default=0.10, ge=0.0, le=1.0)
49 augment: bool = True53 augment: bool = True
50 label_source: str = "rendered" # rendered *_lines.png | vector *_vectors.json | review54 # rendered *_lines.png | vector *_vectors.json | reviewer-confirmed tiles
51 label_stroke_px: int | float | dict[str, float] = 4 # scalar or {solid, dashed}55 label_source: Literal["rendered", "vector", "review"] = "rendered"
52 review_statuses: list = field(default_factory=lambda: ["ok"]) # label_source: review56 # scalar or {solid, dashed}; int stays int so reports render "5", not "5.0"
5357 label_stroke_px: int | float | dict[str, int | float] = 4
5458 review_statuses: list[str] = ["ok"] # label_source: review
55_STROKE_KINDS = frozenset({"solid", "dashed"})
5659
5760 @pydantic.field_validator("label_stroke_px", mode="before")
58def _parse_label_stroke_px(value: Any) -> int | float | dict[str, float]:61 @classmethod
59 """Scalar > 0, or exactly ``{solid, dashed}`` with positive numeric values."""62 def _check_label_stroke_px(cls, value: Any) -> Any:
60 if isinstance(value, (int, float)) and not isinstance(value, bool):63 """Scalar > 0, or exactly ``{solid, dashed}`` with positive numeric values."""
61 if value <= 0:64 if isinstance(value, (int, float)) and not isinstance(value, bool):
62 raise ValueError(f"data.label_stroke_px must be > 0, got {value}")65 if value <= 0:
63 return value66 raise ValueError(f"data.label_stroke_px must be > 0, got {value}")
64 if isinstance(value, dict):67 return value
65 keys = set(value)68 if isinstance(value, Mapping):
66 if keys != _STROKE_KINDS:69 keys = set(value)
67 raise ValueError(70 if keys != _STROKE_KINDS:
68 f"data.label_stroke_px mapping must have exactly the keys "
69 f"{sorted(_STROKE_KINDS)}, got {sorted(keys)}")
70 out: dict[str, float] = {}
71 for kind, width in value.items():
72 if (not isinstance(width, (int, float)) or isinstance(width, bool)
73 or width <= 0):
74 raise ValueError(71 raise ValueError(
75 f"data.label_stroke_px[{kind!r}] must be a positive number, "72 f"data.label_stroke_px mapping must have exactly the keys "
76 f"got {width!r}")73 f"{sorted(_STROKE_KINDS)}, got {sorted(keys)}")
77 out[kind] = width74 for kind, width in value.items():
78 return out75 if (not isinstance(width, (int, float)) or isinstance(width, bool)
79 raise ValueError(76 or width <= 0):
80 f"data.label_stroke_px must be a positive number or a "77 raise ValueError(
81 f"{{solid, dashed}} mapping, got {type(value).__name__}: {value!r}")78 f"data.label_stroke_px[{kind!r}] must be a positive number, "
79 f"got {width!r}")
80 return dict(value)
81 raise ValueError(
82 f"data.label_stroke_px must be a positive number or a "
83 f"{{solid, dashed}} mapping, got {type(value).__name__}: {value!r}")
8284
8385
84def _reroot_path(value: Any, root: Path) -> Any:86def _reroot_path(value: Any, root: Path) -> Any:
87 """Return *value* as an absolute path string, joined onto *root* if relative."""
85 path = Path(value)88 path = Path(value)
86 if path.is_absolute():89 if path.is_absolute():
87 return str(path)90 return str(path)
88 return str(root / path)91 return str(root / path)
8992
9093
91def _reroot_path_value(value: Any, root: Path) -> Any:94def _reroot_path_value(value: Any, root: Path) -> Any:
95 """Re-root every path-like leaf of a scalar/list/tuple/dict value."""
92 if isinstance(value, (str, Path)):96 if isinstance(value, (str, Path)):
93 return _reroot_path(value, root)97 return _reroot_path(value, root)
94 if isinstance(value, list):98 if isinstance(value, list):
95 return [_reroot_path_value(item, root) for item in value]99 return [_reroot_path_value(item, root) for item in value]
Importance #21: src/train/config.py @@ -100,85 +104,103 @@
100 return value104 return value
101105
102106
103def reroot_data_paths(data_cfg: DataConfig, root: str | Path) -> DataConfig:107def reroot_data_paths(data_cfg: DataConfig, root: str | Path) -> DataConfig:
104 """Re-root relative dataset paths in a DataConfig onto root."""108 """Re-root relative dataset paths in a DataConfig onto root.
109
110 Config models are frozen, so this returns an updated copy instead of
111 mutating ``data_cfg`` in place.
112
113 Args:
114 data_cfg: The data section to re-root; never mutated.
115 root: Directory relative paths are joined onto. Absolute paths are kept.
116
117 Returns:
118 A copy of ``data_cfg`` with the pair lists and every path-like field
119 (name ending in path/paths/dir/dirs/root/roots) made absolute.
120 """
105 root = Path(root)121 root = Path(root)
106 for pair_list_name in ("pairs", "val_pairs", "test_pairs"):122 updates: dict[str, Any] = {}
107 for pair in getattr(data_cfg, pair_list_name):123 for name in _PAIR_LIST_FIELDS:
108 pair.images = _reroot_path(pair.images, root)124 specs = getattr(data_cfg, name)
109 pair.masks = _reroot_path(pair.masks, root)125 if specs:
110126 updates[name] = [
111 path_field_suffixes = ("path", "paths", "dir", "dirs", "root", "roots")127 spec.model_copy(update={
112 for field_info in fields(data_cfg):128 "images": _reroot_path(spec.images, root),
113 name = field_info.name129 "masks": _reroot_path(spec.masks, root)})
114 if name in {"pairs", "val_pairs", "test_pairs"}:130 for spec in specs]
131 for name in type(data_cfg).model_fields:
132 if name in _PAIR_LIST_FIELDS or not name.endswith(_PATH_FIELD_SUFFIXES):
115 continue133 continue
116 if name.endswith(path_field_suffixes):134 updates[name] = _reroot_path_value(getattr(data_cfg, name), root)
117 setattr(data_cfg, name, _reroot_path_value(getattr(data_cfg, name), root))135 return data_cfg.model_copy(update=updates)
118 return data_cfg
119136
120137
121@dataclass138class ModelConfig(config_loader.ConfigModel):
122class ModelConfig:139 """Segmentation-models-pytorch architecture/encoder selection."""
123 name: str = "unet"140 name: str = "unet"
124 encoder_name: str = "resnet18"141 encoder_name: str = "resnet18"
125 encoder_weights: str | None = "imagenet" # None = train from scratch142 encoder_weights: str | None = "imagenet" # None = train from scratch
126 in_channels: int = 1143 in_channels: int = pydantic.Field(default=1, gt=0)
127 num_classes: int = 3144 num_classes: int = pydantic.Field(default=3, gt=0)
128 extra: dict = field(default_factory=dict) # passed through to the model factory145 extra: dict[str, Any] = {} # passed through to the model factory
129146
130147
131@dataclass148class LossConfig(config_loader.ConfigModel):
132class LossConfig:149 """Loss selection by registry name plus factory keyword arguments."""
133 name: str = "dice_focal"150 name: str = "dice_focal"
134 args: dict = field(default_factory=dict)151 args: dict[str, Any] = {}
135152
136153
137@dataclass154class TrainerConfig(config_loader.ConfigModel):
138class TrainerConfig:155 """Lightning trainer, logger, and callback knobs."""
139 max_epochs: int = -1 # -1 = no cap; early stopping ends training instead156 max_epochs: int = pydantic.Field(default=-1, ge=-1) # -1 = no cap
140 lr: float = 3.0e-4157 lr: float = pydantic.Field(default=3.0e-4, gt=0)
141 weight_decay: float = 1.0e-4158 weight_decay: float = pydantic.Field(default=1.0e-4, ge=0)
142 precision: str = "auto" # auto -> 16-mixed on CUDA, 32-true on CPU159 precision: str = "auto" # auto -> 16-mixed on CUDA, 32-true on CPU
143 accumulate_grad_batches: int = 1 # match effective batch across experiments160 accumulate_grad_batches: int = pydantic.Field(default=1, ge=1)
144 accelerator: str = "auto"161 accelerator: str = "auto"
145 devices: int | str = 1162 devices: int | str = 1
146 viz_every_n_epochs: int = 2163 viz_every_n_epochs: int = pydantic.Field(default=2, ge=0)
147 viz_samples: int = 4164 viz_samples: int = pydantic.Field(default=4, ge=0)
148 monitor: str = "val/f1_mean_fg"165 monitor: str = "val/f1_mean_fg"
149 monitor_mode: str = "max"166 monitor_mode: Literal["max", "min"] = "max"
150 early_stop_monitor: str = "val/loss" # stop when this stops improving167 early_stop_monitor: str = "val/loss" # stop when this stops improving
151 early_stop_mode: str = "min"168 early_stop_mode: Literal["min", "max"] = "min"
152 early_stop_patience: int = 4 # epochs without improvement before stopping; 0 disables169 # epochs without improvement before stopping; 0 disables
170 early_stop_patience: int = pydantic.Field(default=4, ge=0)
153 log_dir: str = "runs"171 log_dir: str = "runs"
154 log_every_n_steps: int = 10172 log_every_n_steps: int = pydantic.Field(default=10, ge=1)
155173
156174
157@dataclass175class HarnessConfig(config_loader.ConfigModel):
158class HarnessConfig:176 """Top-level training config: one YAML file, one instance."""
159 experiment: str = "experiment"177 experiment: str = "experiment"
160 seed: int = 1337178 seed: int = 1337
161 data: DataConfig = field(default_factory=DataConfig)179 data: DataConfig = DataConfig()
162 model: ModelConfig = field(default_factory=ModelConfig)180 model: ModelConfig = ModelConfig()
163 loss: LossConfig = field(default_factory=LossConfig)181 loss: LossConfig = LossConfig()
164 train: TrainerConfig = field(default_factory=TrainerConfig)182 train: TrainerConfig = TrainerConfig()
165183
166 @classmethod184 @classmethod
167 def from_yaml(cls, path: str | Path) -> "HarnessConfig":185 def from_yaml(cls, path: str | Path) -> "HarnessConfig":
168 raw = yaml.safe_load(Path(path).read_text()) or {}186 """Loads and validates a harness YAML config.
169 data = _build(DataConfig, raw.pop("data", {}), "data")187
170 data.pairs = [_build(PairSpec, p, "data.pairs[]") for p in data.pairs]188 Args:
171 data.val_pairs = [_build(PairSpec, p, "data.val_pairs[]") for p in data.val_pairs]189 path: Path of the YAML file, read as UTF-8.
172 data.test_pairs = [_build(PairSpec, p, "data.test_pairs[]") for p in data.test_pairs]190
173 data.label_stroke_px = _parse_label_stroke_px(data.label_stroke_px)191 Returns:
174 cfg = cls(192 The validated, frozen config.
175 experiment=raw.pop("experiment", cls.experiment),193
176 seed=raw.pop("seed", cls.seed),194 Raises:
177 data=data,195 FileNotFoundError: If ``path`` does not exist.
178 model=_build(ModelConfig, raw.pop("model", {}), "model"),196 ConfigError: If the document is not a mapping, holds an unknown key,
179 loss=_build(LossConfig, raw.pop("loss", {}), "loss"),197 or holds a value invalid for its field. Derives from
180 train=_build(TrainerConfig, raw.pop("train", {}), "train"),198 ``ValueError``.
181 )199 """
182 if raw:200 raw = yaml.safe_load(Path(path).read_text(encoding="utf-8")) or {}
183 raise KeyError(f"unknown top-level config key(s) {sorted(raw)} in {path}")201 if not isinstance(raw, Mapping):
184 return cfg202 raise ConfigError(
203 f"config {str(path)!r} must contain a top-level mapping, "
204 f"got {type(raw).__name__}")
205 return config_loader.validate_config(
206 cls, raw, context=str(path), error_cls=ConfigError)
Importance #22: test/test_ml_harness.py @@ -185,9 +185,9 @@
185 cfg = HarnessConfig.from_yaml(good)185 cfg = HarnessConfig.from_yaml(good)
186 assert cfg.experiment == "t" and cfg.model.name == "unet"186 assert cfg.experiment == "t" and cfg.model.name == "unet"
187 bad = tmp_path / "bad.yaml"187 bad = tmp_path / "bad.yaml"
188 bad.write_text("model:\n encoder: oops\n")188 bad.write_text("model:\n encoder: oops\n")
189 with pytest.raises(KeyError, match="encoder"):189 with pytest.raises(ValueError, match="encoder"):
190 HarnessConfig.from_yaml(bad)190 HarnessConfig.from_yaml(bad)
191191
192192
193def test_shipped_baseline_config_parses() -> None:193def test_shipped_baseline_config_parses() -> None:
Importance #23: test/test_train_overrides.py @@ -1,17 +1,45 @@
1"""Unit tests for train.py configuration overrides."""1"""Unit tests for train.py configuration overrides."""
2from dataclasses import dataclass, field, fields2import argparse
3import importlib.util
4import sys
3from pathlib import Path5from pathlib import Path
46
5from src.train.config import DataConfig, PairSpec, reroot_data_paths7import pytest
8
9from src.train.config import DataConfig, HarnessConfig, PairSpec, reroot_data_paths
10
11_REPO_ROOT = Path(__file__).resolve().parents[1]
12_CLI_FLAGS = ("log_dir", "num_workers", "max_epochs", "batch_size", "model", "encoder")
13
14
15def _train_script():
16 """Imports scripts/train.py as a module; skips when the ml extra is missing."""
17 pytest.importorskip("torch")
18 pytest.importorskip("lightning")
19 pytest.importorskip("segmentation_models_pytorch")
20 pytest.importorskip("albumentations")
21 if "train_script" in sys.modules:
22 return sys.modules["train_script"]
23 spec = importlib.util.spec_from_file_location(
24 "train_script", _REPO_ROOT / "scripts" / "train.py")
25 module = importlib.util.module_from_spec(spec)
26 sys.modules["train_script"] = module
27 spec.loader.exec_module(module)
28 return module
29
30
31def _args(**overrides: object) -> argparse.Namespace:
32 """Builds a parsed-CLI namespace where unset flags are None."""
33 return argparse.Namespace(**{name: overrides.get(name) for name in _CLI_FLAGS})
634
735
8def test_reroot_data_paths_rewrites_relative_pair_dirs(tmp_path: Path) -> None:36def test_reroot_data_paths_rewrites_relative_pair_dirs(tmp_path: Path) -> None:
9 cfg = DataConfig(pairs=[37 cfg = DataConfig(pairs=[
10 PairSpec(images="data/train/images", masks="data/train/masks"),38 PairSpec(images="data/train/images", masks="data/train/masks"),
11 ])39 ])
1240
13 reroot_data_paths(cfg, tmp_path)41 cfg = reroot_data_paths(cfg, tmp_path)
1442
15 assert cfg.pairs[0].images == str(tmp_path / "data/train/images")43 assert cfg.pairs[0].images == str(tmp_path / "data/train/images")
16 assert cfg.pairs[0].masks == str(tmp_path / "data/train/masks")44 assert cfg.pairs[0].masks == str(tmp_path / "data/train/masks")
1745
Importance #24: test/test_train_overrides.py @@ -20,9 +48,9 @@
20 cfg = DataConfig(pairs=[48 cfg = DataConfig(pairs=[
21 PairSpec(images="/mnt/input/images", masks="/mnt/input/masks"),49 PairSpec(images="/mnt/input/images", masks="/mnt/input/masks"),
22 ])50 ])
2351
24 reroot_data_paths(cfg, tmp_path)52 cfg = reroot_data_paths(cfg, tmp_path)
2553
26 assert cfg.pairs[0].images == "/mnt/input/images"54 assert cfg.pairs[0].images == "/mnt/input/images"
27 assert cfg.pairs[0].masks == "/mnt/input/masks"55 assert cfg.pairs[0].masks == "/mnt/input/masks"
2856
Importance #25: test/test_train_overrides.py @@ -32,28 +60,27 @@
32 val_pairs=[PairSpec(images="val/images", masks="val/masks")],60 val_pairs=[PairSpec(images="val/images", masks="val/masks")],
33 test_pairs=[PairSpec(images="test/images", masks="test/masks")],61 test_pairs=[PairSpec(images="test/images", masks="test/masks")],
34 )62 )
3563
36 reroot_data_paths(cfg, tmp_path)64 cfg = reroot_data_paths(cfg, tmp_path)
3765
38 assert cfg.val_pairs[0].images == str(tmp_path / "val/images")66 assert cfg.val_pairs[0].images == str(tmp_path / "val/images")
39 assert cfg.val_pairs[0].masks == str(tmp_path / "val/masks")67 assert cfg.val_pairs[0].masks == str(tmp_path / "val/masks")
40 assert cfg.test_pairs[0].images == str(tmp_path / "test/images")68 assert cfg.test_pairs[0].images == str(tmp_path / "test/images")
41 assert cfg.test_pairs[0].masks == str(tmp_path / "test/masks")69 assert cfg.test_pairs[0].masks == str(tmp_path / "test/masks")
4270
4371
44def test_reroot_data_paths_covers_extra_path_fields(tmp_path: Path) -> None:72def test_reroot_data_paths_covers_extra_path_fields(tmp_path: Path) -> None:
45 @dataclass
46 class ExtendedDataConfig(DataConfig):73 class ExtendedDataConfig(DataConfig):
47 review_sidecar_path: str = "review/dataset_review.sidecar.json"74 review_sidecar_path: str = "review/dataset_review.sidecar.json"
48 materialised_split_dirs: list = field(default_factory=lambda: [75 materialised_split_dirs: list[str] = [
49 "splits/train",76 "splits/train",
50 "/mnt/splits/val",77 "/mnt/splits/val",
51 ])78 ]
5279
53 cfg = ExtendedDataConfig()80 cfg = ExtendedDataConfig()
5481
55 reroot_data_paths(cfg, tmp_path)82 cfg = reroot_data_paths(cfg, tmp_path)
5683
57 assert cfg.review_sidecar_path == str(tmp_path / "review/dataset_review.sidecar.json")84 assert cfg.review_sidecar_path == str(tmp_path / "review/dataset_review.sidecar.json")
58 assert cfg.materialised_split_dirs == [85 assert cfg.materialised_split_dirs == [
59 str(tmp_path / "splits/train"),86 str(tmp_path / "splits/train"),
Importance #26: test/test_train_overrides.py @@ -63,11 +90,11 @@
6390
64def test_current_data_path_fields_are_covered_by_pair_lists() -> None:91def test_current_data_path_fields_are_covered_by_pair_lists() -> None:
65 pair_fields = {"pairs", "val_pairs", "test_pairs"}92 pair_fields = {"pairs", "val_pairs", "test_pairs"}
66 path_like_fields = {93 path_like_fields = {
67 field.name94 name
68 for field in fields(DataConfig)95 for name in DataConfig.model_fields
69 if field.name.endswith(("path", "paths", "dir", "dirs", "root", "roots"))96 if name.endswith(("path", "paths", "dir", "dirs", "root", "roots"))
70 }97 }
7198
72 assert path_like_fields <= pair_fields99 assert path_like_fields <= pair_fields
73100
Importance #27: test/test_train_overrides.py @@ -84,4 +111,38 @@
84 assert cfg.val_pairs[0].images == "val/images"111 assert cfg.val_pairs[0].images == "val/images"
85 assert cfg.val_pairs[0].masks == "val/masks"112 assert cfg.val_pairs[0].masks == "val/masks"
86 assert cfg.test_pairs[0].images == "/abs/test/images"113 assert cfg.test_pairs[0].images == "/abs/test/images"
87 assert cfg.test_pairs[0].masks == "/abs/test/masks"114 assert cfg.test_pairs[0].masks == "/abs/test/masks"
115
116
117def test_apply_cli_overrides_copies_every_touched_section() -> None:
118 cfg = HarnessConfig()
119 args = _args(log_dir="tb", num_workers=2, max_epochs=7, batch_size=3,
120 model="unetplusplus", encoder="resnet34")
121
122 updated = _train_script().apply_cli_overrides(cfg, args)
123
124 assert updated.train.log_dir == "tb" and updated.train.max_epochs == 7
125 assert updated.data.num_workers == 2 and updated.data.batch_size == 3
126 assert updated.model.name == "unetplusplus"
127 assert updated.model.encoder_name == "resnet34"
128 # untouched fields of the copied sections survive
129 assert updated.train.lr == cfg.train.lr
130 assert updated.data.crop_size == cfg.data.crop_size
131 assert updated.model.num_classes == cfg.model.num_classes
132 assert updated.loss == cfg.loss and updated.seed == cfg.seed
133 # the frozen input is never mutated
134 assert cfg.train.log_dir == "runs" and cfg.data.batch_size == 16
135
136
137def test_apply_cli_overrides_without_flags_returns_the_input() -> None:
138 cfg = HarnessConfig()
139
140 assert _train_script().apply_cli_overrides(cfg, _args()) is cfg
141
142
143def test_apply_cli_overrides_honours_zero_valued_flags() -> None:
144 cfg = HarnessConfig()
145
146 updated = _train_script().apply_cli_overrides(cfg, _args(num_workers=0, max_epochs=0))
147
148 assert updated.data.num_workers == 0 and updated.train.max_epochs == 0