flexlock 0.8.2__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
flexlock/api.py ADDED
@@ -0,0 +1,1444 @@
1
+ """Python API for FlexLock."""
2
+
3
+ from enum import Enum
4
+ from pathlib import Path
5
+ from omegaconf import OmegaConf, DictConfig, open_dict
6
+ from loguru import logger
7
+ from typing import List, Dict, Any, Optional
8
+ import yaml
9
+ import json
10
+ from .utils import (
11
+ instantiate,
12
+ load_python_defaults,
13
+ extract_tracking_info,
14
+ select_and_freeze_root_refs,
15
+ )
16
+ from .freeze import freeze_deferred
17
+ from .snapshot import snapshot, RunTracker
18
+ from .run_record import RunRecord
19
+ from .fingerprint import fingerprint as compute_fingerprint
20
+ from . import index
21
+ from .diff import RunDiff
22
+ from . import config as flexlock_config
23
+ from .exceptions import FlexLockExecutionError
24
+
25
+
26
+ def _leaf_paths(container, prefix=""):
27
+ """Yield ``(dotpath, value)`` for every leaf in a plain dict/list container.
28
+
29
+ Dotpaths use OmegaConf's ``a.b.0`` selector form so they can be fed back to
30
+ ``OmegaConf.select`` — used by :meth:`Project.check` to probe each leaf.
31
+ """
32
+ if isinstance(container, dict):
33
+ for k, v in container.items():
34
+ yield from _leaf_paths(v, f"{prefix}.{k}" if prefix else str(k))
35
+ elif isinstance(container, list):
36
+ for i, v in enumerate(container):
37
+ yield from _leaf_paths(v, f"{prefix}.{i}" if prefix else str(i))
38
+ else:
39
+ yield prefix, container
40
+
41
+
42
+ def _print_compiled_config(cfg):
43
+ """Print the resolved config + the target function's docstring, matching
44
+ what the docs promise for ``--print-config`` / ``print_config=True``.
45
+ """
46
+ print("=== COMPILED CONFIG ===")
47
+ print(OmegaConf.to_yaml(cfg))
48
+ target = cfg.get("_target_") if isinstance(cfg, DictConfig) else None
49
+ if not target:
50
+ return
51
+ print("=== TARGET FUNCTION DOCSTRING ===")
52
+ print(f"Target: {target}")
53
+ try:
54
+ import importlib
55
+
56
+ module_name, func_name = target.rsplit(".", 1)
57
+ module = importlib.import_module(module_name)
58
+ func = getattr(module, func_name)
59
+ doc = getattr(func, "__doc__", None)
60
+ print(f"Docstring:\n{doc}" if doc else "No docstring available.")
61
+ except (ImportError, AttributeError, ValueError) as e:
62
+ print(f"Could not import target function '{target}': {e}")
63
+
64
+
65
+ class Status(str, Enum):
66
+ """Terminal state of a run. A ``str`` subclass so ``status == "SUCCESS"``
67
+ and f-string formatting keep working for existing callers."""
68
+
69
+ SUCCESS = "SUCCESS"
70
+ CACHED = "CACHED"
71
+ FAILED = "FAILED"
72
+ INTERRUPTED = "INTERRUPTED"
73
+ SUBMITTED = "SUBMITTED"
74
+ SKIPPED = "SKIPPED"
75
+
76
+ def __str__(self) -> str:
77
+ return self.value
78
+
79
+
80
+ def _coerce_status(s):
81
+ if isinstance(s, Status):
82
+ return s
83
+ try:
84
+ return Status(s)
85
+ except ValueError:
86
+ return s # tolerate unknown strings rather than raising
87
+
88
+
89
+ # Real attributes that a result-dict key must never shadow (issue 6).
90
+ _RESERVED_RESULT_KEYS = frozenset(
91
+ {"save_dir", "status", "result", "metrics", "cfg", "error", "get",
92
+ "raise_on_failure", "is_success"}
93
+ )
94
+
95
+
96
+ class ExecutionResult:
97
+ """Typed result object from a run.
98
+
99
+ The function's return payload is kept in ``result`` (alias ``metrics``);
100
+ its keys are still reachable as attributes (``res.accuracy``) via
101
+ ``__getattr__`` — but only when they don't collide with a real attribute,
102
+ so a payload key named ``status``/``get`` no longer clobbers the object's
103
+ own API (issue 6).
104
+ """
105
+
106
+ def __init__(
107
+ self,
108
+ save_dir: str,
109
+ status: "str | Status",
110
+ result: Any = None,
111
+ cfg: DictConfig = None,
112
+ error: "str | None" = None,
113
+ ):
114
+ """
115
+ Args:
116
+ save_dir: Directory where results are saved
117
+ status: One of the :class:`Status` values
118
+ result: The actual return value from the function
119
+ cfg: Configuration used for execution
120
+ error: Error/traceback string for FAILED items (else None)
121
+ """
122
+ self.save_dir = save_dir
123
+ self.status = _coerce_status(status)
124
+ self.result = result
125
+ self.cfg = cfg
126
+ self.error = error
127
+
128
+ @property
129
+ def metrics(self):
130
+ """The return payload (alias for ``result``)."""
131
+ return self.result
132
+
133
+ @property
134
+ def is_success(self) -> bool:
135
+ return self.status in (Status.SUCCESS, Status.CACHED)
136
+
137
+ def __getitem__(self, key):
138
+ """Allow dict-like access to result."""
139
+ if isinstance(self.result, dict):
140
+ return self.result[key]
141
+ raise TypeError(f"Result is not a dict: {type(self.result)}")
142
+
143
+ def get(self, key, default=None):
144
+ """Dict-like get method."""
145
+ if isinstance(self.result, dict):
146
+ return self.result.get(key, default)
147
+ return default
148
+
149
+ def __getattr__(self, name):
150
+ """Expose result-dict keys as attributes without clobbering real ones.
151
+
152
+ Only invoked when normal attribute lookup fails, so real attributes
153
+ (status, save_dir, ...) always win over same-named payload keys.
154
+ """
155
+ if name.startswith("__") or name in _RESERVED_RESULT_KEYS:
156
+ raise AttributeError(name)
157
+ result = self.__dict__.get("result")
158
+ if isinstance(result, dict) and name in result:
159
+ return result[name]
160
+ raise AttributeError(name)
161
+
162
+ def raise_on_failure(self) -> "ExecutionResult":
163
+ """Raise if the run did not succeed; return self otherwise (chainable)."""
164
+ if self.status in (Status.FAILED, Status.INTERRUPTED):
165
+ raise FlexLockExecutionError(
166
+ f"Run at {self.save_dir} ended {self.status}: {self.error}"
167
+ )
168
+ return self
169
+
170
+ def __repr__(self):
171
+ return f"ExecutionResult(save_dir={self.save_dir}, status={self.status})"
172
+
173
+
174
+ class ChainedResult:
175
+ """Return type for :meth:`Project.submit_chained`.
176
+
177
+ Attributes:
178
+ sweep: List of sweep ``ExecutionResult`` objects.
179
+ downstream: List of lists; ``downstream[i][j]`` is the result of
180
+ the ``j``-th downstream stage for the ``i``-th sweep item.
181
+ """
182
+
183
+ def __init__(
184
+ self,
185
+ sweep: List[ExecutionResult],
186
+ downstream: List[List[ExecutionResult]],
187
+ ):
188
+ self.sweep = sweep
189
+ self.downstream = downstream
190
+
191
+ def __iter__(self):
192
+ """Iterate (sweep_result, downstream_results) pairs."""
193
+ return iter(zip(self.sweep, self.downstream))
194
+
195
+ def __len__(self):
196
+ return len(self.sweep)
197
+
198
+ def __repr__(self):
199
+ n_down = len(self.downstream[0]) if self.downstream else 0
200
+ return (
201
+ f"ChainedResult(sweep={len(self.sweep)}, "
202
+ f"downstream_per_item={n_down})"
203
+ )
204
+
205
+
206
+ class Project:
207
+ def __init__(self, defaults: "str | DictConfig | dict | None" = None):
208
+ """Initialize a FlexLock project.
209
+
210
+ Args:
211
+ defaults: One of:
212
+ - Python import path string (``'pkg.config.defaults'`` or
213
+ ``'path/to/file.py:defaults'``)
214
+ - A pre-built ``DictConfig`` or plain dict
215
+ - ``None`` for a project with no defaults — useful for
216
+ one-off submissions of an explicit config.
217
+ """
218
+ if defaults is None:
219
+ self.defaults_str = None
220
+ self.defaults = OmegaConf.create({})
221
+ elif isinstance(defaults, str):
222
+ self.defaults_str = defaults
223
+ loaded = load_python_defaults(defaults)
224
+ self.defaults = (
225
+ loaded if isinstance(loaded, DictConfig) else OmegaConf.create(loaded)
226
+ )
227
+ else:
228
+ self.defaults_str = None
229
+ self.defaults = (
230
+ defaults
231
+ if isinstance(defaults, DictConfig)
232
+ else OmegaConf.create(defaults)
233
+ )
234
+
235
+ def get(self, key: str):
236
+ """
237
+ Get a configuration by key from the defaults.
238
+
239
+ The returned config is self-contained: root-scope ``${...}`` references
240
+ are frozen into concrete values, while resolver calls (``${vinc:}``,
241
+ ``${latest:}``, ``${run_lock:}``, etc.) and intra-sub-tree refs are
242
+ preserved for resolution at submit time. This means the returned
243
+ config can be pickled, modified, and submitted without losing root
244
+ context (essential for HPC submission).
245
+
246
+ Args:
247
+ key: Dot-path to select a specific node from the defaults.
248
+
249
+ Returns:
250
+ The selected configuration (as DictConfig).
251
+ """
252
+ return select_and_freeze_root_refs(self.defaults, key)
253
+
254
+ def _generate_fingerprint(self, cfg: DictConfig) -> dict:
255
+ """
256
+ Generate a fingerprint (proposed snapshot) for the given config.
257
+
258
+ This is used for smart run logic to check if a run already exists.
259
+ """
260
+ repos, data, prevs = extract_tracking_info(cfg)
261
+
262
+ # Use RunTracker to generate snapshot without writing to disk
263
+ # We pass a dummy save_dir since we won't actually save
264
+ tracker = RunTracker(save_dir=Path("outputs/dummy-ref-smart-run"))
265
+
266
+ # Record environment and data
267
+ if repos:
268
+ tracker.record_env(repos)
269
+ if data:
270
+ tracker.record_data(data)
271
+
272
+ # Finalize to get the snapshot dict
273
+ fingerprint = tracker.finalize(cfg)
274
+
275
+ return fingerprint
276
+
277
+ def _find_matching_run(
278
+ self,
279
+ cfg: DictConfig,
280
+ search_dirs: List[str] = None,
281
+ match_include: List[str] = None,
282
+ match_exclude: List[str] = None,
283
+ ) -> Optional[Path]:
284
+ """
285
+ Search for an existing run that matches the given configuration.
286
+
287
+ Args:
288
+ cfg: Configuration to match
289
+ search_dirs: List of directories to search (defaults to parent of cfg.save_dir)
290
+ match_include: Override include patterns for git comparison
291
+ match_exclude: Override exclude patterns for git comparison
292
+
293
+ Returns:
294
+ Path to matching run directory, or None if no match found
295
+ """
296
+ # Auto-populate match_include from _target_ modules if not provided
297
+ if match_include is None:
298
+ from .utils import collect_target_include_patterns
299
+
300
+ match_include = collect_target_include_patterns(cfg) or None
301
+
302
+ # Compute the pure fingerprint digest — the index key. Both sweep tasks
303
+ # and serial runs record under this same key, so a config first run as a
304
+ # sweep task hits here when re-run serially (sweep items are first-class).
305
+ repos, data, _ = extract_tracking_info(cfg)
306
+ try:
307
+ fp = compute_fingerprint(cfg, repos=repos, data=data)
308
+ except Exception as e:
309
+ logger.warning(f"Could not compute fingerprint for lookup: {e}")
310
+ fp = None
311
+
312
+ # Determine where to search
313
+ if search_dirs is None:
314
+ if flexlock_config.WARN_SMART_RUN_NO_SEARCH_DIRS:
315
+ logger.warning(
316
+ "smart_run=True but search_dirs=None. "
317
+ "Defaulting to parent of save_dir. "
318
+ "This may not find all cached runs. "
319
+ "Set FLEXLOCK_WARN_SMART_RUN=0 to disable this warning."
320
+ )
321
+ if "save_dir" in cfg:
322
+ search_dirs = [str(Path(cfg.save_dir).parent)]
323
+ else:
324
+ logger.warning("No save_dir in config and no search_dirs provided")
325
+ return None
326
+
327
+ fallback = flexlock_config.get_env_bool("FLEXLOCK_INDEX_FALLBACK", True)
328
+
329
+ for search_root in search_dirs:
330
+ logger.debug(f"Searching for matching runs in: {search_root}")
331
+ root_path = Path(search_root)
332
+
333
+ # 1. Fast path: single indexed lookup by fingerprint.
334
+ if fp:
335
+ row = index.lookup(root_path, fp)
336
+ if row is not None:
337
+ resolved = index.verify_and_resolve(root_path, row)
338
+ if resolved is not None:
339
+ logger.success(f"⚡ Cache Hit (index)! {resolved}")
340
+ return resolved
341
+
342
+ # 2. Fallback: legacy glob scan for runs predating the index. On a
343
+ # hit, backfill the index so the slow path self-eliminates.
344
+ if fallback and root_path.exists():
345
+ match = self._glob_scan_match(
346
+ cfg, root_path, match_include, match_exclude
347
+ )
348
+ if match is not None:
349
+ if fp:
350
+ index.record_run_lock(match, fp)
351
+ logger.success(f"⚡ Cache Hit (glob)! {match}")
352
+ return match
353
+
354
+ return None
355
+
356
+ def _glob_scan_match(
357
+ self, cfg, root_path, match_include, match_exclude
358
+ ) -> Optional[Path]:
359
+ """Legacy O(N) content scan used only as an index fallback (RunDiff)."""
360
+ fingerprint = self._generate_fingerprint(cfg)
361
+ for lock_file in Path(root_path).glob("**/run.lock"):
362
+ run_dir = Path(lock_file).parent
363
+ try:
364
+ with open(lock_file, "r") as f:
365
+ candidate_snapshot = yaml.safe_load(f)
366
+
367
+ proposed_save_dir = fingerprint.get("config", {}).get("save_dir")
368
+ candidate_save_dir = candidate_snapshot.get("config", {}).get(
369
+ "save_dir"
370
+ )
371
+
372
+ differ = RunDiff(
373
+ current=fingerprint,
374
+ target=candidate_snapshot,
375
+ current_save_dir=proposed_save_dir,
376
+ target_save_dir=candidate_save_dir,
377
+ ignore_keys=["_snapshot_"],
378
+ match_include=match_include,
379
+ match_exclude=match_exclude,
380
+ )
381
+
382
+ if differ.is_match():
383
+ # Require run.complete — interrupted runs are not cache hits.
384
+ if not (run_dir / "run.complete").exists():
385
+ logger.debug(
386
+ f"Match at {run_dir} has no run.complete; skipping"
387
+ )
388
+ continue
389
+ return run_dir
390
+ else:
391
+ logger.debug(f"No match for run at: {run_dir}: {differ.diffs}")
392
+ except Exception as e:
393
+ logger.debug(f"Failed to read/compare {lock_file}: {e}")
394
+ continue
395
+ return None
396
+
397
+ def exists(self, cfg: DictConfig, search_dirs: List[str] = None) -> bool:
398
+ """
399
+ Check if a run with the given configuration already exists.
400
+
401
+ Args:
402
+ cfg: Configuration to check
403
+ search_dirs: Optional list of directories to search
404
+
405
+ Returns:
406
+ True if matching run exists, False otherwise
407
+ """
408
+ return self._find_matching_run(cfg, search_dirs) is not None
409
+
410
+ def get_result(
411
+ self, cfg: DictConfig, search_dirs: List[str] = None
412
+ ) -> ExecutionResult:
413
+ """
414
+ Retrieve results from a previously completed run.
415
+
416
+ Args:
417
+ cfg: Configuration to match
418
+ search_dirs: Optional list of directories to search
419
+
420
+ Returns:
421
+ ExecutionResult object with cached results
422
+
423
+ Raises:
424
+ ValueError: If no matching run is found
425
+ """
426
+ match_dir = self._find_matching_run(cfg, search_dirs)
427
+
428
+ if match_dir is None:
429
+ raise ValueError("No matching run found. Use exists() to check first.")
430
+
431
+ return self._load_cached_result(match_dir, cfg)
432
+
433
+ def _load_cached_result(self, match_dir: Path, cfg: DictConfig) -> ExecutionResult:
434
+ """Build a CACHED ExecutionResult from an existing run directory."""
435
+ result_data = None
436
+
437
+ # Try results.json
438
+ results_file = match_dir / "results.json"
439
+ if results_file.exists():
440
+ with open(results_file, "r") as f:
441
+ result_data = json.load(f)
442
+
443
+ # Try loading from run.lock
444
+ lock_file = match_dir / "run.lock"
445
+ if result_data is None and lock_file.exists():
446
+ with open(lock_file, "r") as f:
447
+ lock_data = yaml.safe_load(f)
448
+ result_data = lock_data.get("result", {})
449
+
450
+ return ExecutionResult(
451
+ save_dir=str(match_dir), status="CACHED", result=result_data, cfg=cfg
452
+ )
453
+
454
+ def run_stage(
455
+ self, cfg, stage_name=None, smart_run=True, search_dirs=None, **submit_kwargs
456
+ ):
457
+ """
458
+ Run a single stage with automatic search_dirs and save_dir propagation.
459
+
460
+ Args:
461
+ cfg: Stage configuration (DictConfig)
462
+ stage_name: Name of the stage (inferred from save_dir if None)
463
+ smart_run: If True, checks for cached runs
464
+ search_dirs: Directories to search for cached runs (auto-discovered if None)
465
+ **submit_kwargs: Additional arguments passed to submit()
466
+
467
+ Returns:
468
+ ExecutionResult (or list for sweeps)
469
+ """
470
+ # Infer stage name from save_dir
471
+ if stage_name is None and "save_dir" in cfg:
472
+ stage_name = Path(cfg.save_dir).name
473
+
474
+ # Auto-discover search_dirs: look for sibling experiment dirs with same stage name
475
+ if search_dirs is None and smart_run and stage_name and "save_dir" in cfg:
476
+ parent = Path(cfg.save_dir).parent.parent
477
+ if parent.exists():
478
+ search_dirs = [
479
+ str(p) for p in parent.glob(f"*/{stage_name}") if p.is_dir()
480
+ ]
481
+
482
+ result = self.submit(
483
+ cfg, smart_run=smart_run, search_dirs=search_dirs, **submit_kwargs
484
+ )
485
+
486
+ # For single runs, propagate save_dir back into cfg
487
+ if isinstance(result, list):
488
+ return result
489
+ with open_dict(cfg):
490
+ cfg.save_dir = result.save_dir
491
+ return result
492
+
493
+ def save_snapshot(self, save_dir):
494
+ """
495
+ Save the current project defaults as pipeline.yaml.
496
+
497
+ Args:
498
+ save_dir: Directory to save the pipeline snapshot to
499
+ """
500
+ save_path = Path(save_dir)
501
+ save_path.mkdir(parents=True, exist_ok=True)
502
+ (save_path / "pipeline.yaml").write_text(OmegaConf.to_yaml(self.defaults))
503
+
504
+ def submit(
505
+ self,
506
+ config: "DictConfig | str | None" = None,
507
+ sweep: List[Dict] = None,
508
+ sweep_target: str = None,
509
+ sweep_root: "str | None" = None,
510
+ n_jobs: int = 1,
511
+ smart_run: bool = True,
512
+ search_dirs: List[str] = None,
513
+ wait: bool = True,
514
+ pbs_config: str = None,
515
+ slurm_config: str = None,
516
+ sweep_dir_suffix: bool = False,
517
+ match_include: List[str] = None,
518
+ match_exclude: List[str] = None,
519
+ isolated: bool = False,
520
+ force: bool = False,
521
+ overrides: "dict | List[str] | None" = None,
522
+ merge: "str | Path | dict | None" = None,
523
+ debug: bool = False,
524
+ print_config: bool = False,
525
+ dry_run: bool = False,
526
+ tag: "str | None" = None,
527
+ timeout: "int | None" = None,
528
+ note: "str | None" = None,
529
+ save_dir_policy: "str | None" = None,
530
+ ) -> "ExecutionResult | List[ExecutionResult] | None":
531
+ """Submit a configuration for execution.
532
+
533
+ Args:
534
+ config: The configuration to execute. Accepts a ``DictConfig``,
535
+ a key string (looked up via :meth:`get`), or ``None`` (uses
536
+ ``self.defaults`` as-is).
537
+ sweep: Optional list of override dicts for parameter sweep.
538
+ sweep_target: Dot-path inside each task config where the sweep
539
+ item is merged. When ``None``, items merge at the root of
540
+ ``config``.
541
+ n_jobs: Number of parallel workers (for sweeps).
542
+ smart_run: If True, checks for an existing cached run before
543
+ executing.
544
+ search_dirs: Directories to search for cached runs.
545
+ wait: If True, blocks until completion.
546
+ pbs_config / slurm_config: Path to HPC backend YAML.
547
+ sweep_dir_suffix: If True, append ``_sweep_{i:04d}`` to each
548
+ sweep item's ``save_dir``.
549
+ match_include / match_exclude: Override git path patterns used
550
+ during ``smart_run`` comparison.
551
+ isolated: If True, run in a spawned subprocess even for a
552
+ single task (use for GPU stages to confine the CUDA
553
+ context).
554
+ force: If True, invalidate the cache marker (``run.complete``)
555
+ for this ``save_dir`` and re-execute. Outputs and
556
+ ``run.lock`` are preserved; the user function overwrites
557
+ in place.
558
+ overrides: Dict (``{'lr': 0.01}``) or dotlist
559
+ (``['lr=0.01']``) merged into ``config`` before execution.
560
+ merge: Path to a YAML file (or a dict) merged into ``config``
561
+ before execution. ``overrides`` is applied after ``merge``.
562
+ debug: Wrap the user function with the post-mortem debugger so
563
+ exceptions drop into PDB.
564
+ print_config: Print the fully-resolved config and return
565
+ ``None`` without executing.
566
+ dry_run: When using an HPC backend, render the would-be Slurm
567
+ or PBS submission script, print it (with any validation
568
+ warnings), and return ``None`` without submitting. No-op
569
+ for local execution.
570
+ tag: Human-readable label for this sweep's rows in the shared task
571
+ DB (e.g. ``"extract"`` or ``"collocate"``). Passed through to
572
+ ``ParallelExecutor`` so workers and the status CLI can scope to
573
+ it. When ``None`` a deterministic hash is auto-generated.
574
+ note: Free-text intent recorded as a top-level ``note:`` key in
575
+ ``run.lock`` (sibling of ``timestamp``/``config``). Never enters
576
+ the fingerprint, so it can't perturb caching. For sweeps the note
577
+ lands on the master ``run.lock`` only; sweep items inherit it for
578
+ display via their ``.flexlock_marker`` → master lookup.
579
+ save_dir_policy: What to do when ``config.save_dir`` is already
580
+ occupied by a previous run (contains ``run.lock``), applied
581
+ once at submit time, after the ``smart_run`` cache check:
582
+ ``None``/``"raise"`` (default) refuses and raises;
583
+ ``"increment"`` versions the dir (``run`` → ``run_0000``,
584
+ claimed atomically); ``"overwrite"`` deletes the occupied
585
+ dir's contents first (never touches a dir without
586
+ ``run.lock``); ``"skip"`` returns the existing complete
587
+ run's result without executing; ``"unsafe"`` runs in place
588
+ with no check; ``"timestamp"`` appends
589
+ ``config.TIMESTAMP_FORMAT``. ``force=True`` bypasses the
590
+ default guard (explicit in-place rerun). Replaces the
591
+ ``${vinc:}`` / ``${now:}`` resolvers. For sweeps the policy
592
+ applies once to the sweep root; items nest beneath it, and
593
+ ``"skip"`` means resume (reuse complete items, run the rest).
594
+
595
+ Returns:
596
+ ``ExecutionResult`` (single), ``List[ExecutionResult]`` (sweep),
597
+ or ``None`` (``print_config=True``).
598
+ """
599
+ # Resolve config from a key, a DictConfig, or default to self.defaults.
600
+ if config is None:
601
+ config = self.defaults
602
+ elif isinstance(config, str):
603
+ config = self.get(config)
604
+ if not isinstance(config, DictConfig):
605
+ config = OmegaConf.create(config)
606
+
607
+ # Apply post-resolution merges and overrides (post-select equivalents).
608
+ if merge is not None:
609
+ if isinstance(merge, (str, Path)):
610
+ config.merge_with(OmegaConf.load(str(merge)))
611
+ else:
612
+ config.merge_with(OmegaConf.create(merge))
613
+ if overrides is not None:
614
+ if isinstance(overrides, dict):
615
+ overrides = [f"{k}={v}" for k, v in overrides.items()]
616
+ config.merge_with(OmegaConf.from_dotlist(overrides))
617
+
618
+ # Bake save_dir to a concrete string exactly once. Done here (not
619
+ # lazily during config reads) so run.lock and run.complete always land
620
+ # in the same directory. The policy itself (naming / collision guard)
621
+ # is applied later, after the smart_run cache check and the
622
+ # side-effect-free previews (print_config / dry_run).
623
+ from .save_dir import SKIP, apply_save_dir_policy, resolve_save_dir
624
+
625
+ resolve_save_dir(config)
626
+
627
+ if print_config:
628
+ # When a sweep is given, preview each item's merged config so
629
+ # the user can verify per-item interpolations resolve correctly
630
+ # before launching.
631
+ if sweep:
632
+ from .utils import merge_task_into_cfg
633
+
634
+ for i, override in enumerate(sweep):
635
+ # Mirror _submit_sweep exactly (merge → freeze) so the
636
+ # preview matches what will actually execute.
637
+ item_cfg = merge_task_into_cfg(config, override, sweep_target)
638
+ item_cfg = freeze_deferred(item_cfg)
639
+ print(f"# --- sweep item {i} ---")
640
+ _print_compiled_config(item_cfg)
641
+ else:
642
+ _print_compiled_config(config)
643
+ return None
644
+
645
+ if force:
646
+ # Invalidate the single-run cache marker. For sweeps the per-item
647
+ # markers and task DB are reset inside _submit_sweep (2.4).
648
+ if not sweep:
649
+ save_dir = Path(config.get("save_dir", "outputs/job"))
650
+ marker = save_dir / "run.complete"
651
+ if marker.exists():
652
+ logger.info(f"Force flag enabled: invalidating cache at {save_dir}")
653
+ marker.unlink()
654
+ smart_run = False
655
+
656
+ # Handle sweep execution
657
+ if sweep:
658
+ # A sweep runs each item in its own worker process; `isolated`
659
+ # (spawn a subprocess for a *single* run) can't be honoured here.
660
+ # Fail loudly rather than silently dropping it (issue 16).
661
+ if isolated:
662
+ from .exceptions import FlexLockConfigError
663
+
664
+ raise FlexLockConfigError(
665
+ "isolated=True is not supported with sweep=... — sweep items "
666
+ "already execute in separate worker processes."
667
+ )
668
+ # A naming policy fixes the sweep root before items are built;
669
+ # guard policies are enforced per item inside _submit_sweep, after
670
+ # the per-item cache check ('skip' there means resume: reuse
671
+ # complete items, rerun crashed ones).
672
+ from .save_dir import NAMING_POLICIES, validate_policy
673
+
674
+ validate_policy(save_dir_policy)
675
+ if save_dir_policy in NAMING_POLICIES:
676
+ apply_save_dir_policy(config, save_dir_policy, force=force)
677
+ # Items are fresh under the freshly named root; keep the
678
+ # default guard as a backstop.
679
+ item_policy = "raise"
680
+ else:
681
+ item_policy = save_dir_policy or "raise"
682
+ return self._submit_sweep(
683
+ config,
684
+ sweep,
685
+ n_jobs,
686
+ smart_run,
687
+ search_dirs,
688
+ pbs_config,
689
+ slurm_config,
690
+ wait,
691
+ sweep_dir_suffix,
692
+ match_include,
693
+ match_exclude,
694
+ sweep_target=sweep_target,
695
+ sweep_root=sweep_root,
696
+ debug=debug,
697
+ tag=tag,
698
+ force=force,
699
+ timeout=timeout,
700
+ note=note,
701
+ item_policy=item_policy,
702
+ )
703
+
704
+ # Single execution path
705
+ # Check for existing run if smart_run is enabled
706
+ if smart_run:
707
+ match_dir = self._find_matching_run(
708
+ config, search_dirs, match_include, match_exclude
709
+ )
710
+ if match_dir:
711
+ logger.info(f"Skipping execution, using cached result from {match_dir}")
712
+ return self.get_result(config, search_dirs)
713
+
714
+ # Check if using HPC backend
715
+ use_hpc = pbs_config is not None or slurm_config is not None
716
+
717
+ if dry_run:
718
+ if not use_hpc:
719
+ logger.info("dry_run is a no-op for local execution.")
720
+ return None
721
+ self._preview_hpc_script(config, slurm_config, pbs_config)
722
+ return None
723
+
724
+ # Naming / collision-guard policy — after the cache check (a hit never
725
+ # claims or cleans anything) and after the previews (no side effects).
726
+ if apply_save_dir_policy(config, save_dir_policy, force=force) == SKIP:
727
+ match_dir = Path(str(config.save_dir))
728
+ logger.info(f"save_dir_policy='skip': reusing result from {match_dir}")
729
+ return self._load_cached_result(match_dir, config)
730
+
731
+ if use_hpc:
732
+ # Execute via HPC backend
733
+ logger.info(f"Submitting to HPC backend...")
734
+
735
+ # Use ParallelExecutor with a single task
736
+ from .parallel import ParallelExecutor
737
+
738
+ save_dir = config.get("save_dir", "outputs/job")
739
+ # Resolve _snapshot_ while config still has its parent chain: the
740
+ # executor_cfg below is re-rooted, so relative refs (${...key})
741
+ # inside _snapshot_ would no longer reach the stage node.
742
+ if "_snapshot_" in config:
743
+ snapshot_resolved = OmegaConf.to_container(
744
+ config._snapshot_, resolve=True
745
+ )
746
+ else:
747
+ snapshot_resolved = {}
748
+ executor_cfg = OmegaConf.create(
749
+ {"save_dir": str(save_dir), "_snapshot_": snapshot_resolved}
750
+ )
751
+
752
+ executor = ParallelExecutor(
753
+ func=instantiate,
754
+ tasks=[config], # Single task as a list
755
+ task_target=None,
756
+ cfg=executor_cfg,
757
+ n_jobs=flexlock_config.DEFAULT_N_JOBS,
758
+ pbs_config=pbs_config,
759
+ slurm_config=slurm_config,
760
+ local_workers=None,
761
+ tag=tag,
762
+ note=note,
763
+ )
764
+
765
+ # Run with wait parameter (executor handles waiting)
766
+ success = executor.run(
767
+ wait=wait, timeout=flexlock_config.DEFAULT_TIMEOUT if wait else None
768
+ )
769
+
770
+ # Load result
771
+ result_data = None
772
+ if "save_dir" in config:
773
+ results_file = Path(config.save_dir) / "results.json"
774
+ if results_file.exists():
775
+ with open(results_file, "r") as f:
776
+ result_data = json.load(f)
777
+
778
+ return ExecutionResult(
779
+ save_dir=str(save_dir),
780
+ status="SUCCESS" if wait else "SUBMITTED",
781
+ result=result_data,
782
+ cfg=config,
783
+ )
784
+
785
+ else:
786
+ # Local execution
787
+ if isolated:
788
+ # Run in an isolated spawned subprocess so that any GPU/CUDA
789
+ # context initialised inside the task stays in the child and
790
+ # never leaks into the parent (which may later fork workers).
791
+ from .parallel import ParallelExecutor
792
+
793
+ save_dir = config.get("save_dir", "outputs/job")
794
+ if "_snapshot_" in config:
795
+ snapshot_resolved = OmegaConf.to_container(
796
+ config._snapshot_, resolve=True
797
+ )
798
+ else:
799
+ snapshot_resolved = {}
800
+ executor_cfg = OmegaConf.create(
801
+ {"save_dir": str(save_dir), "_snapshot_": snapshot_resolved}
802
+ )
803
+ executor = ParallelExecutor(
804
+ func=instantiate,
805
+ tasks=[config],
806
+ task_target=None,
807
+ cfg=executor_cfg,
808
+ n_jobs=1,
809
+ isolated=True,
810
+ tag=tag,
811
+ note=note,
812
+ )
813
+ executor.run(wait=True)
814
+ result_data = None
815
+ if "save_dir" in config:
816
+ results_file = Path(save_dir) / "results.json"
817
+ if results_file.exists():
818
+ with open(results_file, "r") as f:
819
+ result_data = json.load(f)
820
+ return ExecutionResult(
821
+ save_dir=str(save_dir),
822
+ status="SUCCESS",
823
+ result=result_data,
824
+ cfg=config,
825
+ )
826
+
827
+ # Extract tracking info
828
+ repos, data, prevs = extract_tracking_info(config)
829
+
830
+ # Single deferred-resolution point (local single-run path): fire
831
+ # the deferred resolvers (run_lock/latest) once and detach into a
832
+ # plain config. Re-wrapping prevents instantiate()'s internal
833
+ # config.copy() — a fresh instance with an empty resolver cache —
834
+ # from re-firing anything after snapshot() has run.
835
+ try:
836
+ from .resolvers import resolve_deferred
837
+
838
+ config = resolve_deferred(config)
839
+ except Exception as exc:
840
+ logger.warning(f"Could not fully resolve config before execution: {exc}")
841
+
842
+ # Compute the fingerprint once (index key), stored in run.lock so
843
+ # `flexlock reindex` can rebuild the index from disk.
844
+ run_fp = None
845
+ if "save_dir" in config:
846
+ try:
847
+ run_fp = compute_fingerprint(config, repos=repos, data=data)
848
+ except Exception as exc:
849
+ logger.warning(f"Could not compute run fingerprint: {exc}")
850
+
851
+ # Create snapshot before execution
852
+ if "save_dir" in config:
853
+ snapshot(
854
+ config, repos=repos, data=data, prevs=prevs,
855
+ fingerprint=run_fp, note=note,
856
+ )
857
+
858
+ # Execute the function
859
+ logger.info(f"Executing configuration...")
860
+ run_func = instantiate
861
+ if debug:
862
+ from .debug import debug_on_fail
863
+
864
+ run_func = debug_on_fail(run_func)
865
+ try:
866
+ result = run_func(config)
867
+ except Exception as exc:
868
+ # Record a failure sidecar next to run.lock so triage doesn't
869
+ # need to re-run. The user's exception still propagates unwrapped
870
+ # (docs §12); write_error never raises. KeyboardInterrupt is
871
+ # deliberately not caught — a bare run.lock is the interrupted
872
+ # signature.
873
+ if "save_dir" in config:
874
+ RunRecord(config.save_dir).write_error(exc)
875
+ raise
876
+
877
+ # Save results if save_dir is specified
878
+ save_dir = config.get("save_dir", ".")
879
+ if "save_dir" in config:
880
+ record = RunRecord(save_dir)
881
+ try:
882
+ record.write_results(result)
883
+ except Exception as e:
884
+ logger.warning(
885
+ f"Could not save results to {record.results_path}: {e}"
886
+ )
887
+ record.mark_complete(result=result)
888
+ # Record the completed run in the project-wide index (2.2).
889
+ if run_fp:
890
+ index.record_run_lock(save_dir, run_fp)
891
+
892
+ return ExecutionResult(
893
+ save_dir=str(save_dir), status="SUCCESS", result=result, cfg=config
894
+ )
895
+
896
+ def check(
897
+ self,
898
+ config=None,
899
+ *,
900
+ sweep: "List[Dict] | None" = None,
901
+ sweep_target: "str | None" = None,
902
+ overrides: "dict | List[str] | None" = None,
903
+ merge: "str | Path | dict | None" = None,
904
+ ) -> "List[dict]":
905
+ """Preflight: fully resolve the config (and every sweep item) without
906
+ touching the filesystem or executing anything.
907
+
908
+ Mirrors submit's merge → freeze pipeline, but runs under the freeze
909
+ stubs so no resolver fires (``vinc``/``now`` take no ``mkdir``,
910
+ ``run_lock``/``latest`` read nothing). Every unresolvable interpolation
911
+ is reported — not just the first — as a dict with ``item`` (sweep index
912
+ or ``None``), ``full_key``, and ``error``.
913
+
914
+ Returns an empty list when everything resolves.
915
+ """
916
+ from .utils import merge_task_into_cfg
917
+ from .resolvers import frozen_resolvers
918
+ from omegaconf.errors import OmegaConfBaseException
919
+
920
+ # Normalize config exactly like submit (but never mutate the caller's).
921
+ if config is None:
922
+ config = self.defaults
923
+ elif isinstance(config, str):
924
+ config = self.get(config)
925
+ if not isinstance(config, DictConfig):
926
+ config = OmegaConf.create(config)
927
+ config = config.copy()
928
+ if merge is not None:
929
+ if isinstance(merge, (str, Path)):
930
+ config.merge_with(OmegaConf.load(str(merge)))
931
+ else:
932
+ config.merge_with(OmegaConf.create(merge))
933
+ if overrides is not None:
934
+ if isinstance(overrides, dict):
935
+ overrides = [f"{k}={v}" for k, v in overrides.items()]
936
+ config.merge_with(OmegaConf.from_dotlist(overrides))
937
+
938
+ if sweep:
939
+ items = [
940
+ (i, merge_task_into_cfg(config, ov, sweep_target))
941
+ for i, ov in enumerate(sweep)
942
+ ]
943
+ else:
944
+ items = [(None, config)]
945
+
946
+ errors: List[dict] = []
947
+ with frozen_resolvers():
948
+ for idx, item in items:
949
+ detached = OmegaConf.create(
950
+ OmegaConf.to_container(item, resolve=False, throw_on_missing=False)
951
+ )
952
+ # Resolve leaf-by-leaf so every failure is reported, not only
953
+ # the first one OmegaConf.resolve would raise on.
954
+ for path, val in _leaf_paths(
955
+ OmegaConf.to_container(detached, resolve=False, throw_on_missing=False)
956
+ ):
957
+ if not (isinstance(val, str) and "${" in val):
958
+ continue
959
+ try:
960
+ OmegaConf.select(detached, path, throw_on_missing=True)
961
+ except OmegaConfBaseException as e:
962
+ errors.append(
963
+ {
964
+ "item": idx,
965
+ "full_key": getattr(e, "full_key", None) or path,
966
+ "error": str(e),
967
+ }
968
+ )
969
+ return errors
970
+
971
+ @staticmethod
972
+ def _preview_hpc_script(config, slurm_config, pbs_config):
973
+ """Render the would-be HPC submission script and print it.
974
+
975
+ Used by ``submit(..., dry_run=True)``. Loads the backend YAML the
976
+ same way :class:`ParallelExecutor` would, instantiates the backend
977
+ targeting a temporary folder, and prints the rendered script plus
978
+ any validation warnings.
979
+ """
980
+ import tempfile
981
+
982
+ save_dir = Path(config.get("save_dir", "outputs/job"))
983
+ with tempfile.TemporaryDirectory() as tmp:
984
+ folder = Path(tmp)
985
+ if slurm_config:
986
+ from .backends.slurm import SlurmBackend, validate_slurm_script
987
+
988
+ params = OmegaConf.to_container(
989
+ OmegaConf.load(slurm_config), resolve=True
990
+ )
991
+ # Match ParallelExecutor's folder convention so paths in
992
+ # the preview look like the real submission.
993
+ params_folder = save_dir / "slurm_logs"
994
+ backend = SlurmBackend(folder=folder, **params)
995
+ script = backend.render_script()
996
+ # Rewrite the temp folder path to the would-be real one so
997
+ # the preview is faithful.
998
+ script = script.replace(str(folder), str(params_folder))
999
+
1000
+ print("# === Slurm submission script (dry run) ===")
1001
+ print(script)
1002
+ print("# === end ===")
1003
+ warnings = validate_slurm_script(script)
1004
+ if warnings:
1005
+ print("\n# Warnings:")
1006
+ for w in warnings:
1007
+ print(f"# - {w}")
1008
+ elif pbs_config:
1009
+ from .backends.pbs import PBSBackend
1010
+
1011
+ params = OmegaConf.to_container(
1012
+ OmegaConf.load(pbs_config), resolve=True
1013
+ )
1014
+ backend = PBSBackend(folder=folder, **params)
1015
+ # PBS backend may or may not expose render_script — fall
1016
+ # back to printing the YAML if not.
1017
+ if hasattr(backend, "render_script"):
1018
+ print("# === PBS submission script (dry run) ===")
1019
+ print(backend.render_script())
1020
+ print("# === end ===")
1021
+ else:
1022
+ print("# === PBS config (dry run) ===")
1023
+ print(OmegaConf.to_yaml(OmegaConf.create(params)))
1024
+ print("# === end ===")
1025
+
1026
+ @staticmethod
1027
+ def _validate_sweep_save_dirs(base_config, merged_items, sweep_root=None):
1028
+ """Ensure every sweep item's save_dir nests under the sweep root.
1029
+
1030
+ The sweep root is taken from ``sweep_root`` when provided, otherwise
1031
+ from the parent of the first (already merged and frozen) item's
1032
+ save_dir — the same directory the tasks DB lives in. Validation only
1033
+ ever touches the merged items, never ``base_config.save_dir`` (which
1034
+ may still carry a per-item ``${.variable}`` that is unresolvable until
1035
+ an item supplies it).
1036
+ """
1037
+ from .exceptions import FlexLockValidationError
1038
+
1039
+ if not merged_items:
1040
+ return
1041
+
1042
+ # When the user explicitly provides a sweep_root they've opted out of
1043
+ # the containment constraint — trust them and skip validation.
1044
+ if sweep_root is not None:
1045
+ return
1046
+
1047
+ # Derive the sweep root from the merged items (concrete save_dirs).
1048
+ first_save = merged_items[0][1].get("save_dir")
1049
+ if first_save is None:
1050
+ return # nothing to validate against
1051
+ effective_root = Path(first_save).resolve().parent
1052
+
1053
+ offenders = []
1054
+ for i, sweep_cfg in merged_items:
1055
+ sd = sweep_cfg.get("save_dir")
1056
+ if sd is None:
1057
+ continue
1058
+ try:
1059
+ Path(sd).resolve().relative_to(effective_root)
1060
+ except ValueError:
1061
+ offenders.append((i, sd))
1062
+
1063
+ if offenders:
1064
+ lines = [
1065
+ f" item {i}: save_dir={sd!r}" for i, sd in offenders
1066
+ ]
1067
+ raise FlexLockValidationError(
1068
+ f"Sweep item save_dir(s) must nest under the sweep root "
1069
+ f"({effective_root}); the tasks DB and lineage markers live "
1070
+ f"there. Offending items:\n"
1071
+ + "\n".join(lines)
1072
+ + f"\n\nTo keep the current save_dirs, use --sweep-root "
1073
+ f"<common_parent> (e.g. --sweep-root {Path(offenders[0][1]).parent})."
1074
+ )
1075
+
1076
+ def submit_chained(
1077
+ self,
1078
+ config=None,
1079
+ *,
1080
+ sweep: "List[Dict] | None" = None,
1081
+ downstream: "List[tuple] | None" = None,
1082
+ sweep_kwargs: "dict | None" = None,
1083
+ downstream_kwargs: "dict | None" = None,
1084
+ ) -> "ChainedResult":
1085
+ """Run a sweep, then chain a sequence of downstream stages per result.
1086
+
1087
+ For each item in ``sweep``, this method submits the base config and
1088
+ then, for each ``(stage_key, anchor_wiring)`` in ``downstream``,
1089
+ propagates the sweep result's attributes into ``proj.defaults`` and
1090
+ submits the downstream stage.
1091
+
1092
+ ``anchor_wiring`` maps an anchor name in ``proj.defaults`` to an
1093
+ attribute of the parent ``ExecutionResult`` (typically
1094
+ ``'save_dir'``).
1095
+
1096
+ Args:
1097
+ config: Base config for the sweep (passed to :meth:`submit`).
1098
+ sweep: List of override dicts for the parameter sweep.
1099
+ downstream: List of ``(stage_key, anchor_wiring)`` tuples
1100
+ describing the stages to run per sweep result.
1101
+ sweep_kwargs: Extra kwargs forwarded to ``submit`` for the
1102
+ sweep itself (e.g. ``slurm_config``, ``smart_run``).
1103
+ downstream_kwargs: Extra kwargs forwarded to ``submit`` for
1104
+ every downstream stage. Defaults to ``smart_run=False`` to
1105
+ avoid stale-cache false hits between iterations.
1106
+
1107
+ Returns:
1108
+ ``ChainedResult`` with ``sweep`` (list) and ``downstream`` (a
1109
+ list of lists, one per sweep item, each inner list aligned to
1110
+ the order of ``downstream``).
1111
+ """
1112
+ if sweep is None or downstream is None:
1113
+ raise ValueError(
1114
+ "submit_chained requires both `sweep` and `downstream` lists. "
1115
+ "Use submit() if you don't need post-sweep chaining."
1116
+ )
1117
+
1118
+ sweep_kwargs = dict(sweep_kwargs or {})
1119
+ downstream_kwargs = {"smart_run": False, **(downstream_kwargs or {})}
1120
+
1121
+ sweep_results = self.submit(config, sweep=sweep, **sweep_kwargs)
1122
+ if not isinstance(sweep_results, list):
1123
+ sweep_results = [sweep_results]
1124
+
1125
+ downstream_results: list[list[ExecutionResult]] = []
1126
+ for parent in sweep_results:
1127
+ per_parent: list[ExecutionResult] = []
1128
+ for entry in downstream:
1129
+ stage_key, anchor_wiring = entry
1130
+ for anchor_name, attr in anchor_wiring.items():
1131
+ value = getattr(parent, attr, None)
1132
+ if value is None:
1133
+ logger.warning(
1134
+ f"submit_chained: parent result has no '{attr}' "
1135
+ f"attribute; anchor '{anchor_name}' not updated."
1136
+ )
1137
+ continue
1138
+ OmegaConf.update(self.defaults, anchor_name, value)
1139
+ result = self.submit(stage_key, **downstream_kwargs)
1140
+ per_parent.append(result)
1141
+ downstream_results.append(per_parent)
1142
+
1143
+ return ChainedResult(sweep=sweep_results, downstream=downstream_results)
1144
+
1145
+ @classmethod
1146
+ def submit_config(cls, config=None, **kwargs):
1147
+ """Shortcut for one-off submissions without holding a Project instance.
1148
+
1149
+ Equivalent to ``Project().submit(config, **kwargs)``. Use this when
1150
+ you already have a ``DictConfig`` (from ``py2cfg`` or otherwise) and
1151
+ don't need ``proj.get``/``proj.exists``/``proj.defaults`` plumbing.
1152
+
1153
+ See :func:`flexlock.submit` for the module-level alias.
1154
+ """
1155
+ return cls().submit(config, **kwargs)
1156
+
1157
+ @staticmethod
1158
+ def collect_results(indices, task_configs, db_path, tag) -> list:
1159
+ """Build per-item ExecutionResults from the task DB's terminal state.
1160
+
1161
+ Replaces the old blanket ``status="SUCCESS"`` — a task that raised is
1162
+ now reported as ``FAILED`` with its traceback, so a sweep result set
1163
+ faithfully reflects what happened (issue 2).
1164
+ """
1165
+ from .taskdb import get_all_tasks, _hash_task
1166
+
1167
+ by_id = {
1168
+ t["task_id"]: t
1169
+ for t in get_all_tasks(db_path, tags=[tag] if tag else None)
1170
+ }
1171
+ status_map = {
1172
+ "done": "SUCCESS",
1173
+ "failed": "FAILED",
1174
+ "interrupted": "INTERRUPTED",
1175
+ }
1176
+ out = []
1177
+ for idx, cfg in zip(indices, task_configs):
1178
+ save_dir = str(cfg.get("save_dir", "."))
1179
+ row = by_id.get(_hash_task(cfg))
1180
+ db_status = row["status"] if row else None
1181
+ status = status_map.get(db_status, "SUBMITTED")
1182
+
1183
+ result_data = None
1184
+ error = None
1185
+ if status == "SUCCESS":
1186
+ result_data = RunRecord(save_dir).load_results()
1187
+ if result_data is None and row:
1188
+ result_data = row.get("result") or None
1189
+ elif status == "FAILED":
1190
+ error = row.get("error") if row else None
1191
+
1192
+ out.append(
1193
+ (
1194
+ idx,
1195
+ ExecutionResult(
1196
+ save_dir=save_dir,
1197
+ status=status,
1198
+ result=result_data,
1199
+ cfg=cfg,
1200
+ error=error,
1201
+ ),
1202
+ )
1203
+ )
1204
+ return out
1205
+
1206
+ @staticmethod
1207
+ def _sweep_db_dir(merged_items, sweep_root) -> "Path | None":
1208
+ """Directory that hosts the sweep's task DB (mirrors the parallel path)."""
1209
+ if sweep_root is not None:
1210
+ return Path(sweep_root)
1211
+ for _, cfg in merged_items:
1212
+ if "save_dir" in cfg:
1213
+ return Path(cfg.save_dir).parent
1214
+ return None
1215
+
1216
+ def _reset_sweep_for_force(self, merged_items, sweep_root) -> None:
1217
+ """Invalidate per-item markers and the task DB so a forced sweep reruns."""
1218
+ # Per-item completion markers + results (so RunRecord/index see fresh runs).
1219
+ for _, cfg in merged_items:
1220
+ if "save_dir" in cfg:
1221
+ d = Path(cfg.save_dir)
1222
+ (d / "run.complete").unlink(missing_ok=True)
1223
+ (d / "results.json").unlink(missing_ok=True)
1224
+
1225
+ # Task DB (its pending_count==0 short-circuit would otherwise resume).
1226
+ db_dir = self._sweep_db_dir(merged_items, sweep_root)
1227
+ if db_dir is not None:
1228
+ for suffix in ("", "-wal", "-shm"):
1229
+ (db_dir / f"run.lock.tasks.db{suffix}").unlink(missing_ok=True)
1230
+ logger.info(f"Force flag enabled: reset sweep task DB under {db_dir}")
1231
+
1232
+ def _submit_sweep(
1233
+ self,
1234
+ base_config: DictConfig,
1235
+ sweep: List[Dict],
1236
+ n_jobs: int,
1237
+ smart_run: bool,
1238
+ search_dirs: List[str],
1239
+ pbs_config: str = None,
1240
+ slurm_config: str = None,
1241
+ wait: bool = True,
1242
+ dir_suffix: bool = False,
1243
+ match_include: List[str] = None,
1244
+ match_exclude: List[str] = None,
1245
+ sweep_target: str = None,
1246
+ sweep_root: "str | None" = None,
1247
+ debug: bool = False,
1248
+ tag: "str | None" = None,
1249
+ force: bool = False,
1250
+ timeout: "int | None" = None,
1251
+ note: "str | None" = None,
1252
+ item_policy: str = "raise",
1253
+ ) -> List[ExecutionResult]:
1254
+ """
1255
+ Execute a parameter sweep.
1256
+
1257
+ Args:
1258
+ base_config: Base configuration
1259
+ sweep: List of override dictionaries
1260
+ n_jobs: Number of parallel workers
1261
+ smart_run: Whether to check for cached runs
1262
+ search_dirs: Directories to search for cached runs
1263
+ pbs_config: Path to PBS configuration YAML
1264
+ slurm_config: Path to Slurm configuration YAML
1265
+ wait: Whether to wait for jobs to complete
1266
+ match_include: Override include patterns for git comparison
1267
+ match_exclude: Override exclude patterns for git comparison
1268
+
1269
+ Returns:
1270
+ List of ExecutionResult objects
1271
+ """
1272
+ from .parallel import ParallelExecutor
1273
+ from .utils import merge_task_into_cfg
1274
+
1275
+ results = []
1276
+ configs_to_run = []
1277
+ cached_results = []
1278
+
1279
+ # Build sweep configs first so we can validate save_dir containment
1280
+ # in one place, before any execution.
1281
+ merged_items = []
1282
+ for i, override in enumerate(sweep):
1283
+ # Merge the sweep item into the base FIRST, so item-injected keys
1284
+ # (e.g. `variable`) exist before any resolution — then freeze.
1285
+ sweep_cfg = merge_task_into_cfg(base_config, override, sweep_target)
1286
+ # Eager-resolve everything self-contained for DB serialization,
1287
+ # while preserving deferred resolvers (run_lock/latest) as call
1288
+ # strings so they fire once, on the worker, at stage start.
1289
+ sweep_cfg = freeze_deferred(sweep_cfg)
1290
+ if dir_suffix and "save_dir" in sweep_cfg:
1291
+ # Nest each item under the base save_dir (the sweep root) so
1292
+ # tasks DB and lineage markers stay inside the same tree.
1293
+ # Pre-fix this produced siblings (e.g. train_sweep_0000 next
1294
+ # to train/), which always tripped the containment check.
1295
+ base_save_dir = Path(sweep_cfg.save_dir)
1296
+ sweep_cfg.save_dir = str(base_save_dir / f"sweep_{i:04d}")
1297
+ merged_items.append((i, sweep_cfg))
1298
+
1299
+ # Validate per-item save_dir containment up front. The tasks DB lives
1300
+ # at <sweep_root>/run.lock.tasks.db and each task records its path
1301
+ # relative to its parent dir. If items sit outside the sweep tree the
1302
+ # worker either fails opaquely (pre-validation) or can't form a
1303
+ # relative path. Surface a clear error before we queue anything.
1304
+ self._validate_sweep_save_dirs(base_config, merged_items, sweep_root=sweep_root)
1305
+
1306
+ # Force: reset the per-item completion markers *and* the task DB so every
1307
+ # item re-executes. Unlinking only the base marker (the old behaviour)
1308
+ # missed both the per-item markers and the DB's pending_count==0
1309
+ # resume short-circuit, so a forced sweep silently re-used cached
1310
+ # tasks (issue 15).
1311
+ if force:
1312
+ self._reset_sweep_for_force(merged_items, sweep_root)
1313
+
1314
+ # Check each sweep config for cached results
1315
+ for i, sweep_cfg in merged_items:
1316
+ if smart_run:
1317
+ match_dir = self._find_matching_run(
1318
+ sweep_cfg, search_dirs, match_include, match_exclude
1319
+ )
1320
+ if match_dir:
1321
+ logger.info(f"Sweep {i}: Using cached result from {match_dir}")
1322
+ cached_results.append((i, self.get_result(sweep_cfg, search_dirs)))
1323
+ continue
1324
+
1325
+ configs_to_run.append((i, sweep_cfg))
1326
+
1327
+ # Collision guard, per item, after the cache check (a cache hit never
1328
+ # trips the guard). Applied here — before anything is queued — so a
1329
+ # refusal aborts the whole sweep up front. Items whose save_dir is
1330
+ # occupied by a previous run are handled per `item_policy`; items of
1331
+ # *this* sweep sharing one save_dir don't trip it (nothing is occupied
1332
+ # until execution starts).
1333
+ if item_policy != "unsafe":
1334
+ from .save_dir import clean_run_dir, is_complete, is_occupied, occupied_error
1335
+
1336
+ still_to_run = []
1337
+ for i, sweep_cfg in configs_to_run:
1338
+ item_dir = sweep_cfg.get("save_dir")
1339
+ if item_dir is None or not is_occupied(item_dir):
1340
+ still_to_run.append((i, sweep_cfg))
1341
+ elif item_policy == "overwrite":
1342
+ clean_run_dir(item_dir)
1343
+ still_to_run.append((i, sweep_cfg))
1344
+ elif item_policy == "skip":
1345
+ # Resume: reuse complete items, rerun crashed ones in place.
1346
+ if is_complete(item_dir):
1347
+ logger.info(f"Sweep {i}: skip — reusing {item_dir}")
1348
+ cached_results.append(
1349
+ (i, self._load_cached_result(Path(item_dir), sweep_cfg))
1350
+ )
1351
+ else:
1352
+ still_to_run.append((i, sweep_cfg))
1353
+ elif force:
1354
+ still_to_run.append((i, sweep_cfg))
1355
+ else:
1356
+ raise occupied_error(str(item_dir), sweep_item=i)
1357
+ configs_to_run = still_to_run
1358
+
1359
+ # Execute remaining configs
1360
+ if configs_to_run:
1361
+ # Decide whether to use HPC backend or local execution
1362
+ use_hpc = pbs_config is not None or slurm_config is not None
1363
+
1364
+ if use_hpc or (n_jobs > 1 and len(configs_to_run) > 1):
1365
+ # Parallel execution using ParallelExecutor (local or HPC)
1366
+ logger.info(
1367
+ f"Executing {len(configs_to_run)} sweep configs with ParallelExecutor"
1368
+ )
1369
+
1370
+ # Extract just configs for parallel execution
1371
+ task_configs = [cfg for _, cfg in configs_to_run]
1372
+ indices = [i for i, _ in configs_to_run]
1373
+
1374
+ # Prepare a common save_dir for the sweep master (tasks DB lives here).
1375
+ if sweep_root is not None:
1376
+ sweep_save_dir = Path(sweep_root)
1377
+ elif "save_dir" in task_configs[0]:
1378
+ sweep_save_dir = Path(task_configs[0].save_dir).parent
1379
+ else:
1380
+ sweep_save_dir = Path("outputs/sweep")
1381
+
1382
+ # Create a wrapper config that ParallelExecutor can work with
1383
+ # The task configs are what ParallelExecutor will execute
1384
+ # Resolve _snapshot_ while base_config still has its parent chain
1385
+ # so that OmegaConf interpolations (e.g. ${...key}) can resolve
1386
+ if "_snapshot_" in base_config:
1387
+ snapshot_resolved = OmegaConf.to_container(
1388
+ base_config._snapshot_, resolve=True
1389
+ )
1390
+ else:
1391
+ snapshot_resolved = {}
1392
+
1393
+ executor_cfg = OmegaConf.create(
1394
+ {
1395
+ "save_dir": str(sweep_save_dir),
1396
+ "_snapshot_": snapshot_resolved,
1397
+ }
1398
+ )
1399
+
1400
+ # Use ParallelExecutor with backend support
1401
+ executor = ParallelExecutor(
1402
+ func=instantiate, # The function to execute
1403
+ tasks=task_configs, # List of configs to execute
1404
+ task_target=None, # Each task is already a complete config
1405
+ cfg=executor_cfg, # Master config for tracking
1406
+ n_jobs=n_jobs,
1407
+ pbs_config=pbs_config,
1408
+ slurm_config=slurm_config,
1409
+ local_workers=n_jobs if not use_hpc else None,
1410
+ tag=tag,
1411
+ note=note,
1412
+ )
1413
+
1414
+ # Run the sweep (executor handles waiting based on wait parameter)
1415
+ success = executor.run(wait=wait, timeout=timeout)
1416
+
1417
+ # Collect real per-task statuses from the task DB (issue 2).
1418
+ results.extend(
1419
+ self.collect_results(
1420
+ indices, task_configs, executor.db_path, executor.tag
1421
+ )
1422
+ )
1423
+ else:
1424
+ # Sequential execution (no backend, n_jobs=1)
1425
+ for i, cfg in configs_to_run:
1426
+ logger.info(f"Executing sweep {i}/{len(sweep)}")
1427
+ # Guard already applied above — run items unguarded so
1428
+ # same-save_dir items of one sweep behave as before.
1429
+ result = self.submit(
1430
+ cfg,
1431
+ sweep=None,
1432
+ smart_run=False,
1433
+ wait=True,
1434
+ debug=debug,
1435
+ force=force,
1436
+ save_dir_policy="unsafe",
1437
+ )
1438
+ results.append((i, result))
1439
+
1440
+ # Combine cached and new results, sorted by index
1441
+ all_results = cached_results + results
1442
+ all_results.sort(key=lambda x: x[0])
1443
+
1444
+ return [result for _, result in all_results]