diff --git a/exca/map.py b/exca/map.py index 53be2efb..0374b0ad 100644 --- a/exca/map.py +++ b/exca/map.py @@ -372,7 +372,8 @@ def _method_override(self, *args: tp.Any, **kwargs: tp.Any) -> tp.Iterator[tp.An items = next(iter(kwargs.values())) # specific function for thread and process pool executors if self.cluster in [None, "threadpool", "processpool"]: - return self._method_override_futures(items) + with self._work_env(): + return self._method_override_futures(items) uid_func = imethod.item_uid # we need to keep order for output: uid_items = [(uid_func(item), item) for item in items] diff --git a/exca/slurm.py b/exca/slurm.py index b6720bc8..f1a961b0 100644 --- a/exca/slurm.py +++ b/exca/slurm.py @@ -127,7 +127,7 @@ def model_post_init(self, log__: tp.Any) -> None: f"cluster={self.cluster} requires a folder to be provided, " "only cluster=None works without folder" ) - if self.workdir is not None: + if self.workdir is not None and self.workdir.folder is None: raise ValueError("Workdir requires a folder") if self.tasks_per_node > 1 and not self.slurm_use_srun: if self.cluster in ["slurm", "auto"]: @@ -205,7 +205,8 @@ def _work_env(self) -> tp.Iterator[None]: if not isinstance(self, base.BaseInfra): raise RuntimeError("SubmititMixin should be set a BaseInfra mixin") with contextlib.ExitStack() as estack: - estack.enter_context(submitit.helpers.clean_env()) + if self.executor is not None: + estack.enter_context(submitit.helpers.clean_env()) if self.workdir is not None: if self.workdir.folder is None: if self.folder is None: diff --git a/exca/task.py b/exca/task.py index ba0e1cdf..da609847 100644 --- a/exca/task.py +++ b/exca/task.py @@ -337,7 +337,8 @@ def job(self) -> submitit.Job[tp.Any] | LocalJob: # submit job if it does not exist executor = self.executor() if executor is None: - job = LocalJob(self._run_method) + with self._work_env(): + job = LocalJob(self._run_method) job._name = self._factory() # for better logging message else: executor.folder.mkdir(exist_ok=True, parents=True) diff --git a/exca/test_base.py b/exca/test_base.py index 1e4d664b..ae12de8c 100644 --- a/exca/test_base.py +++ b/exca/test_base.py @@ -388,6 +388,18 @@ def test_tricky_update(tmp_path: Path) -> None: assert isinstance(wxp.infra.workdir, WorkDir) +def test_workdir_no_cache(tmp_path: Path) -> None: + # pb in confdict for subconfig + infra: tp.Any = { + "folder": None, + "workdir": {"folder": tmp_path, "copied": [Path(__file__).parent]}, + } + xp = Base(infra=infra) + assert xp.func() == 24 + folders = [x.name for x in tmp_path.iterdir()] + assert folders == ["exca"] + + def test_missing_base_model() -> None: with pytest.raises(RuntimeError):