From 34af2dffbf6f29178d222776bf9057a080747703 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Korbinian=20P=C3=B6ppel?= Date: Sat, 1 Aug 2026 18:21:38 +0200 Subject: [PATCH 1/4] Cache Hydra composition so sweep builds stop redoing the same work Mirrors the change in hydra_staged_sweep into the vendored copy. Building a sweep composes the same config tree once per sweep point, and Hydra keeps nothing between compose() calls: every point re-read and re-parsed the same YAML, re-walked the same defaults list, re-merged the same configs and re-parsed the same interpolation and override strings. hydra_staged_sweep/config/cache.py adds six in-memory caches over Hydra and OmegaConf -- parsed config files, config-group lookups, merged defaults lists, interpolation parse trees, override parses, and a libyaml-backed YAML loader. Entries carry a stat() fingerprint and are revalidated on every lookup, so editing a config invalidates exactly what depends on it. They install when hydra_staged_sweep/config/loader.py is imported; HYDRA_STAGED_SWEEP_CACHE=0 turns them off. Staged sweeps also stop recomposing siblings: resolve_sweep_with_dag reuses the sibling's already-resolved JobPlan config, builds the sibling context once per stage chain, and merges it into the composed config in one step instead of flattening it into several hundred ++key=value overrides for Hydra to parse and apply one at a time. job_parameters still records those overrides, so a job stays reproducible from the command line. run_autoexp.py --dry-run on the 300-job titan_qwen3_depth_width_scaling sweep goes from 469s to 82s (5.75x), with all 300 rendered sbatch scripts byte-identical to before. --- .../hydra_staged_sweep/config/cache.py | 520 ++++++++++++++++++ .../hydra_staged_sweep/config/loader.py | 34 +- .../hydra_staged_sweep/dag_resolver.py | 92 +++- tests/hydra_staged_sweep/test_config_cache.py | 167 ++++++ 4 files changed, 787 insertions(+), 26 deletions(-) create mode 100644 oellm_autoexp/hydra_staged_sweep/config/cache.py create mode 100644 tests/hydra_staged_sweep/test_config_cache.py diff --git a/oellm_autoexp/hydra_staged_sweep/config/cache.py b/oellm_autoexp/hydra_staged_sweep/config/cache.py new file mode 100644 index 00000000..776dc71f --- /dev/null +++ b/oellm_autoexp/hydra_staged_sweep/config/cache.py @@ -0,0 +1,520 @@ +"""In-memory caches that make repeated Hydra composition cheap. + +Building a sweep composes the same config tree once per sweep point. Hydra +caches nothing between ``compose()`` calls, so every point re-reads and +re-parses the same YAML files, re-walks the same defaults list, re-merges the +same configs and re-parses the same interpolation strings. + +Six independent layers, each individually switchable: + +``fast_yaml`` + Back OmegaConf's YAML loader with libyaml's C scanner (~14x faster + parsing). No-op when PyYAML was built without libyaml. +``repo_cache`` + Keep parsed configs -- and the defaults lists derived from them -- in + memory, keyed by file identity. +``lookup_cache`` + Remember which search path a config or group name resolves to, so repeated + "is this a config group?" questions stop hitting the filesystem. +``compose_cache`` + Memoize merging a defaults list into a config. Sweep points that differ + only in ``++key=value`` overrides share one merge. +``parse_cache`` + Memoize OmegaConf's ANTLR parse trees for interpolation strings. A sweep + typically parses a couple of dozen distinct strings thousands of times. +``override_cache`` + Memoize parsing of command-line override strings. Staged sweeps pass the + whole sibling config down as hundreds of ``++key=value`` overrides per + point, and Hydra builds a fresh ANTLR parser for each one. + +Cache entries are revalidated against ``stat()`` on every lookup, so editing a +config on disk invalidates exactly the entries that depend on it. Composition +results are unchanged; this module only avoids repeating work. +""" + +from __future__ import annotations + +import copy +import logging +import os +import sys +from typing import Any + +import yaml + +LOGGER = logging.getLogger(__name__) + +__all__ = ["clear", "disable", "enable", "reset_stats", "stats"] + +_INSTALL_MARKER = "_hydra_staged_sweep_cache_installed" + +_stats: dict[str, int] = { + "repo_hit": 0, + "repo_miss": 0, + "repo_skip": 0, + "compose_hit": 0, + "compose_miss": 0, + "parse_hit": 0, + "parse_miss": 0, + "override_hit": 0, + "override_miss": 0, + "lookup_hit": 0, + "lookup_miss": 0, +} + + +def stats() -> dict[str, int]: + """Hit/miss counters for each layer.""" + return dict(_stats) + + +def reset_stats() -> None: + for key in _stats: + _stats[key] = 0 + + +# --------------------------------------------------------------------------- +# shared helpers +# --------------------------------------------------------------------------- +def _fingerprint(path: str) -> tuple[int, int] | None: + try: + st = os.stat(path) + except OSError: + return None + return (st.st_mtime_ns, st.st_size) + + +def _resolved_file(source: Any, config_path: str) -> str | None: + """Absolute path backing ``config_path`` in ``source``, if file-backed.""" + if source.scheme() != "file": + return None + return os.path.realpath(os.path.join(source.path, source._normalize_file_name(config_path))) + + +# --------------------------------------------------------------------------- +# 1. libyaml +# --------------------------------------------------------------------------- +def _install_fast_yaml() -> bool: + if not hasattr(yaml, "CSafeLoader"): + LOGGER.debug("PyYAML built without libyaml; keeping the pure-Python loader") + return False + + from omegaconf import _utils + + if getattr(_utils.get_yaml_loader, "_hss_fast", False): + return True + + slow = _utils.get_yaml_loader() + + class OmegaConfCLoader(yaml.CSafeLoader): # type: ignore[misc,valid-type] + pass + + # Carry over OmegaConf's constructors and implicit resolvers so values are + # typed exactly as before (it overrides timestamp and bool handling). + OmegaConfCLoader.yaml_constructors = dict(slow.yaml_constructors) + OmegaConfCLoader.yaml_multi_constructors = dict(slow.yaml_multi_constructors) + OmegaConfCLoader.yaml_implicit_resolvers = { + k: list(v) for k, v in slow.yaml_implicit_resolvers.items() + } + + def get_yaml_loader() -> Any: + return OmegaConfCLoader + + get_yaml_loader._hss_fast = True # type: ignore[attr-defined] + _utils.get_yaml_loader = get_yaml_loader + import omegaconf.omegaconf as _oc + + if hasattr(_oc, "get_yaml_loader"): + _oc.get_yaml_loader = get_yaml_loader + return True + + +# --------------------------------------------------------------------------- +# 2. repository-level cache +# --------------------------------------------------------------------------- +# key -> (file fingerprint, ConfigResult, must_copy) +_repo_cache: dict[Any, tuple[tuple[int, int] | None, Any, bool]] = {} +_orig_repo_load: Any = None + + +def _has_matching_schema(repo: Any, config_path: str) -> bool: + """True if a ConfigStore schema shares this config's name. + + Only then does Hydra's deprecated automatic schema matching run, and that + is the one code path that mutates a loaded config in place (it pops + ``hydra`` out of the primary config). Everywhere else the loaded config is + treated as read-only -- Hydra's own per-compose ``CachingConfigRepository`` + already hands the same object to several callers -- so entries can be + shared instead of copied. + """ + from hydra.plugins.config_source import ConfigSource + + try: + source = repo.get_schema_source() + return bool(source.is_config(ConfigSource._normalize_file_name(config_path))) + except Exception: # noqa: BLE001 - any failure here means "copy, don't share" + return True + + +def _install_repo_cache() -> None: + global _orig_repo_load + from hydra._internal.config_repository import ConfigRepository + + if _orig_repo_load is not None: + return + _orig_repo_load = ConfigRepository.load_config + + def load_config(self: ConfigRepository, config_path: str) -> Any: + from hydra.core.object_type import ObjectType + + source = self._find_object_source(config_path, ObjectType.CONFIG) + if source is None: + return _orig_repo_load(self, config_path) + + # Structured configs live in the ConfigStore, can be mutated at runtime + # and are cheap to build. Leave them alone. + if source.scheme() == "structured": + _stats["repo_skip"] += 1 + return _orig_repo_load(self, config_path) + + path = _resolved_file(source, config_path) + fingerprint = _fingerprint(path) if path is not None else None + if path is not None and fingerprint is None: + return _orig_repo_load(self, config_path) # vanished; let Hydra report it + + key = (config_path, source.scheme(), source.provider, source.path) + entry = _repo_cache.get(key) + if entry is not None and entry[0] == fingerprint: + _stats["repo_hit"] += 1 + return copy.deepcopy(entry[1]) if entry[2] else entry[1] + + _stats["repo_miss"] += 1 + result = _orig_repo_load(self, config_path) + must_copy = _has_matching_schema(self, config_path) + _repo_cache[key] = (fingerprint, copy.deepcopy(result) if must_copy else result, must_copy) + return result + + ConfigRepository.load_config = load_config # type: ignore[assignment] + + +# --------------------------------------------------------------------------- +# 2b. config-group lookup cache +# --------------------------------------------------------------------------- +# Resolving the defaults list asks "is this a config group?" once per override +# per sweep point, and each answer costs a realpath() plus a stat() on every +# search path. The answers only change if the config tree is edited on disk +# mid-run, which a sweep build does not do. +_lookup_cache: dict[Any, Any] = {} +_orig_find_source: Any = None + + +def _install_lookup_cache() -> None: + global _orig_find_source + from hydra._internal.config_repository import ConfigRepository + + if _orig_find_source is not None: + return + _orig_find_source = ConfigRepository._find_object_source + + def _find_object_source(self: ConfigRepository, config_path: str, object_type: Any) -> Any: + key = ( + config_path, + object_type, + tuple((s.scheme(), s.provider, s.path) for s in self.sources), + ) + if key in _lookup_cache: + _stats["lookup_hit"] += 1 + index = _lookup_cache[key] + # Cache the position, not the source object: sources are rebuilt + # for every composition. + return None if index is None else self.sources[index] + _stats["lookup_miss"] += 1 + source = _orig_find_source(self, config_path, object_type) + _lookup_cache[key] = None if source is None else self.sources.index(source) + return source + + ConfigRepository._find_object_source = _find_object_source # type: ignore[assignment] + + +# --------------------------------------------------------------------------- +# 3. composed-config cache +# --------------------------------------------------------------------------- +# key -> (contributing file fingerprints, composed DictConfig) +_compose_cache: dict[Any, tuple[tuple[Any, ...], Any]] = {} +_orig_compose: Any = None + + +# Structured configs are not file-backed, so a composition that merges one in +# cannot be revalidated by stat(). Bump a generation counter whenever the +# ConfigStore changes and key the cache on it instead. +_store_generation = [0] +_orig_store: Any = None + + +def _install_store_watch() -> None: + global _orig_store + from hydra.core.config_store import ConfigStore + + if _orig_store is not None: + return + _orig_store = ConfigStore.store + + def store(self: ConfigStore, *args: Any, **kwargs: Any) -> Any: + _store_generation[0] += 1 + return _orig_store(self, *args, **kwargs) + + ConfigStore.store = store # type: ignore[assignment] + + +def _defaults_key(defaults: list[Any], repo: Any) -> Any: + return ( + tuple((d.config_path, d.parent, d.package, d.is_self, d.primary) for d in defaults), + tuple((s.scheme(), s.provider, s.path) for s in repo.get_sources()), + _store_generation[0], + ) + + +def _contributing_files(defaults: list[Any], repo: Any) -> tuple[Any, ...]: + """Fingerprints of every file that can feed this composition.""" + found = set() + sources = [s for s in repo.get_sources() if s.scheme() == "file"] + for default in defaults: + if default.config_path is None: + continue + for source in sources: + path = _resolved_file(source, default.config_path) + if path is None: + continue + fingerprint = _fingerprint(path) + if fingerprint is not None: + found.add((path, fingerprint)) + return tuple(sorted(found)) + + +def _install_compose_cache() -> None: + global _orig_compose + from hydra._internal.config_loader_impl import ConfigLoaderImpl + + if _orig_compose is not None: + return + _orig_compose = ConfigLoaderImpl._compose_config_from_defaults_list + + def _compose(self: ConfigLoaderImpl, defaults: list[Any], repo: Any) -> Any: + key = _defaults_key(defaults, repo) + files = _contributing_files(defaults, repo) + entry = _compose_cache.get(key) + if entry is not None and entry[0] == files: + _stats["compose_hit"] += 1 + # The caller mutates this (struct flag, overrides, hydra bookkeeping). + return copy.deepcopy(entry[1]) + _stats["compose_miss"] += 1 + cfg = _orig_compose(self, defaults, repo) + _compose_cache[key] = (files, copy.deepcopy(cfg)) + return cfg + + ConfigLoaderImpl._compose_config_from_defaults_list = _compose # type: ignore[assignment] + + +# --------------------------------------------------------------------------- +# 4. interpolation parse-tree cache +# --------------------------------------------------------------------------- +_parse_cache: dict[Any, Any] = {} +_orig_parse: Any = None + + +def _materialize(node: Any) -> None: + """Force every token in a parse tree to hold its own text. + + antlr's ``CommonToken.text`` lazily reads back from the lexer's input + stream, and OmegaConf swaps that stream out on the next parse. Reading the + text once here pins it, which makes the tree safe to keep and reuse. + """ + token = getattr(node, "symbol", None) + if token is not None: + token.text = token.text + for attr in ("start", "stop"): + token = getattr(node, attr, None) + if token is not None and hasattr(token, "text"): + token.text = token.text + for child in getattr(node, "children", None) or (): + _materialize(child) + + +def _install_parse_cache(maxsize: int = 4096) -> None: + global _orig_parse + from omegaconf import grammar_parser + + if _orig_parse is not None: + return + _orig_parse = grammar_parser.parse + + def parse( + value: str, parser_rule: str = "configValue", lexer_mode: str = "DEFAULT_MODE" + ) -> Any: + key = (value, parser_rule, lexer_mode) + tree = _parse_cache.get(key) + if tree is not None: + _stats["parse_hit"] += 1 + return tree + _stats["parse_miss"] += 1 + tree = _orig_parse(value, parser_rule, lexer_mode) + _materialize(tree) + if len(_parse_cache) < maxsize: + _parse_cache[key] = tree + return tree + + grammar_parser.parse = parse + # omegaconf.base and omegaconf._utils imported the symbol directly. + for modname in ("omegaconf.base", "omegaconf._utils"): + module = sys.modules.get(modname) + if module is not None and getattr(module, "parse", None) is _orig_parse: + module.parse = parse + + +# --------------------------------------------------------------------------- +# 5. override-parse cache +# --------------------------------------------------------------------------- +# Hydra builds a fresh ANTLR lexer and parser for *every* command-line +# override. A staged sweep passes the whole sibling config down as a few +# hundred ``++sibling..a.b.c=value`` overrides per point, nearly all of +# them repeated across points. +_override_cache: dict[Any, Any] = {} +_orig_parse_rule: Any = None + + +def _functions_key(parser: Any) -> Any: + """Identify the grammar function set, so custom functions can't collide.""" + key = getattr(parser, "_hss_functions_key", None) + if key is None: + try: + key = tuple(sorted(parser.functions.definitions)) + except (AttributeError, TypeError): + key = id(parser.functions) + parser._hss_functions_key = key + return key + + +def _install_override_cache(maxsize: int = 16384) -> None: + global _orig_parse_rule + from hydra.core.override_parser.overrides_parser import OverridesParser + + if _orig_parse_rule is not None: + return + _orig_parse_rule = OverridesParser.parse_rule + + def parse_rule(self: OverridesParser, s: str, rule_name: str) -> Any: + key = (s, rule_name, _functions_key(self)) + cached = _override_cache.get(key) + if cached is not None: + _stats["override_hit"] += 1 + return copy.deepcopy(cached) + _stats["override_miss"] += 1 + result = _orig_parse_rule(self, s, rule_name) + if len(_override_cache) < maxsize: + _override_cache[key] = copy.deepcopy(result) + return result + + OverridesParser.parse_rule = parse_rule # type: ignore[assignment] + + +# --------------------------------------------------------------------------- +# public API +# --------------------------------------------------------------------------- +_enabled = False + + +def enable( + *, + fast_yaml: bool = True, + repo_cache: bool = True, + compose_cache: bool = True, + parse_cache: bool = True, + override_cache: bool = True, + lookup_cache: bool = True, +) -> None: + """Install the caches. Idempotent; safe to call from every entry point. + + Set ``HYDRA_STAGED_SWEEP_CACHE=0`` to turn the whole thing off without + touching call sites. + """ + global _enabled + if _enabled or os.environ.get("HYDRA_STAGED_SWEEP_CACHE", "1") == "0": + return + # A vendored copy of this module and an installed one would otherwise each + # wrap Hydra, stacking a redundant layer on every call. Mark the target. + import hydra + + if getattr(hydra, _INSTALL_MARKER, False): + _enabled = True + return + setattr(hydra, _INSTALL_MARKER, True) + if fast_yaml: + _install_fast_yaml() + if repo_cache: + _install_repo_cache() + if lookup_cache: + _install_lookup_cache() + if compose_cache: + _install_store_watch() + _install_compose_cache() + if parse_cache: + _install_parse_cache() + if override_cache: + _install_override_cache() + _enabled = True + LOGGER.debug("hydra config caches enabled") + + +def disable() -> None: + """Restore Hydra's and OmegaConf's original behaviour.""" + global _enabled, _orig_repo_load, _orig_compose, _orig_parse, _orig_parse_rule + global _orig_find_source, _orig_store + if _orig_repo_load is not None: + from hydra._internal.config_repository import ConfigRepository + + ConfigRepository.load_config = _orig_repo_load # type: ignore[assignment] + _orig_repo_load = None + if _orig_find_source is not None: + from hydra._internal.config_repository import ConfigRepository + + ConfigRepository._find_object_source = _orig_find_source # type: ignore[assignment] + _orig_find_source = None + if _orig_compose is not None: + from hydra._internal.config_loader_impl import ConfigLoaderImpl + + ConfigLoaderImpl._compose_config_from_defaults_list = _orig_compose # type: ignore[assignment] + _orig_compose = None + if _orig_store is not None: + from hydra.core.config_store import ConfigStore + + ConfigStore.store = _orig_store # type: ignore[assignment] + _orig_store = None + if _orig_parse is not None: + from omegaconf import grammar_parser + + grammar_parser.parse = _orig_parse + for modname in ("omegaconf.base", "omegaconf._utils"): + module = sys.modules.get(modname) + if module is not None: + module.parse = _orig_parse + _orig_parse = None + if _orig_parse_rule is not None: + from hydra.core.override_parser.overrides_parser import OverridesParser + + OverridesParser.parse_rule = _orig_parse_rule # type: ignore[assignment] + _orig_parse_rule = None + import hydra + + if hasattr(hydra, _INSTALL_MARKER): + delattr(hydra, _INSTALL_MARKER) + clear() + _enabled = False + + +def clear() -> None: + """Drop every cached entry (the caches refill on the next composition).""" + _repo_cache.clear() + _compose_cache.clear() + _parse_cache.clear() + _override_cache.clear() + _lookup_cache.clear() diff --git a/oellm_autoexp/hydra_staged_sweep/config/loader.py b/oellm_autoexp/hydra_staged_sweep/config/loader.py index 60c5916c..65d43e53 100644 --- a/oellm_autoexp/hydra_staged_sweep/config/loader.py +++ b/oellm_autoexp/hydra_staged_sweep/config/loader.py @@ -10,9 +10,10 @@ from compoconf import parse_config, ConfigInterface from hydra import compose, initialize_config_dir -from omegaconf import OmegaConf +from omegaconf import OmegaConf, open_dict from . import schema +from .cache import enable as enable_config_cache from .resolvers import register_default_resolvers LOGGER = logging.getLogger(__file__) @@ -20,6 +21,9 @@ T = TypeVar("T", bound=ConfigInterface) register_default_resolvers() +# Sweeps compose the same config tree once per point; without this every point +# re-reads and re-merges the whole tree. See config/cache.py. +enable_config_cache() class ConfigLoaderError(RuntimeError): @@ -62,11 +66,29 @@ def load_config(path: str | Path, config_class: type[T] = schema.StagedSweepRoot return _parse_root(data, config_class, str(path), str(path.parent)) +def _merge_extra_config(cfg: Any, extra_config: Mapping[str, Any] | None) -> None: + """Force-add ``extra_config`` into the composed config, in place. + + Equivalent to passing the same data as ``++key.path=value`` overrides, but + as a single merge. Hydra parses and applies overrides one at a time, which + is far too slow for the few hundred entries a resolved sibling config + expands into. + """ + if not extra_config: + return + # Accepts an already-built container so callers that reuse the same context + # across configs can build it once; merging does not modify the source. + source = extra_config if OmegaConf.is_config(extra_config) else OmegaConf.create(dict(extra_config)) + with open_dict(cfg): + cfg.merge_with(source) + + def load_hydra_config( config_name: str, config_dir: str | Path, overrides: Iterable[str] | None = None, config_class: type[T] = schema.StagedSweepRoot, + extra_config: Mapping[str, Any] | None = None, ) -> T: LOGGER.info(f"Loading Hydra config: {config_name} from {config_dir}") register_default_resolvers() @@ -82,6 +104,8 @@ def load_hydra_config( with initialize_config_dir(version_base=None, config_dir=str(config_dir)): cfg = compose(config_name=config_name, overrides=overrides) + _merge_extra_config(cfg, extra_config) + data = OmegaConf.to_container(cfg, resolve=True) # type: ignore[return-value] if not isinstance(data, Mapping): raise ConfigLoaderError(f"Hydra config {config_name} did not produce a mapping") @@ -95,6 +119,7 @@ def load_config_reference( config_dir: str | Path | None = None, overrides: Iterable[str] | None = None, config_class: type[T] = schema.StagedSweepRoot, + extra_config: Mapping[str, Any] | None = None, ) -> T: if config_name is None and config_path: path = Path(config_path) @@ -116,6 +141,7 @@ def load_config_reference( original_config_dir or config_dir, combined_overrides, config_class=config_class, + extra_config=extra_config, ) except Exception as exc: LOGGER.warning( @@ -126,6 +152,8 @@ def load_config_reference( with initialize_config_dir(version_base=None, config_dir=os.path.abspath(path.parent)): cfg = compose(config_name=path.name[:-5], overrides=overrides) + _merge_extra_config(cfg, extra_config) + data = OmegaConf.to_container(cfg, resolve=True) if not isinstance(data, Mapping): raise ConfigLoaderError(f"Config file {path} did not produce a mapping") @@ -133,7 +161,9 @@ def load_config_reference( return _parse_root(data, config_class, str(path), str(path.parent)) else: return load_config(path, config_class=config_class) - return load_hydra_config(config_name, config_dir, overrides, config_class=config_class) + return load_hydra_config( + config_name, config_dir, overrides, config_class=config_class, extra_config=extra_config + ) __all__ = [ diff --git a/oellm_autoexp/hydra_staged_sweep/dag_resolver.py b/oellm_autoexp/hydra_staged_sweep/dag_resolver.py index 363ddcaf..a478c8ff 100644 --- a/oellm_autoexp/hydra_staged_sweep/dag_resolver.py +++ b/oellm_autoexp/hydra_staged_sweep/dag_resolver.py @@ -283,6 +283,31 @@ def dict_to_cmdlines(dct: dict | list | str | int | float, prefix: str = ""): return cmdline_opts +def drop_cmdline_invisible(value: Any) -> Any: + """Strip what ``config_to_cmdline`` cannot express, so a direct merge matches it. + + An empty mapping flattens to zero overrides, so round-tripping a config + through the command line silently drops it -- and a list element that is an + empty mapping is left as the placeholder index ``config_to_cmdline`` emits + for it. Applying the same pruning before merging keeps the two paths in + exact agreement. + """ + if isinstance(value, Mapping): + pruned = {key: drop_cmdline_invisible(item) for key, item in value.items()} + return { + key: item + for key, item in pruned.items() + if not (isinstance(item, Mapping) and not item) + } + if isinstance(value, (list, ListConfig, Sequence)) and not isinstance(value, (str, bytes)): + out = [] + for index, item in enumerate(value): + pruned = drop_cmdline_invisible(item) + out.append(index if isinstance(pruned, Mapping) and not pruned else pruned) + return out + return value + + def is_config_group(key: str, config_dir: str | Path | None) -> bool: """Check if a parameter key refers to a Hydra config group. @@ -370,6 +395,7 @@ def resolve_sweep_with_dag( resolved_jobs = {} filtered_jobs = {} + sibling_context_cache: dict[tuple, tuple] = {} base_context = asdict(config) base_context = {k: v for k, v in base_context.items() if k not in ("sweep", "sibling")} sweep_filter_expr = config.sweep.filter if isinstance(config.sweep, SweepConfig) else True @@ -410,39 +436,50 @@ def resolve_sweep_with_dag( sibling_patterns = sibling_index.patterns_by_idx.get(point_idx, set()) sibling_jobs = {} + sibling_ids = [] for pattern in sibling_patterns: sibling_point = find_sibling_by_group_path( point, points_dict, pattern, sibling_index=sibling_index ) if sibling_point and sibling_point.index in resolved_jobs: sibling_jobs[pattern] = resolved_jobs[sibling_point.index] - - sibling_job_configs = { - sibling_pattern: asdict( - load_config_reference( - config_dir=config_setup.config_dir, - config_path=config_setup.config_path, - config_name=config_setup.config_name, - overrides=list(config_setup.overrides) + sibling_job.parameters, - config_class=config_class, - ) - ) - for sibling_pattern, sibling_job in sibling_jobs.items() - } - - for sibling_pattern in sibling_job_configs: - if "sweep" in sibling_job_configs[sibling_pattern]: - del sibling_job_configs[sibling_pattern]["sweep"] - - cmdline_overrides_siblings = config_to_cmdline( - { + sibling_ids.append((pattern, sibling_point.index)) + + # Every point in a stage chain sees the same siblings, and flattening a + # resolved sibling config is the most expensive step in this loop, so + # build the context once per distinct set of siblings. + context_key = tuple(sorted(sibling_ids)) + cached_context = sibling_context_cache.get(context_key) + if cached_context is None: + # The sibling was already resolved with exactly these overrides -- + # its JobPlan holds the result. Recomposing it here doubled the + # amount of Hydra composition every staged sweep had to do. + sibling_job_configs = { + sibling_pattern: asdict(sibling_job.config) + for sibling_pattern, sibling_job in sibling_jobs.items() + } + + for sibling_pattern in sibling_job_configs: + if "sweep" in sibling_job_configs[sibling_pattern]: + del sibling_job_configs[sibling_pattern]["sweep"] + + sibling_context = { "sibling": { sibling_job.get("stage", "unknown"): sibling_job for sibling_job in sibling_job_configs.values() } - }, - override="++", - ) + } + # Flatten the context as-is so job_parameters is byte-identical to + # what a command-line run would carry; merge the pruned form, which + # is what those overrides actually produce once applied. + cached_context = ( + config_to_cmdline(sibling_context, override="++"), + OmegaConf.create(drop_cmdline_invisible(sibling_context)) + if sibling_job_configs + else None, + ) + sibling_context_cache[context_key] = cached_context + cmdline_overrides_siblings, sibling_container = cached_context # Generate parameter overrides with smart prefix selection param_overrides = [] @@ -459,6 +496,9 @@ def resolve_sweep_with_dag( param_to_cmdlines(key, value, prefix="++", config_dir=config_setup.config_dir) ) + compose_overrides = ( + list(config_setup.overrides) + [f"++index={point_idx}"] + param_overrides + ) job_parameters = ( list(config_setup.overrides) + cmdline_overrides_siblings @@ -466,11 +506,15 @@ def resolve_sweep_with_dag( + param_overrides ) + # job_parameters still records the sibling context as ++overrides so the + # job stays reproducible from the command line, but composing with it is + # far cheaper as a single merge than as several hundred parsed overrides. resolved = load_config_reference( config_dir=config_setup.config_dir, config_path=config_setup.config_path, config_name=config_setup.config_name, - overrides=job_parameters, + overrides=compose_overrides, + extra_config=sibling_container, config_class=config_class, ) diff --git a/tests/hydra_staged_sweep/test_config_cache.py b/tests/hydra_staged_sweep/test_config_cache.py new file mode 100644 index 00000000..fe60fada --- /dev/null +++ b/tests/hydra_staged_sweep/test_config_cache.py @@ -0,0 +1,167 @@ +"""The composition caches must not change what composition produces.""" + +import textwrap +from dataclasses import dataclass, field +from typing import Any + +import pytest +from omegaconf import OmegaConf + +from oellm_autoexp.hydra_staged_sweep.config import cache +from oellm_autoexp.hydra_staged_sweep.config.loader import load_hydra_config +from oellm_autoexp.hydra_staged_sweep.config.schema import StagedSweepRoot +from oellm_autoexp.hydra_staged_sweep.dag_resolver import config_to_cmdline, drop_cmdline_invisible + + +@dataclass(kw_only=True) +class CacheTestConfig(StagedSweepRoot): + name: str = "" + depth: int = 0 + label: str = "" + nested: dict[str, Any] = field(default_factory=dict) + metadata: dict[str, Any] = field(default_factory=dict) + + +@pytest.fixture +def config_dir(tmp_path): + conf = tmp_path / "conf" + (conf / "group").mkdir(parents=True) + (conf / "group" / "a.yaml").write_text("# @package _global_\ndepth: 1\n") + (conf / "group" / "b.yaml").write_text("# @package _global_\ndepth: 2\n") + (conf / "config.yaml").write_text( + textwrap.dedent("""\ + defaults: + - group: a + - _self_ + + name: base + label: "${name}-${depth}" + nested: + on: yes + off: no + when: 2020-01-01 + num: 1.0e-4 + """) + ) + return conf + + +@pytest.fixture +def fresh_cache(): + """Each test starts from an empty, freshly installed set of caches.""" + cache.disable() + cache.enable() + cache.reset_stats() + yield cache + cache.disable() + + +def _load(config_dir, overrides=None): + return load_hydra_config( + "config", config_dir, overrides or [], config_class=CacheTestConfig + ) + + +def test_cached_composition_matches_uncached(config_dir): + cache.disable() + uncached = [_load(config_dir, [f"++depth={i}"]) for i in range(4)] + + cache.enable() + try: + cached = [_load(config_dir, [f"++depth={i}"]) for i in range(4)] + finally: + cache.disable() + + assert [c.label for c in cached] == [u.label for u in uncached] + assert [c.nested for c in cached] == [u.nested for u in uncached] + + +def test_group_overrides_are_not_conflated(config_dir, fresh_cache): + """Different config groups must not share a cached composition.""" + assert _load(config_dir, ["group=a"]).depth == 1 + assert _load(config_dir, ["group=b"]).depth == 2 + assert _load(config_dir, ["group=a"]).depth == 1 + + +def test_editing_a_config_invalidates_the_cache(config_dir, fresh_cache): + assert _load(config_dir).name == "base" + + target = config_dir / "config.yaml" + target.write_text(target.read_text().replace("name: base", "name: edited") + "\n# pad\n") + + assert _load(config_dir).name == "edited" + + +def test_caches_record_hits(config_dir, fresh_cache): + for _ in range(3): + _load(config_dir) + stats = fresh_cache.stats() + assert stats["repo_hit"] > 0 + assert stats["compose_hit"] > 0 + + +def test_enable_is_idempotent(config_dir, fresh_cache): + before = _load(config_dir).label + fresh_cache.enable() + fresh_cache.enable() + assert _load(config_dir).label == before + + +def test_disable_restores_hydra(config_dir): + cache.enable() + _load(config_dir) + cache.disable() + assert cache.stats() is not None + # Composition still works with the caches removed. + assert _load(config_dir).name == "base" + + +def test_fast_yaml_types_match_the_pure_python_loader(tmp_path): + """libyaml must type scalars exactly as OmegaConf's own loader does.""" + sample = tmp_path / "s.yaml" + sample.write_text( + "a: yes\nb: no\nc: null\nd: 2020-01-01\ne: 1.0e-4\nf: '010'\ng: 010\nh: ~\n" + ) + cache.disable() + plain = OmegaConf.to_container(OmegaConf.load(sample)) + cache.enable() + try: + fast = OmegaConf.to_container(OmegaConf.load(sample)) + finally: + cache.disable() + assert fast == plain + + +def test_repeated_interpolations_resolve_identically(config_dir, fresh_cache): + """The parse-tree cache hands out one shared tree; resolution must be stable.""" + first = _load(config_dir, ["++depth=7"]).label + second = _load(config_dir, ["++depth=8"]).label + assert (first, second) == ("base-7", "base-8") + + +@pytest.mark.parametrize( + "payload", + [ + {"empty": {}}, + {"outer": {"inner": {}}}, + {"kept": 1, "dropped": {}}, + {"items": [1, {}, {"b": 2}]}, + {"items": []}, + {"nothing": None}, + {"text": "value"}, + {"interp": "${name}"}, + {"quoted": 'has "quotes"'}, + ], +) +def test_extra_config_matches_the_override_round_trip(config_dir, fresh_cache, payload): + """Merging the context directly must land exactly where the ++overrides do.""" + value = {"nested": payload} + via_overrides = _load(config_dir, config_to_cmdline(value, override="++")) + via_merge = load_hydra_config( + "config", + config_dir, + [], + config_class=CacheTestConfig, + extra_config=drop_cmdline_invisible(value), + ) + assert via_merge.nested == via_overrides.nested From 2aa08fc759a21d46b4d30dcd4260fc534b2a7fb6 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Korbinian=20P=C3=B6ppel?= Date: Sun, 2 Aug 2026 02:22:36 +0200 Subject: [PATCH 2/4] Resolve independent sweep chains in a process pool Mirrors the change in hydra_staged_sweep into the vendored copy, and fans out script rendering in the orchestrator the same way. Sweep resolution is pure CPU: composing, resolving interpolations and building dataclasses. A sweep splits into independent dependency chains -- a stable stage and the cooldowns branching off it must run in order, but separate chains share nothing -- so resolve_sweep_with_dag hands the chains to a fork pool. Forking means workers inherit the loaded modules, the registered resolvers and the warm config caches, so only each chain and its results cross the process boundary. compoconf configs cannot pickle themselves -- ConfigInterface.__reduce__ reduces to (cls, (), state), so unpickling calls cls() with no arguments, which raises for any config with required fields, and its asdict state flattens nested configs into plain dicts. Plans are pickled by __dict__ instead. Two more caches, both of which the pool then multiplies: the Defaults List is memoized on the overrides that can actually select a config group, and CachingConfigRepository no longer deep-copies the whole repository on every composition. Filter evaluation builds its context lazily, since a filter that is already a bool never reads it. HYDRA_STAGED_SWEEP_WORKERS pins the pool size; 1 keeps everything in-process. run_autoexp.py --dry-run on the 300-job titan_qwen3_depth_width_scaling sweep goes from 82s to 21s -- 470s before any of this work -- with all 300 rendered sbatch scripts still byte-identical. --- .../hydra_staged_sweep/config/cache.py | 158 +++++++++++- .../hydra_staged_sweep/dag_resolver.py | 229 +++++++++++++----- oellm_autoexp/hydra_staged_sweep/parallel.py | 72 ++++++ oellm_autoexp/orchestrator.py | 45 +++- .../test_parallel_resolution.py | 140 +++++++++++ 5 files changed, 565 insertions(+), 79 deletions(-) create mode 100644 oellm_autoexp/hydra_staged_sweep/parallel.py create mode 100644 tests/hydra_staged_sweep/test_parallel_resolution.py diff --git a/oellm_autoexp/hydra_staged_sweep/config/cache.py b/oellm_autoexp/hydra_staged_sweep/config/cache.py index 776dc71f..42326d38 100644 --- a/oellm_autoexp/hydra_staged_sweep/config/cache.py +++ b/oellm_autoexp/hydra_staged_sweep/config/cache.py @@ -5,7 +5,7 @@ re-parses the same YAML files, re-walks the same defaults list, re-merges the same configs and re-parses the same interpolation strings. -Six independent layers, each individually switchable: +Seven independent layers, each individually switchable: ``fast_yaml`` Back OmegaConf's YAML loader with libyaml's C scanner (~14x faster @@ -19,6 +19,9 @@ ``compose_cache`` Memoize merging a defaults list into a config. Sweep points that differ only in ``++key=value`` overrides share one merge. +``defaults_cache`` + Memoize the Defaults List itself, keyed on the overrides that can actually + select a config group. ``parse_cache`` Memoize OmegaConf's ANTLR parse trees for interpolation strings. A sweep typically parses a couple of dozen distinct strings thousands of times. @@ -60,6 +63,8 @@ "override_miss": 0, "lookup_hit": 0, "lookup_miss": 0, + "defaults_hit": 0, + "defaults_miss": 0, } @@ -197,6 +202,45 @@ def load_config(self: ConfigRepository, config_path: str) -> Any: ConfigRepository.load_config = load_config # type: ignore[assignment] +# --------------------------------------------------------------------------- +# 2a. lazy repository copy +# --------------------------------------------------------------------------- +_orig_caching_init: Any = None +_orig_caching_initialize_sources: Any = None + + +def _install_lazy_repo_copy() -> None: + """Stop copying the whole repository on every composition. + + ``CachingConfigRepository`` deep-copies its delegate up front so that + ``initialize_sources()`` cannot mutate the loader's shared repository. That + only happens when a config overrides ``hydra.searchpath``, so defer the copy + until it is actually needed. + """ + global _orig_caching_init, _orig_caching_initialize_sources + from hydra._internal.config_repository import CachingConfigRepository + + if _orig_caching_init is not None: + return + _orig_caching_init = CachingConfigRepository.__init__ + _orig_caching_initialize_sources = CachingConfigRepository.initialize_sources + orig_initialize = _orig_caching_initialize_sources + + def __init__(self: CachingConfigRepository, delegate: Any) -> None: + self.delegate = delegate + self.cache = {} + self._hss_owns_delegate = False + + def initialize_sources(self: CachingConfigRepository, config_search_path: Any) -> None: + if not getattr(self, "_hss_owns_delegate", False): + self.delegate = copy.deepcopy(self.delegate) + self._hss_owns_delegate = True + orig_initialize(self, config_search_path) + + CachingConfigRepository.__init__ = __init__ # type: ignore[assignment] + CachingConfigRepository.initialize_sources = initialize_sources # type: ignore[assignment] + + # --------------------------------------------------------------------------- # 2b. config-group lookup cache # --------------------------------------------------------------------------- @@ -315,6 +359,95 @@ def _compose(self: ConfigLoaderImpl, defaults: list[Any], repo: Any) -> Any: ConfigLoaderImpl._compose_config_from_defaults_list = _compose # type: ignore[assignment] +# --------------------------------------------------------------------------- +# 3b. defaults-list cache +# --------------------------------------------------------------------------- +# Walking the Defaults List means loading every config in the tree to read its +# own defaults and package header. Only overrides that select a config group +# can change the outcome; the ``++key=value`` overrides a sweep varies per +# point cannot, so every point in a sweep rebuilds the same list. +_defaults_list_cache: dict[Any, tuple[tuple[Any, ...], Any]] = {} +_orig_create_defaults_list: Any = None + + +def _install_defaults_list_cache() -> None: + global _orig_create_defaults_list + from hydra._internal import config_loader_impl + from hydra._internal import defaults_list as defaults_list_module + from hydra._internal.defaults_list import DefaultsList, Overrides + + if _orig_create_defaults_list is not None: + return + _orig_create_defaults_list = defaults_list_module.create_defaults_list + + def create_defaults_list( + repo: Any, + config_name: str | None, + overrides_list: list[Any], + prepend_hydra: bool, + skip_missing: bool, + ) -> Any: + # Overrides() is what decides group override vs value override, so build + # it first and key the cache on the group-affecting ones only. + overrides = Overrides(repo=repo, overrides_list=overrides_list) + value_overrides = {id(override) for override in overrides.config_overrides} + selecting = tuple( + override.input_line + for override in overrides_list + if id(override) not in value_overrides + ) + key = ( + config_name, + prepend_hydra, + skip_missing, + selecting, + tuple((s.scheme(), s.provider, s.path) for s in repo.get_sources()), + _store_generation[0], + ) + + entry = _defaults_list_cache.get(key) + if entry is not None: + defaults, tree, known_choices, known_per_group = entry[1] + if entry[0] == _contributing_files(defaults, repo): + _stats["defaults_hit"] += 1 + # Rebuilt per call: these depend on every override, not just the + # selecting ones. known_choices is filled by the tree walk we + # just skipped, so restore it from the cached run. + overrides.known_choices = dict(known_choices) + overrides.known_choices_per_group = { + group: set(choices) for group, choices in known_per_group.items() + } + return DefaultsList( + defaults=defaults, + defaults_tree=tree, + config_overrides=overrides.config_overrides, + overrides=overrides, + ) + + _stats["defaults_miss"] += 1 + # A miss re-runs the real thing, including the validation that reports + # unused overrides -- so a hit can only happen for a set that validated. + result = _orig_create_defaults_list( + repo, config_name, overrides_list, prepend_hydra, skip_missing + ) + _defaults_list_cache[key] = ( + _contributing_files(result.defaults, repo), + ( + result.defaults, + result.defaults_tree, + dict(result.overrides.known_choices), + { + group: set(choices) + for group, choices in result.overrides.known_choices_per_group.items() + }, + ), + ) + return result + + defaults_list_module.create_defaults_list = create_defaults_list + config_loader_impl.create_defaults_list = create_defaults_list + + # --------------------------------------------------------------------------- # 4. interpolation parse-tree cache # --------------------------------------------------------------------------- @@ -428,6 +561,7 @@ def enable( fast_yaml: bool = True, repo_cache: bool = True, compose_cache: bool = True, + defaults_cache: bool = True, parse_cache: bool = True, override_cache: bool = True, lookup_cache: bool = True, @@ -452,11 +586,15 @@ def enable( _install_fast_yaml() if repo_cache: _install_repo_cache() + _install_lazy_repo_copy() if lookup_cache: _install_lookup_cache() if compose_cache: _install_store_watch() _install_compose_cache() + if defaults_cache: + _install_store_watch() + _install_defaults_list_cache() if parse_cache: _install_parse_cache() if override_cache: @@ -468,12 +606,20 @@ def enable( def disable() -> None: """Restore Hydra's and OmegaConf's original behaviour.""" global _enabled, _orig_repo_load, _orig_compose, _orig_parse, _orig_parse_rule - global _orig_find_source, _orig_store + global _orig_find_source, _orig_store, _orig_create_defaults_list, _orig_caching_init if _orig_repo_load is not None: from hydra._internal.config_repository import ConfigRepository ConfigRepository.load_config = _orig_repo_load # type: ignore[assignment] _orig_repo_load = None + if _orig_caching_init is not None: + from hydra._internal.config_repository import CachingConfigRepository + + CachingConfigRepository.__init__ = _orig_caching_init # type: ignore[assignment] + CachingConfigRepository.initialize_sources = ( # type: ignore[assignment] + _orig_caching_initialize_sources + ) + _orig_caching_init = None if _orig_find_source is not None: from hydra._internal.config_repository import ConfigRepository @@ -484,6 +630,13 @@ def disable() -> None: ConfigLoaderImpl._compose_config_from_defaults_list = _orig_compose # type: ignore[assignment] _orig_compose = None + if _orig_create_defaults_list is not None: + from hydra._internal import config_loader_impl + from hydra._internal import defaults_list as defaults_list_module + + defaults_list_module.create_defaults_list = _orig_create_defaults_list + config_loader_impl.create_defaults_list = _orig_create_defaults_list + _orig_create_defaults_list = None if _orig_store is not None: from hydra.core.config_store import ConfigStore @@ -518,3 +671,4 @@ def clear() -> None: _parse_cache.clear() _override_cache.clear() _lookup_cache.clear() + _defaults_list_cache.clear() diff --git a/oellm_autoexp/hydra_staged_sweep/dag_resolver.py b/oellm_autoexp/hydra_staged_sweep/dag_resolver.py index a478c8ff..67600788 100644 --- a/oellm_autoexp/hydra_staged_sweep/dag_resolver.py +++ b/oellm_autoexp/hydra_staged_sweep/dag_resolver.py @@ -11,22 +11,26 @@ from __future__ import annotations +import io import logging +import os +import pickle import re from collections import defaultdict -from collections.abc import Mapping, Sequence +from collections.abc import Callable, Mapping, Sequence from dataclasses import dataclass from itertools import zip_longest from pathlib import Path from typing import Any import networkx as nx -from compoconf import asdict +from compoconf import ConfigInterface, asdict from omegaconf import DictConfig, ListConfig, OmegaConf from .config.schema import StagedSweepRoot, SweepConfig, ConfigSetup from .config.loader import load_config_reference from .expander import SweepPoint +from .parallel import run_chunks, worker_count from .planner import JobPlan LOGGER = logging.getLogger(__file__) @@ -101,13 +105,21 @@ def _build_sibling_index(points: Mapping[int, SweepPoint]) -> SiblingIndex: ) -def _resolve_filter_from_context(filter_expr: Any, context: Mapping[str, Any]) -> bool: +def _resolve_filter_from_context(filter_expr: Any, context: Mapping[str, Any] | Callable) -> bool: + """Evaluate one filter expression. + + ``context`` may be a zero-argument callable, which is only invoked for + filters that actually need it -- building the context means walking the + whole resolved config, and most filters are already a plain bool. + """ if filter_expr is None: return True if isinstance(filter_expr, bool): return filter_expr if not isinstance(filter_expr, str): raise ValueError("sweep.filter must resolve to a bool.") + if callable(context): + context = context() cfg = OmegaConf.create({**context, "sweep": {"filter": filter_expr}}) try: resolved = OmegaConf.to_container(cfg, resolve=True) @@ -368,64 +380,52 @@ def param_to_cmdlines(key: str, val: Any, prefix: str = "", config_dir: str | Pa ) -def resolve_sweep_with_dag( - config: StagedSweepRoot, - points: list[SweepPoint] | dict[int, SweepPoint], - config_setup: ConfigSetup, - config_class: type = StagedSweepRoot, -) -> list[JobPlan]: - """Pure OmegaConf resolution with DAG ordering.""" - LOGGER.info(f"Starting DAG resolution for {len(points)} sweep points") +def _restore_config(cls: type, state: dict) -> Any: + """Rebuild a config from its attribute state, without calling __init__.""" + obj = cls.__new__(cls) + obj.__dict__.update(state) + return obj - if isinstance(points, list): - points_dict = {p.index: p for p in points} - else: - points_dict = points - sibling_index = _build_sibling_index(points_dict) - dag = build_dependency_dag_from_points(points_dict, sibling_index=sibling_index) +class _PlanPickler(pickle.Pickler): + """Pickle resolved plans by state. - if not nx.is_directed_acyclic_graph(dag): - cycles = list(nx.simple_cycles(dag)) - LOGGER.error(f"Circular dependencies detected: {cycles}") - raise ValueError(f"Circular dependencies detected: {cycles}") + compoconf's ``ConfigInterface.__reduce__`` reduces to ``(cls, (), state)``, + so unpickling calls ``cls()`` with no arguments -- which raises for any + config with required fields -- and its state comes from ``asdict``, which + flattens nested configs into plain dicts. Neither survives the trip back + from a worker process, so reduce by ``__dict__`` instead. + """ - ordered_indices = list(nx.topological_sort(dag)) - LOGGER.debug(f"Topological order: {ordered_indices}") + def reducer_override(self, obj: Any) -> Any: + if isinstance(obj, ConfigInterface): + return _restore_config, (type(obj), obj.__dict__) + return NotImplemented - resolved_jobs = {} - filtered_jobs = {} - sibling_context_cache: dict[tuple, tuple] = {} - base_context = asdict(config) - base_context = {k: v for k, v in base_context.items() if k not in ("sweep", "sibling")} - sweep_filter_expr = config.sweep.filter if isinstance(config.sweep, SweepConfig) else True - if not isinstance(config.sweep, SweepConfig): - point = points_dict[list(points_dict)[0]] - resolved = load_config_reference( - config_dir=config_setup.config_dir, - config_path=config_setup.config_path, - config_name=config_setup.config_name, - overrides=point.parameters, - config_class=config_class, - ) +def _dumps(obj: Any) -> bytes: + buffer = io.BytesIO() + _PlanPickler(buffer, protocol=pickle.HIGHEST_PROTOCOL).dump(obj) + return buffer.getvalue() - resolved_dict = asdict(resolved) - context = {k: v for k, v in resolved_dict.items() if k not in ("sweep")} - skip_point = False - stage_name = getattr(resolved, "stage", None) +# Chains of dependent points (a stable stage plus the cooldowns that branch off +# it) must be resolved in order, but separate chains share nothing. They are +# handed to a process pool, which is worth it because resolving a point is pure +# CPU: composing, resolving interpolations and building dataclasses. +_CHAIN_CONTEXT: tuple | None = None - job = JobPlan( - config=resolved, - parameters=point, - sibling_pattern=None, - stage_name=stage_name, - ) - return [job] +def _resolve_chain(chain: list[int]) -> tuple[dict[int, JobPlan], dict[int, bool]]: + """Resolve one dependency chain, in the order given.""" + assert _CHAIN_CONTEXT is not None, "chain context not initialised" + (config, points_dict, config_setup, config_class, sibling_index, sweep_filter_expr) = _CHAIN_CONTEXT - for point_idx in ordered_indices: + resolved_jobs: dict[int, JobPlan] = {} + filtered_jobs: dict[int, bool] = {} + sibling_context_cache: dict[tuple, tuple] = {} + + for point_idx in chain: point = points_dict[point_idx] try: group_filters = _collect_group_filters(config.sweep.groups, point.group_path) @@ -445,7 +445,7 @@ def resolve_sweep_with_dag( sibling_jobs[pattern] = resolved_jobs[sibling_point.index] sibling_ids.append((pattern, sibling_point.index)) - # Every point in a stage chain sees the same siblings, and flattening a + # Every point in a chain sees the same siblings, and flattening a # resolved sibling config is the most expensive step in this loop, so # build the context once per distinct set of siblings. context_key = tuple(sorted(sibling_ids)) @@ -487,18 +487,12 @@ def resolve_sweep_with_dag( # Detect if this is a config group parameter if is_config_group(key, config_setup.config_dir): # Config group: use no prefix (regular override) - param_overrides.extend( - param_to_cmdlines(key, value, prefix="", config_dir=config_setup.config_dir) - ) + param_overrides.extend(param_to_cmdlines(key, value, prefix="", config_dir=config_setup.config_dir)) else: # Regular parameter: use ++ prefix (force-add) - param_overrides.extend( - param_to_cmdlines(key, value, prefix="++", config_dir=config_setup.config_dir) - ) + param_overrides.extend(param_to_cmdlines(key, value, prefix="++", config_dir=config_setup.config_dir)) - compose_overrides = ( - list(config_setup.overrides) + [f"++index={point_idx}"] + param_overrides - ) + compose_overrides = list(config_setup.overrides) + [f"++index={point_idx}"] + param_overrides job_parameters = ( list(config_setup.overrides) + cmdline_overrides_siblings @@ -518,30 +512,131 @@ def resolve_sweep_with_dag( config_class=config_class, ) - resolved_dict = asdict(resolved) - context = {k: v for k, v in resolved_dict.items() if k not in ("sweep")} + cached_context: list = [] + + def filter_context() -> dict: + if not cached_context: + cached_context.append( + {k: v for k, v in asdict(resolved).items() if k not in ("sweep")} + ) + return cached_context[0] + skip_point = False for expr in filter_exprs: - if not _resolve_filter_from_context(expr, context): + if not _resolve_filter_from_context(expr, filter_context): LOGGER.info("Skipping point %s due to sweep.filter", point_idx) skip_point = True break filtered_jobs[point_idx] = skip_point + resolved_jobs[point_idx] = JobPlan( + config=resolved, + parameters=job_parameters, + sibling_pattern=None, + stage_name=getattr(resolved, "stage", None), + ) + + # Returned as bytes: the pool's own pickler cannot handle compoconf configs. + return _dumps((resolved_jobs, filtered_jobs)) + + +def _split_into_chains(dag: "nx.DiGraph", ordered_indices: list[int]) -> list[list[int]]: + """Group points into independent chains, each in topological order.""" + position = {index: order for order, index in enumerate(ordered_indices)} + chains = [ + sorted(component, key=position.__getitem__) + for component in nx.weakly_connected_components(dag) + ] + # Longest first, so the pool does not finish early workers and then wait on + # one long chain started last. + chains.sort(key=len, reverse=True) + return chains + + +def resolve_sweep_with_dag( + config: StagedSweepRoot, + points: list[SweepPoint] | dict[int, SweepPoint], + config_setup: ConfigSetup, + config_class: type = StagedSweepRoot, +) -> list[JobPlan]: + """Pure OmegaConf resolution with DAG ordering.""" + LOGGER.info(f"Starting DAG resolution for {len(points)} sweep points") + + if isinstance(points, list): + points_dict = {p.index: p for p in points} + else: + points_dict = points + + sibling_index = _build_sibling_index(points_dict) + dag = build_dependency_dag_from_points(points_dict, sibling_index=sibling_index) + + if not nx.is_directed_acyclic_graph(dag): + cycles = list(nx.simple_cycles(dag)) + LOGGER.error(f"Circular dependencies detected: {cycles}") + raise ValueError(f"Circular dependencies detected: {cycles}") + + ordered_indices = list(nx.topological_sort(dag)) + LOGGER.debug(f"Topological order: {ordered_indices}") + + resolved_jobs = {} + filtered_jobs = {} + base_context = asdict(config) + base_context = {k: v for k, v in base_context.items() if k not in ("sweep", "sibling")} + sweep_filter_expr = config.sweep.filter if isinstance(config.sweep, SweepConfig) else True + + if not isinstance(config.sweep, SweepConfig): + point = points_dict[list(points_dict)[0]] + resolved = load_config_reference( + config_dir=config_setup.config_dir, + config_path=config_setup.config_path, + config_name=config_setup.config_name, + overrides=point.parameters, + config_class=config_class, + ) + + resolved_dict = asdict(resolved) + context = {k: v for k, v in resolved_dict.items() if k not in ("sweep")} + skip_point = False + stage_name = getattr(resolved, "stage", None) job = JobPlan( config=resolved, - parameters=job_parameters, + parameters=point, sibling_pattern=None, stage_name=stage_name, ) - resolved_jobs[point_idx] = job + return [job] - return list( - resolved_jobs[point_idx] for point_idx in resolved_jobs if not filtered_jobs[point_idx] + chains = _split_into_chains(dag, ordered_indices) + workers = worker_count(len(chains), len(ordered_indices)) + LOGGER.info("Resolving %d chains with %d worker(s)", len(chains), workers) + + global _CHAIN_CONTEXT + _CHAIN_CONTEXT = ( + config, + points_dict, + config_setup, + config_class, + sibling_index, + sweep_filter_expr, ) + try: + chain_results = run_chunks(_resolve_chain, chains, workers) + finally: + _CHAIN_CONTEXT = None + + for chain_resolved, chain_filtered in (pickle.loads(blob) for blob in chain_results): + resolved_jobs.update(chain_resolved) + filtered_jobs.update(chain_filtered) + + # Emit in topological order, matching the order a single-process run built. + return [ + resolved_jobs[point_idx] + for point_idx in ordered_indices + if point_idx in resolved_jobs and not filtered_jobs[point_idx] + ] __all__ = [ diff --git a/oellm_autoexp/hydra_staged_sweep/parallel.py b/oellm_autoexp/hydra_staged_sweep/parallel.py new file mode 100644 index 00000000..6091e02e --- /dev/null +++ b/oellm_autoexp/hydra_staged_sweep/parallel.py @@ -0,0 +1,72 @@ +"""Fork a pool of workers over independent chunks of a sweep. + +Resolving a sweep point and rendering its job script are both pure CPU work on +data that is already in memory, so ``fork`` is the right tool: the children +inherit the loaded modules, the registered resolvers and the warm config +caches, and start doing useful work immediately. + +``HYDRA_STAGED_SWEEP_WORKERS`` controls the pool: unset or ``0`` picks one +worker per available CPU, ``1`` keeps everything in-process. +""" + +from __future__ import annotations + +import logging +import multiprocessing +import os +from collections.abc import Callable, Sequence +from typing import Any, TypeVar + +LOGGER = logging.getLogger(__name__) + +__all__ = ["run_chunks", "split_evenly", "worker_count"] + +T = TypeVar("T") + +WORKERS_ENV = "HYDRA_STAGED_SWEEP_WORKERS" + + +def worker_count(chunks: int, items: int, *, requested: int | None = None) -> int: + """How many processes to use; 1 means stay in-process.""" + if requested is None: + configured = os.environ.get(WORKERS_ENV) + requested = int(configured) if configured else 0 + if requested < 0: + raise ValueError(f"{WORKERS_ENV} must not be negative: {requested}") + if requested == 0: + requested = ( + len(os.sched_getaffinity(0)) + if hasattr(os, "sched_getaffinity") + else (os.cpu_count() or 1) + ) + if "fork" not in multiprocessing.get_all_start_methods(): + # Without fork a worker would have to re-import and re-register + # everything, which costs more than a sweep this size saves. + return 1 + # Below roughly two items per worker the process overhead dominates. + if chunks < 2 or items < 8: + return 1 + return max(1, min(requested, chunks)) + + +def run_chunks(func: Callable[[Any], T], chunks: Sequence[Any], workers: int) -> list[T]: + """Apply ``func`` to each chunk, in a fork pool when ``workers > 1``. + + ``func`` and the data it closes over are inherited through the fork, so + only each chunk and its result cross the process boundary. + """ + if workers <= 1: + return [func(chunk) for chunk in chunks] + context = multiprocessing.get_context("fork") + with context.Pool(processes=workers) as pool: + return pool.map(func, chunks, chunksize=1) + + +def split_evenly(items: Sequence[T], groups: int) -> list[list[T]]: + """Deal ``items`` round-robin into ``groups`` lists, preserving order.""" + if groups <= 1: + return [list(items)] + buckets: list[list[T]] = [[] for _ in range(groups)] + for position, item in enumerate(items): + buckets[position % groups].append(item) + return [bucket for bucket in buckets if bucket] diff --git a/oellm_autoexp/orchestrator.py b/oellm_autoexp/orchestrator.py index 49de97db..7b09bb34 100644 --- a/oellm_autoexp/orchestrator.py +++ b/oellm_autoexp/orchestrator.py @@ -14,6 +14,7 @@ from compoconf import asdict from oellm_autoexp.hydra_staged_sweep import expand_sweep, resolve_sweep_with_dag +from oellm_autoexp.hydra_staged_sweep.parallel import run_chunks, split_evenly, worker_count from oellm_autoexp.hydra_staged_sweep.expander import SweepPoint from oellm_autoexp.hydra_staged_sweep.planner import JobPlan @@ -341,6 +342,25 @@ def _resolve_job_name(config: RootConfig, total_jobs: int = 1) -> str: return f"{base_name}_{index_str}" +_RENDER_CONTEXT: tuple | None = None + + +def _render_chunk(indices: list[int]) -> list[tuple[int, str]]: + """Render the scripts for a slice of the plan, in a worker or in-process.""" + assert _RENDER_CONTEXT is not None, "render context not initialised" + plan, session_id = _RENDER_CONTEXT + rendered: list[tuple[int, str]] = [] + for index in indices: + record = _build_job_record(plan, plan.jobs[index], session_id) + if not isinstance(record.definition, SlurmJobConfig): + LOGGER.info("Skipping script render for non-SLURM job '%s'", record.job_id) + continue + path = generate_script(record.definition.slurm) + LOGGER.info("Rendered script: %s", path) + rendered.append((index, str(path))) + return rendered + + def render_job_scripts(plan: ExecutionPlan, *, session_id: str = "dry-run") -> list[Path]: """Render sbatch scripts for every job in the plan without submitting. @@ -350,16 +370,21 @@ def render_job_scripts(plan: ExecutionPlan, *, session_id: str = "dry-run") -> l Returns: Ordered list of paths to the written ``.sbatch`` files. """ - script_paths: list[Path] = [] - for job in plan.jobs: - record = _build_job_record(plan, job, session_id) - if not isinstance(record.definition, SlurmJobConfig): - LOGGER.info("Skipping script render for non-SLURM job '%s'", record.job_id) - continue - path = generate_script(record.definition.slurm) - LOGGER.info("Rendered script: %s", path) - script_paths.append(Path(path)) - return script_paths + global _RENDER_CONTEXT + + # Each job renders to its own file, so this fans out cleanly. Deal the jobs + # round-robin rather than in blocks so the workers stay balanced. + indices = list(range(len(plan.jobs))) + workers = worker_count(len(indices), len(indices)) + chunks = split_evenly(indices, workers) if workers > 1 else [indices] + + _RENDER_CONTEXT = (plan, session_id) + try: + results = [pair for chunk in run_chunks(_render_chunk, chunks, workers) for pair in chunk] + finally: + _RENDER_CONTEXT = None + # Restore plan order, which round-robin scattering broke. + return [Path(path) for _, path in sorted(results)] def chain_submit_jobs( diff --git a/tests/hydra_staged_sweep/test_parallel_resolution.py b/tests/hydra_staged_sweep/test_parallel_resolution.py new file mode 100644 index 00000000..e9ff4938 --- /dev/null +++ b/tests/hydra_staged_sweep/test_parallel_resolution.py @@ -0,0 +1,140 @@ +"""Resolving a sweep across a process pool must match resolving it in-process.""" + +import os +import textwrap +from dataclasses import dataclass, field +from typing import Any + +import pytest + +from oellm_autoexp.hydra_staged_sweep.config.loader import load_config_reference +from oellm_autoexp.hydra_staged_sweep.config.schema import ConfigSetup, StagedSweepRoot +from oellm_autoexp.hydra_staged_sweep.dag_resolver import _dumps, resolve_sweep_with_dag +from oellm_autoexp.hydra_staged_sweep.expander import expand_sweep +from oellm_autoexp.hydra_staged_sweep.parallel import split_evenly, worker_count + + +@dataclass(kw_only=True) +class Project(StagedSweepRoot): + name: str = "" + out_dir: str = "" + + +@dataclass(kw_only=True) +class PoolTestConfig(StagedSweepRoot): + lr: float = 0.0 + width: int = 0 + steps: int = 0 + stage: str = "" + load_path: str = "" + out_dir: str = "" + metadata: dict[str, Any] = field(default_factory=dict) + + +SWEEP = textwrap.dedent("""\ + lr: 0.001 + width: 128 + steps: 100 + stage: stable + load_path: "" + out_dir: "/tmp/${stage}_w${width}_lr${lr}" + + sweep: + type: product + groups: + - type: product + params: + width: [128, 256, 512, 1024] + lr: [0.001, 0.002] + - type: list + configs: + - stage: stable + steps: 100 + - stage: decay + steps: 20 + load_path: "\\\\${sibling.stable.out_dir}/ckpt" + """) + + +@pytest.fixture +def sweep_dir(tmp_path): + conf = tmp_path / "conf" + conf.mkdir() + (conf / "config.yaml").write_text(SWEEP) + return conf + + +def _resolve(sweep_dir, workers): + setup = ConfigSetup( + config_name="config", config_path=None, config_dir=str(sweep_dir), overrides=[] + ) + root = load_config_reference( + config_name="config", config_dir=str(sweep_dir), config_class=PoolTestConfig + ) + points = expand_sweep(root.sweep) + previous = os.environ.get("HYDRA_STAGED_SWEEP_WORKERS") + os.environ["HYDRA_STAGED_SWEEP_WORKERS"] = str(workers) + try: + return resolve_sweep_with_dag(root, points, setup, config_class=PoolTestConfig) + finally: + if previous is None: + os.environ.pop("HYDRA_STAGED_SWEEP_WORKERS", None) + else: + os.environ["HYDRA_STAGED_SWEEP_WORKERS"] = previous + + +def _fingerprint(jobs): + return [(j.stage_name, j.config.out_dir, j.config.load_path, tuple(j.parameters)) for j in jobs] + + +def test_pooled_resolution_matches_in_process(sweep_dir): + sequential = _resolve(sweep_dir, 1) + pooled = _resolve(sweep_dir, 4) + assert len(sequential) == 16 + assert _fingerprint(pooled) == _fingerprint(sequential) + + +def test_sibling_references_survive_the_pool(sweep_dir): + """A decay stage must still see the stable stage it branched off.""" + jobs = _resolve(sweep_dir, 4) + decay = [j for j in jobs if j.stage_name == "decay"] + assert decay, "expected decay stages in the sweep" + for job in decay: + assert job.config.load_path.endswith("/ckpt") + assert "stable" in job.config.load_path + + +def test_plans_survive_the_pickle_round_trip(sweep_dir): + """compoconf configs cannot pickle themselves; _dumps has to cover for it.""" + import pickle + + jobs = _resolve(sweep_dir, 1) + restored = pickle.loads(_dumps(jobs)) + assert _fingerprint(restored) == _fingerprint(jobs) + assert type(restored[0].config) is type(jobs[0].config) + + +@pytest.mark.parametrize( + ("chunks", "items", "expected"), + [ + (1, 100, 1), # nothing to spread + (10, 4, 1), # too little work to be worth forking + (3, 100, 3), # never more workers than chunks + ], +) +def test_worker_count_declines_pointless_pools(chunks, items, expected): + assert worker_count(chunks, items, requested=8) == expected + + +def test_worker_count_honours_explicit_request(): + assert worker_count(50, 500, requested=3) == 3 + assert worker_count(50, 500, requested=1) == 1 + with pytest.raises(ValueError): + worker_count(50, 500, requested=-1) + + +def test_split_evenly_keeps_every_item_once(): + items = list(range(10)) + buckets = split_evenly(items, 3) + assert sorted(x for bucket in buckets for x in bucket) == items + assert all(buckets) From 2c5a309cf8589826144321863b3df197cc05df29 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Korbinian=20P=C3=B6ppel?= Date: Sun, 2 Aug 2026 23:32:26 +0200 Subject: [PATCH 3/4] Drop the compoconf pickling shim compoconf 0.2.1 fixes ConfigInterface pickling: configs now use the default dataclass reduction, which restores state onto a new instance without calling __init__ and keeps nested configs typed. That is exactly what _PlanPickler did by hand, so plans can cross the process boundary on the pool's own pickler and _restore_config / _PlanPickler / _dumps all go away. The floor is compoconf>=0.2.2, which is the first release carrying both that fix and the tests covering it. 0.2.1 is the oldest release that works; 0.2.0 and earlier cannot unpickle a worker's results at all, because cls() with no arguments raises for configs with required fields, and nested configs come back as plain dicts. While here, the filter context becomes a small class instead of a closure over loop variables, and the dead context/skip_point assignments in the no-sweep branch go -- they were unreachable leftovers, and moving the resolve loop into _resolve_chain made that visible. --- .../hydra_staged_sweep/dag_resolver.py | 64 ++++++------------- pyproject.toml | 2 +- .../test_parallel_resolution.py | 9 ++- 3 files changed, 27 insertions(+), 48 deletions(-) diff --git a/oellm_autoexp/hydra_staged_sweep/dag_resolver.py b/oellm_autoexp/hydra_staged_sweep/dag_resolver.py index 67600788..544afb92 100644 --- a/oellm_autoexp/hydra_staged_sweep/dag_resolver.py +++ b/oellm_autoexp/hydra_staged_sweep/dag_resolver.py @@ -11,10 +11,7 @@ from __future__ import annotations -import io import logging -import os -import pickle import re from collections import defaultdict from collections.abc import Callable, Mapping, Sequence @@ -24,7 +21,7 @@ from typing import Any import networkx as nx -from compoconf import ConfigInterface, asdict +from compoconf import asdict from omegaconf import DictConfig, ListConfig, OmegaConf from .config.schema import StagedSweepRoot, SweepConfig, ConfigSetup @@ -380,33 +377,25 @@ def param_to_cmdlines(key: str, val: Any, prefix: str = "", config_dir: str | Pa ) -def _restore_config(cls: type, state: dict) -> Any: - """Rebuild a config from its attribute state, without calling __init__.""" - obj = cls.__new__(cls) - obj.__dict__.update(state) - return obj +class _LazyFilterContext: + """Build the filter context only when a filter actually reads it. - -class _PlanPickler(pickle.Pickler): - """Pickle resolved plans by state. - - compoconf's ``ConfigInterface.__reduce__`` reduces to ``(cls, (), state)``, - so unpickling calls ``cls()`` with no arguments -- which raises for any - config with required fields -- and its state comes from ``asdict``, which - flattens nested configs into plain dicts. Neither survives the trip back - from a worker process, so reduce by ``__dict__`` instead. + Flattening a resolved config is not cheap, and a filter that is already a + bool never looks at the context. """ - def reducer_override(self, obj: Any) -> Any: - if isinstance(obj, ConfigInterface): - return _restore_config, (type(obj), obj.__dict__) - return NotImplemented - + def __init__(self, resolved: Any) -> None: + self._resolved = resolved + self._context: dict[str, Any] | None = None -def _dumps(obj: Any) -> bytes: - buffer = io.BytesIO() - _PlanPickler(buffer, protocol=pickle.HIGHEST_PROTOCOL).dump(obj) - return buffer.getvalue() + def __call__(self) -> dict[str, Any]: + if self._context is None: + self._context = { + key: value + for key, value in asdict(self._resolved).items() + if key not in ("sweep") + } + return self._context # Chains of dependent points (a stable stage plus the cooldowns that branch off @@ -512,15 +501,7 @@ def _resolve_chain(chain: list[int]) -> tuple[dict[int, JobPlan], dict[int, bool config_class=config_class, ) - cached_context: list = [] - - def filter_context() -> dict: - if not cached_context: - cached_context.append( - {k: v for k, v in asdict(resolved).items() if k not in ("sweep")} - ) - return cached_context[0] - + filter_context = _LazyFilterContext(resolved) skip_point = False for expr in filter_exprs: if not _resolve_filter_from_context(expr, filter_context): @@ -536,11 +517,10 @@ def filter_context() -> dict: stage_name=getattr(resolved, "stage", None), ) - # Returned as bytes: the pool's own pickler cannot handle compoconf configs. - return _dumps((resolved_jobs, filtered_jobs)) + return resolved_jobs, filtered_jobs -def _split_into_chains(dag: "nx.DiGraph", ordered_indices: list[int]) -> list[list[int]]: +def _split_into_chains(dag: nx.DiGraph, ordered_indices: list[int]) -> list[list[int]]: """Group points into independent chains, each in topological order.""" position = {index: order for order, index in enumerate(ordered_indices)} chains = [ @@ -594,10 +574,6 @@ def resolve_sweep_with_dag( config_class=config_class, ) - resolved_dict = asdict(resolved) - context = {k: v for k, v in resolved_dict.items() if k not in ("sweep")} - skip_point = False - stage_name = getattr(resolved, "stage", None) job = JobPlan( @@ -627,7 +603,7 @@ def resolve_sweep_with_dag( finally: _CHAIN_CONTEXT = None - for chain_resolved, chain_filtered in (pickle.loads(blob) for blob in chain_results): + for chain_resolved, chain_filtered in chain_results: resolved_jobs.update(chain_resolved) filtered_jobs.update(chain_filtered) diff --git a/pyproject.toml b/pyproject.toml index c47c77a4..a71ce6a4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -13,7 +13,7 @@ authors = [ license = { text = "Apache-2.0" } requires-python = ">=3.10" dependencies = [ - "compoconf>=0.1.13", + "compoconf>=0.2.2", "hydra-core>=1.3", "omegaconf>=2.3", "pydantic>=2.0", diff --git a/tests/hydra_staged_sweep/test_parallel_resolution.py b/tests/hydra_staged_sweep/test_parallel_resolution.py index e9ff4938..69570fe3 100644 --- a/tests/hydra_staged_sweep/test_parallel_resolution.py +++ b/tests/hydra_staged_sweep/test_parallel_resolution.py @@ -9,7 +9,7 @@ from oellm_autoexp.hydra_staged_sweep.config.loader import load_config_reference from oellm_autoexp.hydra_staged_sweep.config.schema import ConfigSetup, StagedSweepRoot -from oellm_autoexp.hydra_staged_sweep.dag_resolver import _dumps, resolve_sweep_with_dag +from oellm_autoexp.hydra_staged_sweep.dag_resolver import resolve_sweep_with_dag from oellm_autoexp.hydra_staged_sweep.expander import expand_sweep from oellm_autoexp.hydra_staged_sweep.parallel import split_evenly, worker_count @@ -105,13 +105,16 @@ def test_sibling_references_survive_the_pool(sweep_dir): def test_plans_survive_the_pickle_round_trip(sweep_dir): - """compoconf configs cannot pickle themselves; _dumps has to cover for it.""" + """Plans cross the process boundary by pickle, so they have to survive it.""" import pickle jobs = _resolve(sweep_dir, 1) - restored = pickle.loads(_dumps(jobs)) + restored = pickle.loads(pickle.dumps(jobs, protocol=pickle.HIGHEST_PROTOCOL)) assert _fingerprint(restored) == _fingerprint(jobs) + # compoconf < 0.2.1 rebuilt configs through __init__ and flattened nested + # configs to dicts, which lost both of these. assert type(restored[0].config) is type(jobs[0].config) + assert restored[0].config == jobs[0].config @pytest.mark.parametrize( From 5125f76d7a48ae67028b4c2a0b9a8b273d51acf9 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Korbinian=20P=C3=B6ppel?= Date: Mon, 3 Aug 2026 00:34:39 +0200 Subject: [PATCH 4/4] Fix: Formatting. --- .../hydra_staged_sweep/config/loader.py | 4 ++- .../hydra_staged_sweep/dag_resolver.py | 27 ++++++++++++------- oellm_autoexp/orchestrator.py | 3 ++- tests/hydra_staged_sweep/test_config_cache.py | 16 +++++------ .../test_parallel_resolution.py | 6 +++-- 5 files changed, 33 insertions(+), 23 deletions(-) diff --git a/oellm_autoexp/hydra_staged_sweep/config/loader.py b/oellm_autoexp/hydra_staged_sweep/config/loader.py index 65d43e53..c19df60b 100644 --- a/oellm_autoexp/hydra_staged_sweep/config/loader.py +++ b/oellm_autoexp/hydra_staged_sweep/config/loader.py @@ -78,7 +78,9 @@ def _merge_extra_config(cfg: Any, extra_config: Mapping[str, Any] | None) -> Non return # Accepts an already-built container so callers that reuse the same context # across configs can build it once; merging does not modify the source. - source = extra_config if OmegaConf.is_config(extra_config) else OmegaConf.create(dict(extra_config)) + source = ( + extra_config if OmegaConf.is_config(extra_config) else OmegaConf.create(dict(extra_config)) + ) with open_dict(cfg): cfg.merge_with(source) diff --git a/oellm_autoexp/hydra_staged_sweep/dag_resolver.py b/oellm_autoexp/hydra_staged_sweep/dag_resolver.py index 544afb92..913c2378 100644 --- a/oellm_autoexp/hydra_staged_sweep/dag_resolver.py +++ b/oellm_autoexp/hydra_staged_sweep/dag_resolver.py @@ -293,7 +293,8 @@ def dict_to_cmdlines(dct: dict | list | str | int | float, prefix: str = ""): def drop_cmdline_invisible(value: Any) -> Any: - """Strip what ``config_to_cmdline`` cannot express, so a direct merge matches it. + """Strip what ``config_to_cmdline`` cannot express, so a direct merge + matches it. An empty mapping flattens to zero overrides, so round-tripping a config through the command line silently drops it -- and a list element that is an @@ -380,8 +381,8 @@ def param_to_cmdlines(key: str, val: Any, prefix: str = "", config_dir: str | Pa class _LazyFilterContext: """Build the filter context only when a filter actually reads it. - Flattening a resolved config is not cheap, and a filter that is already a - bool never looks at the context. + Flattening a resolved config is not cheap, and a filter that is + already a bool never looks at the context. """ def __init__(self, resolved: Any) -> None: @@ -391,9 +392,7 @@ def __init__(self, resolved: Any) -> None: def __call__(self) -> dict[str, Any]: if self._context is None: self._context = { - key: value - for key, value in asdict(self._resolved).items() - if key not in ("sweep") + key: value for key, value in asdict(self._resolved).items() if key not in ("sweep") } return self._context @@ -408,7 +407,9 @@ def __call__(self) -> dict[str, Any]: def _resolve_chain(chain: list[int]) -> tuple[dict[int, JobPlan], dict[int, bool]]: """Resolve one dependency chain, in the order given.""" assert _CHAIN_CONTEXT is not None, "chain context not initialised" - (config, points_dict, config_setup, config_class, sibling_index, sweep_filter_expr) = _CHAIN_CONTEXT + (config, points_dict, config_setup, config_class, sibling_index, sweep_filter_expr) = ( + _CHAIN_CONTEXT + ) resolved_jobs: dict[int, JobPlan] = {} filtered_jobs: dict[int, bool] = {} @@ -476,12 +477,18 @@ def _resolve_chain(chain: list[int]) -> tuple[dict[int, JobPlan], dict[int, bool # Detect if this is a config group parameter if is_config_group(key, config_setup.config_dir): # Config group: use no prefix (regular override) - param_overrides.extend(param_to_cmdlines(key, value, prefix="", config_dir=config_setup.config_dir)) + param_overrides.extend( + param_to_cmdlines(key, value, prefix="", config_dir=config_setup.config_dir) + ) else: # Regular parameter: use ++ prefix (force-add) - param_overrides.extend(param_to_cmdlines(key, value, prefix="++", config_dir=config_setup.config_dir)) + param_overrides.extend( + param_to_cmdlines(key, value, prefix="++", config_dir=config_setup.config_dir) + ) - compose_overrides = list(config_setup.overrides) + [f"++index={point_idx}"] + param_overrides + compose_overrides = ( + list(config_setup.overrides) + [f"++index={point_idx}"] + param_overrides + ) job_parameters = ( list(config_setup.overrides) + cmdline_overrides_siblings diff --git a/oellm_autoexp/orchestrator.py b/oellm_autoexp/orchestrator.py index 7b09bb34..74b2504b 100644 --- a/oellm_autoexp/orchestrator.py +++ b/oellm_autoexp/orchestrator.py @@ -346,7 +346,8 @@ def _resolve_job_name(config: RootConfig, total_jobs: int = 1) -> str: def _render_chunk(indices: list[int]) -> list[tuple[int, str]]: - """Render the scripts for a slice of the plan, in a worker or in-process.""" + """Render the scripts for a slice of the plan, in a worker or in- + process.""" assert _RENDER_CONTEXT is not None, "render context not initialised" plan, session_id = _RENDER_CONTEXT rendered: list[tuple[int, str]] = [] diff --git a/tests/hydra_staged_sweep/test_config_cache.py b/tests/hydra_staged_sweep/test_config_cache.py index fe60fada..b3d8e028 100644 --- a/tests/hydra_staged_sweep/test_config_cache.py +++ b/tests/hydra_staged_sweep/test_config_cache.py @@ -57,9 +57,7 @@ def fresh_cache(): def _load(config_dir, overrides=None): - return load_hydra_config( - "config", config_dir, overrides or [], config_class=CacheTestConfig - ) + return load_hydra_config("config", config_dir, overrides or [], config_class=CacheTestConfig) def test_cached_composition_matches_uncached(config_dir): @@ -117,11 +115,9 @@ def test_disable_restores_hydra(config_dir): def test_fast_yaml_types_match_the_pure_python_loader(tmp_path): - """libyaml must type scalars exactly as OmegaConf's own loader does.""" + """Libyaml must type scalars exactly as OmegaConf's own loader does.""" sample = tmp_path / "s.yaml" - sample.write_text( - "a: yes\nb: no\nc: null\nd: 2020-01-01\ne: 1.0e-4\nf: '010'\ng: 010\nh: ~\n" - ) + sample.write_text("a: yes\nb: no\nc: null\nd: 2020-01-01\ne: 1.0e-4\nf: '010'\ng: 010\nh: ~\n") cache.disable() plain = OmegaConf.to_container(OmegaConf.load(sample)) cache.enable() @@ -133,7 +129,8 @@ def test_fast_yaml_types_match_the_pure_python_loader(tmp_path): def test_repeated_interpolations_resolve_identically(config_dir, fresh_cache): - """The parse-tree cache hands out one shared tree; resolution must be stable.""" + """The parse-tree cache hands out one shared tree; resolution must be + stable.""" first = _load(config_dir, ["++depth=7"]).label second = _load(config_dir, ["++depth=8"]).label assert (first, second) == ("base-7", "base-8") @@ -154,7 +151,8 @@ def test_repeated_interpolations_resolve_identically(config_dir, fresh_cache): ], ) def test_extra_config_matches_the_override_round_trip(config_dir, fresh_cache, payload): - """Merging the context directly must land exactly where the ++overrides do.""" + """Merging the context directly must land exactly where the ++overrides + do.""" value = {"nested": payload} via_overrides = _load(config_dir, config_to_cmdline(value, override="++")) via_merge = load_hydra_config( diff --git a/tests/hydra_staged_sweep/test_parallel_resolution.py b/tests/hydra_staged_sweep/test_parallel_resolution.py index 69570fe3..35bcde3f 100644 --- a/tests/hydra_staged_sweep/test_parallel_resolution.py +++ b/tests/hydra_staged_sweep/test_parallel_resolution.py @@ -1,4 +1,5 @@ -"""Resolving a sweep across a process pool must match resolving it in-process.""" +"""Resolving a sweep across a process pool must match resolving it in- +process.""" import os import textwrap @@ -105,7 +106,8 @@ def test_sibling_references_survive_the_pool(sweep_dir): def test_plans_survive_the_pickle_round_trip(sweep_dir): - """Plans cross the process boundary by pickle, so they have to survive it.""" + """Plans cross the process boundary by pickle, so they have to survive + it.""" import pickle jobs = _resolve(sweep_dir, 1)