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..42326d38 --- /dev/null +++ b/oellm_autoexp/hydra_staged_sweep/config/cache.py @@ -0,0 +1,674 @@ +"""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. + +Seven 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. +``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. +``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, + "defaults_hit": 0, + "defaults_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] + + +# --------------------------------------------------------------------------- +# 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 +# --------------------------------------------------------------------------- +# 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] + + +# --------------------------------------------------------------------------- +# 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 +# --------------------------------------------------------------------------- +_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, + defaults_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() + _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: + _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, _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 + + 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_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 + + 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() + _defaults_list_cache.clear() diff --git a/oellm_autoexp/hydra_staged_sweep/config/loader.py b/oellm_autoexp/hydra_staged_sweep/config/loader.py index 60c5916c..c19df60b 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,31 @@ 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 +106,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 +121,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 +143,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 +154,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 +163,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..913c2378 100644 --- a/oellm_autoexp/hydra_staged_sweep/dag_resolver.py +++ b/oellm_autoexp/hydra_staged_sweep/dag_resolver.py @@ -14,7 +14,7 @@ import logging 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 @@ -27,6 +27,7 @@ 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 +102,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) @@ -283,6 +292,32 @@ 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. @@ -343,63 +378,44 @@ 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") - - 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 _LazyFilterContext: + """Build the filter context only when a filter actually reads it. - 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}") + Flattening a resolved config is not cheap, and a filter that is + already a bool never looks at the context. + """ - ordered_indices = list(nx.topological_sort(dag)) - LOGGER.debug(f"Topological order: {ordered_indices}") + def __init__(self, resolved: Any) -> None: + self._resolved = resolved + self._context: dict[str, Any] | None = None - 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 + 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 - 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 +# 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 - stage_name = getattr(resolved, "stage", None) - job = JobPlan( - config=resolved, - parameters=point, - sibling_pattern=None, - stage_name=stage_name, - ) +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 + ) - return [job] + resolved_jobs: dict[int, JobPlan] = {} + filtered_jobs: dict[int, bool] = {} + sibling_context_cache: dict[tuple, tuple] = {} - for point_idx in ordered_indices: + for point_idx in chain: point = points_dict[point_idx] try: group_filters = _collect_group_filters(config.sweep.groups, point.group_path) @@ -410,39 +426,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 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 +486,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,38 +496,130 @@ 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, ) - resolved_dict = asdict(resolved) - context = {k: v for k, v in resolved_dict.items() if k not in ("sweep")} + filter_context = _LazyFilterContext(resolved) 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), + ) + + return 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, + ) + 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 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..74b2504b 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,26 @@ 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 +371,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/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_config_cache.py b/tests/hydra_staged_sweep/test_config_cache.py new file mode 100644 index 00000000..b3d8e028 --- /dev/null +++ b/tests/hydra_staged_sweep/test_config_cache.py @@ -0,0 +1,165 @@ +"""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 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..35bcde3f --- /dev/null +++ b/tests/hydra_staged_sweep/test_parallel_resolution.py @@ -0,0 +1,145 @@ +"""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 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): + """Plans cross the process boundary by pickle, so they have to survive + it.""" + import pickle + + jobs = _resolve(sweep_dir, 1) + 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( + ("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)