Skip to content

Cluster helpers

patchworks.make_local_cluster(use_gpu: bool = False, n_workers: int | None = None, threads_per_worker: int = 1, memory_limit: str | int | None = 'auto', **cluster_kwargs)

Create a process-based Dask cluster for tiled processing.

Always uses worker subprocesses (processes=True). An in-process (threaded) worker breaks the label merge when segment_fn holds the GIL — see the patchworks docs for details.

For GPU work defaults to a single worker (one CUDA context, no contention). For CPU scales to available cores.

Parameters:

Name Type Description Default
use_gpu bool

Single-worker cluster for GPU. When False, use multiple CPU workers.

False
n_workers int | None

Override the worker count. Defaults to 1 for GPU, min(8, cpu_count).

None
threads_per_worker int

Keep at 1 so a GIL-holding tile function doesn't block heartbeats.

1
memory_limit str | int | None

Per-worker memory cap (e.g. "8GB"). "auto" (default) splits the memory this job may use (SLURM, cgroup and free RAM, see safe_worker_count) evenly across the workers. A limit is what lets a worker spill, pause and restart before the OOM killer ends the whole job; None disables it and all of that with it. (Not distributed's own "auto", which scales by threads per core: one single-threaded GPU worker on a 32-core node would get 1/32 of it.)

'auto'
**cluster_kwargs

Extra arguments forwarded to dask.distributed.LocalCluster.

{}

Returns:

Type Description
(client, cluster)

Examples:

>>> client, cluster = make_local_cluster(use_gpu=True)
>>> print("dashboard:", client.dashboard_link)
>>> result = tile_process("image.zarr", fn, write_to="labels.zarr")
>>> client.close(); cluster.close()
Source code in src/patchworks/_cluster.py
def make_local_cluster(
    use_gpu: bool = False,
    n_workers: int | None = None,
    threads_per_worker: int = 1,
    memory_limit: str | int | None = "auto",
    **cluster_kwargs,
):
    """Create a process-based Dask cluster for tiled processing.

    Always uses worker subprocesses (``processes=True``). An in-process
    (threaded) worker breaks the label merge when ``segment_fn`` holds the
    GIL — see the patchworks docs for details.

    For GPU work defaults to a single worker (one CUDA context, no contention).
    For CPU scales to available cores.

    Parameters
    ----------
    use_gpu:
        Single-worker cluster for GPU. When False, use multiple CPU workers.
    n_workers:
        Override the worker count. Defaults to 1 for GPU, min(8, cpu_count).
    threads_per_worker:
        Keep at 1 so a GIL-holding tile function doesn't block heartbeats.
    memory_limit:
        Per-worker memory cap (e.g. ``"8GB"``). ``"auto"`` (default) splits
        the memory this job may use (SLURM, cgroup and free RAM, see
        ``safe_worker_count``) evenly across the workers. A limit is what
        lets a worker spill, pause and restart before the OOM killer ends the
        whole job; ``None`` disables it and all of that with it. (Not
        distributed's own ``"auto"``, which scales by threads per core: one
        single-threaded GPU worker on a 32-core node would get 1/32 of it.)
    **cluster_kwargs:
        Extra arguments forwarded to ``dask.distributed.LocalCluster``.

    Returns
    -------
    (client, cluster)

    Examples
    --------
    >>> client, cluster = make_local_cluster(use_gpu=True)  # doctest: +SKIP
    >>> print("dashboard:", client.dashboard_link)  # doctest: +SKIP
    >>> result = tile_process("image.zarr", fn, write_to="labels.zarr")  # doctest: +SKIP
    >>> client.close(); cluster.close()  # doctest: +SKIP
    """
    from dask.distributed import Client, LocalCluster

    if n_workers is None:
        n_workers = 1 if use_gpu else min(8, cpu_allocation())
    if memory_limit == "auto":
        memory_limit = max(1, _get_available_memory() // n_workers)

    cluster = LocalCluster(
        processes=True,
        n_workers=n_workers,
        threads_per_worker=threads_per_worker,
        memory_limit=memory_limit,
        **cluster_kwargs,
    )
    client = Client(cluster)
    logger.info(
        "Started %d-worker process cluster (use_gpu=%s). Dashboard: %s",
        n_workers,
        use_gpu,
        client.dashboard_link,
    )
    return client, cluster