diff --git a/audmodel/core/api.py b/audmodel/core/api.py index 607284b..5901833 100644 --- a/audmodel/core/api.py +++ b/audmodel/core/api.py @@ -278,6 +278,7 @@ def load( *, cache_root: str | None = None, timeout: float = 14400, # 4 h + num_workers: int = 1, verbose: bool = False, ) -> str: r"""Download a model by its unique ID. @@ -294,6 +295,7 @@ def load( If not set :meth:`audmodel.default_cache_root` is used timeout: maximum time in seconds before giving up acquiring a lock + num_workers: number of parallel jobs verbose: show debug messages Returns: @@ -313,7 +315,7 @@ def load( """ cache_root = audeer.safe_path(cache_root or default_cache_root()) short_id, version = split_uid(uid, cache_root) - return get_archive(short_id, version, cache_root, timeout, verbose) + return get_archive(short_id, version, cache_root, timeout, num_workers, verbose) def meta( diff --git a/audmodel/core/backend.py b/audmodel/core/backend.py index 0f04c08..e69d61d 100644 --- a/audmodel/core/backend.py +++ b/audmodel/core/backend.py @@ -60,6 +60,7 @@ def get_archive( version: str, cache_root: str, timeout: float, + num_workers: int, verbose: bool, ) -> str: r"""Return backend and local archive path. @@ -70,6 +71,7 @@ def get_archive( cache_root: path of cache root timeout: maximum time in seconds before giving up acquiring a lock + num_workers: number of parallel jobs verbose: if ``True`` show message or progress bar when downloading file @@ -106,6 +108,7 @@ def get_archive( src_path, dst_path, version, + num_workers=num_workers, verbose=verbose, ) @@ -114,6 +117,7 @@ def get_archive( dst_path, tmp_root, keep_archive=False, + num_workers=num_workers, verbose=verbose, ) diff --git a/pyproject.toml b/pyproject.toml index a1f4f08..2ec6aa9 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -31,7 +31,8 @@ classifiers = [ ] requires-python = '>=3.10' dependencies = [ - 'audbackend[all] >=2.2.2', + 'audbackend[all] >=2.3.0', + 'audeer @ git+https://github.com/audeering/audeer.git@archive-num-workers-2', 'filelock', 'oyaml', ]