diff --git a/examples/advanced_example_config.yaml b/examples/advanced_example_config.yaml index e6cfb8f..103bedf 100644 --- a/examples/advanced_example_config.yaml +++ b/examples/advanced_example_config.yaml @@ -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 diff --git a/examples/advanced_example_experiment.py b/examples/advanced_example_experiment.py index 999c536..655ddb0 100644 --- a/examples/advanced_example_experiment.py +++ b/examples/advanced_example_experiment.py @@ -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 @@ -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. @@ -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): diff --git a/examples/artifacts/something b/examples/artifacts/something new file mode 100644 index 0000000..31d4233 --- /dev/null +++ b/examples/artifacts/something @@ -0,0 +1 @@ + \ No newline at end of file diff --git a/src/seml/document.py b/src/seml/document.py index 5af3793..13ce894 100644 --- a/src/seml/document.py +++ b/src/seml/document.py @@ -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): diff --git a/src/seml/experiment/sources.py b/src/seml/experiment/sources.py index 9c78cc5..5e34f5b 100644 --- a/src/seml/experiment/sources.py +++ b/src/seml/experiment/sources.py @@ -14,6 +14,7 @@ from seml.utils import ( assert_package_installed, is_local_file, + recursively_list_files, s_if, working_directory, ) @@ -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. @@ -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. @@ -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 @@ -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()) diff --git a/src/seml/settings.py b/src/seml/settings.py index e586a13..fdb2d31 100644 --- a/src/seml/settings.py +++ b/src/seml/settings.py @@ -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': [ diff --git a/src/seml/utils/__init__.py b/src/seml/utils/__init__.py index 691a77b..8c369a4 100644 --- a/src/seml/utils/__init__.py +++ b/src/seml/utils/__init__.py @@ -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.')