Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions examples/advanced_example_config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,8 @@ seml:
description: "An advanced example configuration. We can also use variable interpolation here: ${config.model.model_type}"
reschedule_timeout: 300 # The time (in seconds) that are left on the job before SEML will try to reschedule unfinished experiments.
# Note that you have to implement a `reschedule_hook` to use this feature.
additional_artifacts:
- artifacts/something

slurm:
- experiments_per_job: 1
Expand Down
8 changes: 8 additions & 0 deletions examples/advanced_example_experiment.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,8 @@
are parsed by a specific method. This avoids having one large "main" function which takes all parameters as input.
"""

import logging

import numpy as np
from seml import Experiment

Expand Down Expand Up @@ -132,6 +134,11 @@ def init_preprocessing(self, mean: float, std: float):
def init_augmentation(self, flip: bool):
self.augmentation_parameters = (flip,)

def init_artifacts(self):
# Load token from artifact specified in `seml.additional_artifacts
with open("artifacts/something") as f:
logging.info(f"Loaded artifact {f.read().strip()}")

def init_all(self):
"""
Sequentially run the sub-initializers of the experiment.
Expand All @@ -141,6 +148,7 @@ def init_all(self):
self.init_optimizer()
self.init_preprocessing()
self.init_augmentation()
self.init_artifacts()

@ex.capture(prefix="training")
def train(self, patience, num_epochs):
Expand Down
1 change: 1 addition & 0 deletions examples/artifacts/something
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
<file content of the artifact>
1 change: 1 addition & 0 deletions src/seml/document.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@ class SemlDocBase(TypedDict, total=False):
name: str
stash_all_py_files: bool
reschedule_timeout: int | None
additional_artifacts: list[str]


class SemlFileConfig(SemlDocBase, total=False):
Expand Down
23 changes: 21 additions & 2 deletions src/seml/experiment/sources.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
from seml.utils import (
assert_package_installed,
is_local_file,
recursively_list_files,
s_if,
working_directory,
)
Expand Down Expand Up @@ -75,7 +76,12 @@ def import_exe(executable: str, conda_env: str | None, working_dir: str):


def get_imported_sources(
executable, root_dir, conda_env, working_dir, stash_all_py_files: bool
executable,
root_dir,
conda_env,
working_dir,
stash_all_py_files: bool,
additional_artifacts: list[str] | None = None,
) -> set[str]:
"""Get the sources imported by the given executable.

Expand All @@ -85,6 +91,7 @@ def get_imported_sources(
conda_env (_type_): The experiment's Anaconda environment.
working_dir (_type_): The working directory of the experiment.
stash_all_py_files (_type_): Whether to stash all .py files in the working directory.
additional_artifacts: list[str] | None: Additional artifacts to put into the source code files.

Returns:
List[str]: The sources imported by the given executable.
Expand Down Expand Up @@ -114,6 +121,18 @@ def get_imported_sources(
if is_local_file(file, root_path):
sources.add(str(file))

for artifact in set().union(
*(recursively_list_files(path) for path in additional_artifacts or [])
):
artifact = artifact.expanduser().resolve()
# Check that the artifact is in `working_dir`
if artifact.is_file() and is_local_file(str(artifact), root_path):
sources.add(str(artifact))
else:
logging.warning(
f'Additional artifact {artifact} is not a subpath of the root directory '
f'{root_path} and will be ignored.'
)
return sources


Expand All @@ -124,13 +143,13 @@ def upload_sources(

with working_directory(seml_config['working_dir']):
root_dir = str(Path(seml_config['working_dir']).expanduser().resolve())

sources = get_imported_sources(
seml_config['executable'],
root_dir=root_dir,
conda_env=seml_config['conda_environment'],
working_dir=seml_config['working_dir'],
stash_all_py_files=seml_config.get('stash_all_py_files', False),
additional_artifacts=seml_config.get('additional_artifacts', []),
)
executable_abs = str(Path(seml_config['executable']).expanduser().resolve())

Expand Down
1 change: 1 addition & 0 deletions src/seml/settings.py
Original file line number Diff line number Diff line change
Expand Up @@ -237,6 +237,7 @@ class Settings(SettingsDict):
'description',
'stash_all_py_files',
'reschedule_timeout',
'additional_artifacts',
],
'SEML_CONFIG_VALUE_VERSION': 'version',
'VALID_SLURM_CONFIG_VALUES': [
Expand Down
16 changes: 16 additions & 0 deletions src/seml/utils/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -788,3 +788,19 @@ def drop_typeddict_difference(obj: TD1, cls: type[TD1], cls2: type[TD2]) -> TD2:
if k in result:
del result[k]
return result # type: ignore


def recursively_list_files(path: Path | str) -> set[Path]:
"""Recursively lists all (resolved) files in the directory.

Also supports `glob` patterns.
"""
if any(c in str(path) for c in '*?[]'):
return {p.expanduser() for p in Path().rglob(str(path))}
path = Path(path)
if path.expanduser().resolve().is_file():
return {path.expanduser().resolve()}
elif path.expanduser().resolve().is_dir():
return {p.expanduser().resolve() for p in path.rglob('*') if p.is_file()}
else:
raise ValueError(f'Path {path} is neither a file nor a directory.')