From 30760069e1f222e0c15a03cb0f1637926f7ac6c7 Mon Sep 17 00:00:00 2001 From: iback Date: Fri, 21 Aug 2026 06:17:31 +0000 Subject: [PATCH] feat: add the sectioned config schema smauglab/config.py, with no consumers yet: nothing reads it until the builder lands, and the shipped configs are untouched by this commit. A config becomes a sectioned document -- GPU, CPU, pipeline, and '_'-prefixed comments -- where a key is a class name exactly and a parameter is a constructor argument exactly, both checked against the registry. A flat, section-less document is rejected rather than migrated. It used to be read as "GPU or CPU, whichever the keys look like", and the two namespaces overlap enough that GaussianBlurTransform meant a different transform depending on which builder read it. Every problem in a file is reported at once, with "did you mean" suggestions. The old behaviour surfaced one per run, which for a 30-key config is 30 edit-run cycles. pipeline.order is the opt-in discussed on the registry commit. The default, `registry`, is registry.PIPELINE_ORDER: fixed, and the same for every config. `config` takes the order the keys appear in the file instead. It is not the default because it makes a pipeline sensitive to something people reasonably treat as cosmetic -- reordering or reformatting a config would silently change what it does -- so opting in makes that intent explicit and greppable. It is independent of pipeline.mode. PipelineMode, OrderSource and validate_section live here rather than in the builder that pr56 put them in. They describe a config document, so config.py depends only on the registry and the builder will import them from here. That is also what lets this land on its own. pipeline.mode replaces encoding the arrangement in the trainer *class*, which baked the choice into the name of every run directory. Co-Authored-By: Claude Opus 5 --- smauglab/config.py | 377 +++++++++++++++++++++++++++++++ unit_tests/test_config_schema.py | 148 ++++++++++++ 2 files changed, 525 insertions(+) create mode 100644 smauglab/config.py create mode 100644 unit_tests/test_config_schema.py diff --git a/smauglab/config.py b/smauglab/config.py new file mode 100644 index 0000000..8af0963 --- /dev/null +++ b/smauglab/config.py @@ -0,0 +1,377 @@ +"""Load and validate a SmaugLab augmentation config. + +A config is a JSON document with reserved top-level sections: + + { + "_comment": "...", // any key starting with '_' is ignored + "GPU": { "": { }, ... }, + "CPU": { ... }, + "pipeline": { "random_choose": { ... } } + } + +A key is a class name, exactly, and a parameter is a constructor argument, exactly. +Both are checked against the registry, and every problem in a file is reported at +once rather than one per run. + +A flat document with no section is rejected: it used to be interpreted as "GPU or +CPU, whichever the keys look like", and the two namespaces overlapped enough that +`GaussianBlurTransform` meant different transforms depending on which builder read +it. See `migration/` in the repository to bring an old file forward -- `MIGRATE_HINT` +below is the single source of truth for that pointer, and the error messages quote it. +""" + +from __future__ import annotations + +import copy +import difflib +import functools +import json +from enum import Enum +from pathlib import Path +from typing import Any + +from smauglab import registry +from smauglab.registry import Backend, InvalidConfigError + + +class PipelineMode(str, Enum): + """How the GPU pipeline arranges what the registry gives it. + + This used to be encoded in the *trainer class* -- a separate subclass per + arrangement -- which baked the choice into the name of every run directory. + """ + + #: Everything in pipeline order. What `AugTransformsGPU` has always done. + SEQUENTIAL = "sequential" + #: Geometry in order, then TA and GE each shuffled inside a RandomChooseX. + RANDOM_ORDER = "random_order" + #: As above but the TA bucket keeps its order; GE is not bucketed separately. + RANDOM_ORDER_TA = "random_order_ta" + + +class OrderSource(str, Enum): + """Where a pipeline's transform order comes from. + + `registry` -- the default -- is `registry.PIPELINE_ORDER`, which is fixed and the + same for every config. `config` takes the order the keys appear in the file + instead, for callers who want to control the sequence per experiment. + + Not the default, because it makes the pipeline sensitive to something people + reasonably treat as cosmetic: reordering or reformatting a config would silently + change what it does. Opting in makes that intent explicit and greppable. + """ + + REGISTRY = "registry" + CONFIG = "config" + + +#: Keys a config section may hold that are not augmentations. +NON_AUGMENTATION_KEYS = ("_",) + + +def validate_section(section: dict, backend: Backend, *, source: str = "") -> list[str]: + """Return every problem in one backend's section. Empty means it will build.""" + problems: list[str] = [] + for name, params in section.items(): + if name.startswith(NON_AUGMENTATION_KEYS): + continue + try: + entry = registry.get(name, backend) + except registry.UnknownAugmentationError as exc: + problems.append(str(exc).replace("\n", "\n ")) + continue + if not isinstance(params, dict): + problems.append(f"{name}: expected a block of parameters, got {type(params).__name__}") + continue + + accepted = registry.accepted_params(entry) + problems.extend(registry.unknown_parameter_message(entry, key).replace("\n", "\n ") for key in params if key not in accepted) + missing = registry.required_params(entry) - set(params) - set(entry.context_params) + if missing: + problems.append(f"{name}: missing required parameter(s) {', '.join(sorted(missing))}") + _ = source + return problems + + +#: Top-level keys that are not augmentation sections. +RESERVED_SECTIONS = ("pipeline",) + +#: Keys the `pipeline` section may hold. +PIPELINE_KEYS = ("mode", "order", "random_choose") + +#: How to bring a pre-registry config forward. The migrator is a one-time tool kept +#: in the repository rather than shipped in the wheel, so this points at the repo. +MIGRATE_HINT = "see migration/ in the SmaugLab repository to bring an old config forward." + + +class SmaugConfig: + """A parsed, validated config document.""" + + def __init__(self, payload: dict, source: str = "") -> None: + self.payload = payload + self.source = source + self.validate() + + # -- construction ------------------------------------------------------------ + + @classmethod + def from_path(cls, path: str | Path) -> SmaugConfig: + path = Path(path) + return cls(json.loads(path.read_text()), source=path.name) + + # -- access ------------------------------------------------------------------ + + def section(self, backend: Backend) -> dict[str, Any]: + """The augmentation blocks for a backend. + + A deepcopy, because `load_config` caches documents and a caller that mutated + what it got back would poison every later read of the same file. + """ + return copy.deepcopy(self.payload.get(backend.value, {})) + + def pipeline_options(self, name: str) -> dict[str, Any]: + """Options for a named pipeline feature, e.g. `random_choose`. + + Always returns a dict. The old builder did `config.get("RandomChooseXTransforms")` + and then `.get()` on the result, which raised AttributeError on every config + that omitted the block. + """ + return copy.deepcopy(self.payload.get("pipeline", {}).get(name, {})) + + def pipeline_mode(self) -> PipelineMode: + """How the GPU pipeline arranges its transforms. + + This used to be encoded in the *trainer class* -- a separate subclass per + arrangement -- which meant the choice was baked into the name of every run + directory and could not be varied without a new class. It is a property of + the augmentation setup, so it belongs in the config next to it. + """ + raw = self.payload.get("pipeline", {}).get("mode") + return PipelineMode(raw) if raw else PipelineMode.SEQUENTIAL + + def order_source(self) -> OrderSource: + """Where this config's transform order comes from. + + Defaults to the registry's fixed PIPELINE_ORDER. `"pipeline": {"order": + "config"}` switches to the order the keys appear in the file. + """ + raw = self.payload.get("pipeline", {}).get("order") + return OrderSource(raw) if raw else OrderSource.REGISTRY + + def names(self, backend: Backend) -> list[str]: + return [k for k in self.section(backend) if not k.startswith("_")] + + # -- validation -------------------------------------------------------------- + + def validate(self) -> None: + problems: list[str] = [] + + known_sections = {b.value for b in Backend} | set(RESERVED_SECTIONS) + for key in self.payload: + if key.startswith("_") or key in known_sections: + continue + problems.append( + f"unknown top-level key {key!r}. Expected one of " + f"{', '.join(sorted(known_sections))}, or a '_'-prefixed comment. " + f"A flat config without a backend section is no longer accepted -- {MIGRATE_HINT}" + ) + + if not any(b.value in self.payload for b in Backend): + problems.append(f"no GPU or CPU section; this looks like a pre-registry config. {MIGRATE_HINT}") + + problems.extend(self._pipeline_problems()) + + for backend in Backend: + section = self.payload.get(backend.value) + if section is None: + continue + if not isinstance(section, dict): + problems.append(f"{backend.value}: expected an object of augmentation blocks") + continue + problems.extend(f"{backend.value}.{p}" for p in validate_section(section, backend, source=self.source)) + + if problems: + raise InvalidConfigError(self.source, problems) + + def _pipeline_problems(self) -> list[str]: + """Check the `pipeline` section, which was previously accepted unchecked.""" + pipeline = self.payload.get("pipeline") + if pipeline is None: + return [] + if not isinstance(pipeline, dict): + return ["pipeline: expected an object"] + + problems = [] + for key in pipeline: + if key.startswith("_") or key in PIPELINE_KEYS: + continue + close = difflib.get_close_matches(key, PIPELINE_KEYS, n=2, cutoff=0.6) + hint = f" Did you mean: {', '.join(close)}?" if close else "" + problems.append(f"pipeline: unknown key {key!r}. Accepted: {', '.join(PIPELINE_KEYS)}.{hint}") + + order = pipeline.get("order") + if order is not None: + valid_orders = [o.value for o in OrderSource] + if order not in valid_orders: + close = difflib.get_close_matches(str(order), valid_orders, n=2, cutoff=0.5) + hint = f" Did you mean: {', '.join(close)}?" if close else "" + problems.append(f"pipeline.order: unknown source {order!r}. Accepted: {', '.join(valid_orders)}.{hint}") + + mode = pipeline.get("mode") + if mode is not None: + valid = [m.value for m in PipelineMode] + if mode not in valid: + close = difflib.get_close_matches(str(mode), valid, n=2, cutoff=0.5) + hint = f" Did you mean: {', '.join(close)}?" if close else "" + problems.append(f"pipeline.mode: unknown mode {mode!r}. Accepted: {', '.join(valid)}.{hint}") + return problems + + +@functools.lru_cache(maxsize=8) +def load_config(path: str) -> SmaugConfig: + """Parse and validate a config, once per path. + + Cached because the nnU-Net trainer reads the same file twice: its + `get_training_transforms` is a staticmethod (nnU-Net's contract), so it cannot + reach the instance's already-parsed config and has to open the file itself. + `SmaugConfig.section` hands out copies, so sharing the parsed document is safe. + """ + return SmaugConfig.from_path(path) + + +def validate_file(path: str | Path) -> list[str]: + """Every problem in a config file, without raising. Empty means it is valid.""" + try: + SmaugConfig.from_path(path) + except InvalidConfigError as exc: + return exc.problems + except json.JSONDecodeError as exc: + return [f"not valid JSON: {exc}"] + return [] + + +def config_hash(payload: dict, algo: str = "sha256") -> str: + """Content-addressed identity for a config. + + Byte-for-byte the same canonicalisation segtransferaug has always used, because + experiment directories are named `...-aug--c-` and + changing it would orphan every existing run folder. + """ + import hashlib + + canonical = json.dumps(payload, sort_keys=True, separators=(",", ":"), ensure_ascii=False).encode("utf-8") + digest = hashlib.new(algo) + digest.update(canonical) + return digest.hexdigest() + + +def file_hash(path: str | Path, algo: str = "sha256") -> str: + """Hash a source file's text. + + Used downstream to name experiment directories after the implementation that + produced them, so a change to the transforms is visible in the run name. + """ + import hashlib + + digest = hashlib.new(algo) + digest.update(Path(path).read_text().encode("utf-8")) + return digest.hexdigest() + + +def registered_names(backend: Backend) -> list[str]: + """Convenience re-export so callers need not import the registry directly.""" + return registry.names(backend) + + +# --- config manipulation ---------------------------------------------------------- +# +# Upstreamed from segtransferaug/utils/smauglab_config.py, which drove the sweep +# scripts. Three module-level absolute paths went away (the packaged config is +# resolved through importlib.resources now), and the hardcoded AUG2GROUP table +# became the registry's `group` field, so a new augmentation no longer has to be +# added to a dict in a different repository before the sweeps can see it. + + +def default_config_path() -> Path: + """The packaged default GPU config.""" + import importlib.resources + + from smauglab import configs + + return Path(str(importlib.resources.files(configs))) / "transform_params_gpu.json" + + +def load_json(path: str | Path) -> dict: + return json.loads(Path(path).read_text()) + + +def transform_names(payload: dict, backend: Backend = Backend.GPU) -> list[str]: + """The augmentations a config actually names.""" + return [k for k in payload.get(backend.value, {}) if not k.startswith("_")] + + +def single_transform_config(name: str, payload: dict, backend: Backend = Backend.GPU) -> dict: + """A copy of the config with only one augmentation left enabled.""" + registry.get(name, backend) # raises with a suggestion if the name is wrong + out = copy.deepcopy(payload) + section = out.get(backend.value, {}) + out[backend.value] = {k: v for k, v in section.items() if k == name or k.startswith("_")} + return out + + +def drop_zero_probability(payload: dict, backend: Backend = Backend.GPU) -> dict: + """Remove augmentations that would never fire, so the config says what it does.""" + out = copy.deepcopy(payload) + section = out.get(backend.value, {}) + out[backend.value] = {k: v for k, v in section.items() if k.startswith("_") or v.get("p", 1.0) != 0} + return out + + +def filter_by_group(payload: dict, group, backend: Backend = Backend.GPU) -> dict: + """Keep only the augmentations in one GEO/GE/TA group.""" + keep = set(registry.names(backend, group=group)) + out = copy.deepcopy(payload) + section = out.get(backend.value, {}) + out[backend.value] = {k: v for k, v in section.items() if k in keep or k.startswith("_")} + return out + + +def set_probabilities(payload: dict, p: float, group=None, backend: Backend = Backend.GPU) -> dict: + """Set `p` on every augmentation, or only on one group.""" + targets = set(registry.names(backend, group=group)) if group is not None else None + out = copy.deepcopy(payload) + for name, block in out.get(backend.value, {}).items(): + if name.startswith("_") or not isinstance(block, dict): + continue + if targets is None or name in targets: + block["p"] = p + return out + + +def write_temp_config(payload: dict, directory: str | Path | None = None) -> str: + """Materialise a config so it can be handed to a subprocess by path. + + Named after its content hash, so the same config reuses the same file and a + sweep does not fill the directory with near-duplicates. + """ + import tempfile + + target_dir = Path(directory) if directory else Path(tempfile.gettempdir()) / "smauglab_configs" + target_dir.mkdir(parents=True, exist_ok=True) + path = target_dir / f"transform_params_{config_hash(payload)[:8]}.json" + path.write_text(json.dumps(payload, indent=4) + "\n") + return str(path) + + +def remove_temp_configs(directory: str | Path | None = None) -> int: + """Delete configs written by `write_temp_config`. Returns how many went.""" + import tempfile + + target_dir = Path(directory) if directory else Path(tempfile.gettempdir()) / "smauglab_configs" + if not target_dir.is_dir(): + return 0 + removed = 0 + for path in target_dir.glob("transform_params_*.json"): + path.unlink() + removed += 1 + return removed diff --git a/unit_tests/test_config_schema.py b/unit_tests/test_config_schema.py new file mode 100644 index 0000000..7723fc2 --- /dev/null +++ b/unit_tests/test_config_schema.py @@ -0,0 +1,148 @@ +"""The sectioned config document: what it accepts, and what it says when it does not. + +A config used to be a flat mapping of augmentation name to parameters, read as "GPU or +CPU, whichever the keys look like". The two namespaces overlap enough that +`GaussianBlurTransform` meant a different transform depending on which builder read it, +so sections are mandatory now and a flat document is rejected outright. + +Every problem in a file is reported at once. The old behaviour surfaced one per run, +which for a 30-key config is 30 edit-run cycles. +""" + +from __future__ import annotations + +import json +import tempfile +import unittest +from pathlib import Path + +from smauglab.config import OrderSource, PipelineMode, SmaugConfig, validate_file, validate_section +from smauglab.registry import Backend, InvalidConfigError + + +def write(payload: dict) -> Path: + path = Path(tempfile.mkdtemp()) / "config.json" + path.write_text(json.dumps(payload)) + return path + + +class TestSections(unittest.TestCase): + def test_a_minimal_gpu_config_validates(self): + config = SmaugConfig({"GPU": {"RandomFlipTransformGPU": {"p": 0.5}}}) + self.assertEqual(config.names(Backend.GPU), ["RandomFlipTransformGPU"]) + + def test_a_flat_section_less_config_is_rejected(self): + """The hard break. It used to be guessed at from the key names.""" + with self.assertRaises(InvalidConfigError) as caught: + SmaugConfig({"FlipTransform": {"probability": 0.5}}) + self.assertIn("pre-registry config", str(caught.exception)) + + def test_an_unknown_top_level_key_is_rejected(self): + with self.assertRaises(InvalidConfigError) as caught: + SmaugConfig({"GPU": {}, "GPUU": {}}) + self.assertIn("unknown top-level key", str(caught.exception)) + + def test_underscore_keys_are_comments(self): + config = SmaugConfig({"_comment": "anything", "GPU": {"_note": "ignored", "RandomFlipTransformGPU": {}}}) + self.assertEqual(config.names(Backend.GPU), ["RandomFlipTransformGPU"]) + + def test_sections_are_handed_out_as_copies(self): + """load_config caches documents, so a caller that mutates must not poison it.""" + config = SmaugConfig({"GPU": {"RandomFlipTransformGPU": {"p": 0.5}}}) + config.section(Backend.GPU)["RandomFlipTransformGPU"]["p"] = 99 + self.assertEqual(config.section(Backend.GPU)["RandomFlipTransformGPU"]["p"], 0.5) + + def test_a_missing_section_is_empty_not_an_error(self): + self.assertEqual(SmaugConfig({"GPU": {}}).section(Backend.CPU), {}) + + +class TestProblemsAreReportedTogether(unittest.TestCase): + def test_every_problem_in_a_file_is_reported_at_once(self): + with self.assertRaises(InvalidConfigError) as caught: + SmaugConfig( + { + "GPU": { + "NotATransform": {}, + "RandomFlipTransformGPU": {"nonsense": 1}, + "RandomScharrGPU": {"probability": 0.5}, + } + } + ) + self.assertEqual(len(caught.exception.problems), 3, caught.exception.problems) + + def test_an_unknown_augmentation_suggests_a_close_match(self): + problems = validate_section({"RandomFlipTransformGP": {}}, Backend.GPU) + self.assertTrue(any("RandomFlipTransformGPU" in p for p in problems), problems) + + def test_a_renamed_parameter_is_pointed_at_its_replacement(self): + """`probability` -> `p` is the commonest migration mistake, and difflib + cannot bridge it: the two strings score ~0.17.""" + problems = validate_section({"RandomFlipTransformGPU": {"probability": 0.5}}, Backend.GPU) + self.assertTrue(any("'probability' -> p" in p for p in problems), problems) + + def test_a_block_that_is_not_an_object_is_reported(self): + problems = validate_section({"RandomFlipTransformGPU": 0.5}, Backend.GPU) + self.assertTrue(any("block of parameters" in p for p in problems), problems) + + def test_validate_file_returns_problems_without_raising(self): + self.assertEqual(validate_file(write({"GPU": {"RandomFlipTransformGPU": {}}})), []) + self.assertTrue(validate_file(write({"GPU": {"Nope": {}}}))) + + def test_invalid_json_is_reported_as_such(self): + path = Path(tempfile.mkdtemp()) / "broken.json" + path.write_text("{not json") + self.assertTrue(any("not valid JSON" in p for p in validate_file(path))) + + +class TestPipelineSection(unittest.TestCase): + def test_mode_defaults_to_sequential(self): + self.assertIs(SmaugConfig({"GPU": {}}).pipeline_mode(), PipelineMode.SEQUENTIAL) + + def test_mode_is_read_from_the_config(self): + config = SmaugConfig({"GPU": {}, "pipeline": {"mode": "random_order"}}) + self.assertIs(config.pipeline_mode(), PipelineMode.RANDOM_ORDER) + + def test_an_unknown_mode_is_rejected_with_a_suggestion(self): + with self.assertRaises(InvalidConfigError) as caught: + SmaugConfig({"GPU": {}, "pipeline": {"mode": "random_ordr"}}) + self.assertIn("random_order", str(caught.exception)) + + def test_an_unknown_pipeline_key_is_rejected(self): + with self.assertRaises(InvalidConfigError) as caught: + SmaugConfig({"GPU": {}, "pipeline": {"modes": "sequential"}}) + self.assertIn("unknown key", str(caught.exception)) + + def test_pipeline_options_is_always_a_dict(self): + """The old builder called .get() on the result and raised AttributeError + on every config that omitted the block.""" + self.assertEqual(SmaugConfig({"GPU": {}}).pipeline_options("random_choose"), {}) + + +class TestOrderSource(unittest.TestCase): + """`pipeline.order`: registry order by default, config key order on request.""" + + def test_it_defaults_to_the_registry(self): + self.assertIs(SmaugConfig({"GPU": {}}).order_source(), OrderSource.REGISTRY) + + def test_config_key_order_can_be_asked_for(self): + config = SmaugConfig({"GPU": {}, "pipeline": {"order": "config"}}) + self.assertIs(config.order_source(), OrderSource.CONFIG) + + def test_the_registry_default_can_be_stated_explicitly(self): + config = SmaugConfig({"GPU": {}, "pipeline": {"order": "registry"}}) + self.assertIs(config.order_source(), OrderSource.REGISTRY) + + def test_an_unknown_order_source_is_rejected_with_a_suggestion(self): + with self.assertRaises(InvalidConfigError) as caught: + SmaugConfig({"GPU": {}, "pipeline": {"order": "confgi"}}) + self.assertIn("pipeline.order", str(caught.exception)) + self.assertIn("config", str(caught.exception)) + + def test_order_and_mode_are_independent(self): + config = SmaugConfig({"GPU": {}, "pipeline": {"mode": "random_order", "order": "config"}}) + self.assertIs(config.pipeline_mode(), PipelineMode.RANDOM_ORDER) + self.assertIs(config.order_source(), OrderSource.CONFIG) + + +if __name__ == "__main__": + unittest.main()