dask-setup 2.1.0__tar.gz → 2.2.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (62) hide show
  1. {dask_setup-2.1.0/src/dask_setup.egg-info → dask_setup-2.2.0}/PKG-INFO +31 -7
  2. dask_setup-2.1.0/PKG-INFO → dask_setup-2.2.0/README.md +26 -26
  3. {dask_setup-2.1.0 → dask_setup-2.2.0}/pyproject.toml +17 -2
  4. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/__init__.py +1 -1
  5. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/benchmark.py +192 -38
  6. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/cli.py +9 -4
  7. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/client.py +374 -137
  8. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/cluster.py +89 -14
  9. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/config.py +35 -0
  10. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/config_manager.py +56 -28
  11. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/dashboard.py +32 -1
  12. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/error_handling.py +18 -4
  13. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/legacy.py +3 -1
  14. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/logging.py +18 -9
  15. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/multinode.py +103 -29
  16. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/rechunk.py +42 -3
  17. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/reporting.py +31 -7
  18. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/resources.py +28 -5
  19. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/tempdir.py +31 -2
  20. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/tune.py +84 -29
  21. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/xarray.py +32 -16
  22. dask_setup-2.1.0/README.md → dask_setup-2.2.0/src/dask_setup.egg-info/PKG-INFO +50 -5
  23. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup.egg-info/SOURCES.txt +4 -0
  24. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup.egg-info/requires.txt +5 -1
  25. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_benchmark.py +205 -0
  26. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_cli.py +14 -3
  27. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_client.py +465 -39
  28. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_cluster.py +102 -0
  29. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_config_manager.py +151 -6
  30. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_dashboard.py +47 -0
  31. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_error_handling.py +61 -0
  32. dask_setup-2.2.0/tests/test_logging.py +105 -0
  33. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_multinode.py +327 -3
  34. dask_setup-2.2.0/tests/test_rechunk.py +70 -0
  35. dask_setup-2.2.0/tests/test_reporting.py +103 -0
  36. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_resources.py +101 -11
  37. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_setup_dask_client.py +74 -59
  38. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_tempdir.py +62 -0
  39. dask_setup-2.2.0/tests/test_tune.py +167 -0
  40. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_xarray_chunks.py +128 -0
  41. {dask_setup-2.1.0 → dask_setup-2.2.0}/LICENSE +0 -0
  42. {dask_setup-2.1.0 → dask_setup-2.2.0}/setup.cfg +0 -0
  43. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/callbacks.py +0 -0
  44. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/environment.py +0 -0
  45. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/exceptions.py +0 -0
  46. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/io_patterns.py +0 -0
  47. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/parquet.py +0 -0
  48. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/py.typed +0 -0
  49. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/schema/__init__.py +0 -0
  50. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/schema/profile_schema.json +0 -0
  51. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/topology.py +0 -0
  52. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/types.py +0 -0
  53. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup/workload.py +0 -0
  54. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup.egg-info/dependency_links.txt +0 -0
  55. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup.egg-info/entry_points.txt +0 -0
  56. {dask_setup-2.1.0 → dask_setup-2.2.0}/src/dask_setup.egg-info/top_level.txt +0 -0
  57. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_compression.py +0 -0
  58. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_config.py +0 -0
  59. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_exceptions.py +0 -0
  60. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_io_patterns.py +0 -0
  61. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_topology.py +0 -0
  62. {dask_setup-2.1.0 → dask_setup-2.2.0}/tests/test_types.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: dask_setup
3
- Version: 2.1.0
3
+ Version: 2.2.0
4
4
  Summary: HPC-tuned Dask helpers for single-node runs on NCI Gadi.
5
5
  Author: Sam Green
6
6
  Project-URL: Homepage, https://github.com/21centuryweather/dask_setup
@@ -11,12 +11,15 @@ Requires-Dist: dask>=2024.1.0
11
11
  Requires-Dist: distributed>=2024.1.0
12
12
  Requires-Dist: psutil>=5.9
13
13
  Requires-Dist: pyyaml>=6.0
14
+ Provides-Extra: multinode
15
+ Requires-Dist: dask-jobqueue>=0.8; extra == "multinode"
14
16
  Provides-Extra: dev
15
- Requires-Dist: ruff~=0.4; extra == "dev"
17
+ Requires-Dist: ruff==0.16.4; extra == "dev"
16
18
  Requires-Dist: pytest~=8.2; extra == "dev"
17
19
  Requires-Dist: pytest-cov~=5.0; extra == "dev"
18
20
  Requires-Dist: xarray; extra == "dev"
19
21
  Requires-Dist: numpy; extra == "dev"
22
+ Requires-Dist: dask-jobqueue>=0.8; extra == "dev"
20
23
  Dynamic: license-file
21
24
 
22
25
  # dask_setup
@@ -35,9 +38,9 @@ HPC-tuned Dask helpers for **NCI Gadi** and other PBS/SLURM systems. Wraps `dask
35
38
  from dask_setup import setup_dask_client
36
39
 
37
40
  # Pick a workload type and go
38
- client, cluster, dask_tmp = setup_dask_client("cpu") # heavy compute
39
- client, cluster, dask_tmp = setup_dask_client("io") # heavy file I/O
40
- client, cluster, dask_tmp = setup_dask_client("mixed") # both
41
+ client, cluster, dask_tmp = setup_dask_client(mode="interactive", workload_type="cpu") # heavy compute
42
+ client, cluster, dask_tmp = setup_dask_client(mode="interactive", workload_type="io") # heavy file I/O
43
+ client, cluster, dask_tmp = setup_dask_client(mode="interactive", workload_type="mixed") # both
41
44
  ```
42
45
 
43
46
  `dask_tmp` is the path to the spill/temp directory (on `$PBS_JOBFS` if available). Pass it to Rechunker, Zarr, or anywhere else you want fast local I/O.
@@ -95,14 +98,27 @@ client, cluster, shared_tmp = setup_dask_client(
95
98
  |-----------|---------|-------------|
96
99
  | `workload_type` | `"io"` | Worker topology: `"cpu"`, `"io"`, `"mixed"`, `"gpu"`, `"auto"` |
97
100
  | `max_workers` | all cores | Hard cap on worker count |
98
- | `reserve_mem_gb` | auto (20% RAM) | Memory held back for OS/cache (GiB) |
101
+ | `reserve_mem_gb` | auto | Memory held back for OS/cache (GiB): 20% of RAM, clamped to [4, 50] |
99
102
  | `max_mem_gb` | total RAM | Upper bound on Dask's total memory use |
100
103
  | `dashboard` | `True` | Start dashboard and print SSH tunnel hint |
101
104
  | `profile` | `None` | Named config profile |
102
- | `config` | `None` | Pre-built `DaskSetupConfig` object — mutually exclusive with `profile` |
105
+ | `config` | `None` | Pre-built `DaskSetupConfig` object — same layer as `profile`; if both are given, `profile` wins |
103
106
  | `mode` | `"auto"` | `"local"`, `"pbs"`, `"slurm"`, or `"auto"` (v2.0) |
104
107
  | `multi_node_config` | `None` | `MultiNodeConfig` for PBS/SLURM multi-node jobs (v2.0) |
105
108
 
109
+ Settings are layered, lowest to highest:
110
+
111
+ ```
112
+ library defaults < config= or profile= < explicit keyword arguments
113
+ ```
114
+
115
+ A parameter you leave unset inherits from the layer below. Passing a value
116
+ always overrides, even when that value happens to equal the default — so
117
+ `reserve_mem_gb=50.0` overrides a profile that says 40.0.
118
+
119
+ `config=` and `profile=` share a layer rather than stacking: pass one or the
120
+ other, and use explicit keyword arguments for the differences.
121
+
106
122
  ---
107
123
 
108
124
  ## Common Patterns
@@ -177,6 +193,14 @@ Tunnel from your laptop:
177
193
  Then open: http://localhost:8787
178
194
  ```
179
195
 
196
+ The login host is inferred from the compute node's DNS domain, so it is correct
197
+ at other sites too. Override it with `$DASK_SETUP_LOGIN_HOST` if the guess is
198
+ wrong:
199
+
200
+ ```bash
201
+ export DASK_SETUP_LOGIN_HOST=login.mycluster.edu
202
+ ```
203
+
180
204
  ---
181
205
 
182
206
  ## CLI
@@ -1,24 +1,3 @@
1
- Metadata-Version: 2.4
2
- Name: dask_setup
3
- Version: 2.1.0
4
- Summary: HPC-tuned Dask helpers for single-node runs on NCI Gadi.
5
- Author: Sam Green
6
- Project-URL: Homepage, https://github.com/21centuryweather/dask_setup
7
- Requires-Python: >=3.11
8
- Description-Content-Type: text/markdown
9
- License-File: LICENSE
10
- Requires-Dist: dask>=2024.1.0
11
- Requires-Dist: distributed>=2024.1.0
12
- Requires-Dist: psutil>=5.9
13
- Requires-Dist: pyyaml>=6.0
14
- Provides-Extra: dev
15
- Requires-Dist: ruff~=0.4; extra == "dev"
16
- Requires-Dist: pytest~=8.2; extra == "dev"
17
- Requires-Dist: pytest-cov~=5.0; extra == "dev"
18
- Requires-Dist: xarray; extra == "dev"
19
- Requires-Dist: numpy; extra == "dev"
20
- Dynamic: license-file
21
-
22
1
  # dask_setup
23
2
 
24
3
  [![CI](https://github.com/21centuryweather/dask_setup/workflows/CI/badge.svg)](https://github.com/21centuryweather/dask_setup/actions)
@@ -35,9 +14,9 @@ HPC-tuned Dask helpers for **NCI Gadi** and other PBS/SLURM systems. Wraps `dask
35
14
  from dask_setup import setup_dask_client
36
15
 
37
16
  # Pick a workload type and go
38
- client, cluster, dask_tmp = setup_dask_client("cpu") # heavy compute
39
- client, cluster, dask_tmp = setup_dask_client("io") # heavy file I/O
40
- client, cluster, dask_tmp = setup_dask_client("mixed") # both
17
+ client, cluster, dask_tmp = setup_dask_client(mode="interactive", workload_type="cpu") # heavy compute
18
+ client, cluster, dask_tmp = setup_dask_client(mode="interactive", workload_type="io") # heavy file I/O
19
+ client, cluster, dask_tmp = setup_dask_client(mode="interactive", workload_type="mixed") # both
41
20
  ```
42
21
 
43
22
  `dask_tmp` is the path to the spill/temp directory (on `$PBS_JOBFS` if available). Pass it to Rechunker, Zarr, or anywhere else you want fast local I/O.
@@ -95,14 +74,27 @@ client, cluster, shared_tmp = setup_dask_client(
95
74
  |-----------|---------|-------------|
96
75
  | `workload_type` | `"io"` | Worker topology: `"cpu"`, `"io"`, `"mixed"`, `"gpu"`, `"auto"` |
97
76
  | `max_workers` | all cores | Hard cap on worker count |
98
- | `reserve_mem_gb` | auto (20% RAM) | Memory held back for OS/cache (GiB) |
77
+ | `reserve_mem_gb` | auto | Memory held back for OS/cache (GiB): 20% of RAM, clamped to [4, 50] |
99
78
  | `max_mem_gb` | total RAM | Upper bound on Dask's total memory use |
100
79
  | `dashboard` | `True` | Start dashboard and print SSH tunnel hint |
101
80
  | `profile` | `None` | Named config profile |
102
- | `config` | `None` | Pre-built `DaskSetupConfig` object — mutually exclusive with `profile` |
81
+ | `config` | `None` | Pre-built `DaskSetupConfig` object — same layer as `profile`; if both are given, `profile` wins |
103
82
  | `mode` | `"auto"` | `"local"`, `"pbs"`, `"slurm"`, or `"auto"` (v2.0) |
104
83
  | `multi_node_config` | `None` | `MultiNodeConfig` for PBS/SLURM multi-node jobs (v2.0) |
105
84
 
85
+ Settings are layered, lowest to highest:
86
+
87
+ ```
88
+ library defaults < config= or profile= < explicit keyword arguments
89
+ ```
90
+
91
+ A parameter you leave unset inherits from the layer below. Passing a value
92
+ always overrides, even when that value happens to equal the default — so
93
+ `reserve_mem_gb=50.0` overrides a profile that says 40.0.
94
+
95
+ `config=` and `profile=` share a layer rather than stacking: pass one or the
96
+ other, and use explicit keyword arguments for the differences.
97
+
106
98
  ---
107
99
 
108
100
  ## Common Patterns
@@ -177,6 +169,14 @@ Tunnel from your laptop:
177
169
  Then open: http://localhost:8787
178
170
  ```
179
171
 
172
+ The login host is inferred from the compute node's DNS domain, so it is correct
173
+ at other sites too. Override it with `$DASK_SETUP_LOGIN_HOST` if the guess is
174
+ wrong:
175
+
176
+ ```bash
177
+ export DASK_SETUP_LOGIN_HOST=login.mycluster.edu
178
+ ```
179
+
180
180
  ---
181
181
 
182
182
  ## CLI
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "dask_setup"
7
- version = "2.1.0"
7
+ version = "2.2.0"
8
8
  description = "HPC-tuned Dask helpers for single-node runs on NCI Gadi."
9
9
  authors = [{name="Sam Green"}]
10
10
  readme = "README.md"
@@ -23,13 +23,28 @@ Homepage = "https://github.com/21centuryweather/dask_setup"
23
23
  dask-setup = "dask_setup.cli:main"
24
24
 
25
25
  [project.optional-dependencies]
26
+ # Multi-node PBS/SLURM cluster support. Optional: the single-node path never
27
+ # imports it, and setup_pbs_cluster()/setup_slurm_cluster() raise a helpful
28
+ # ImportError if it is missing.
29
+ multinode = [
30
+ "dask-jobqueue>=0.8",
31
+ ]
32
+
26
33
  dev = [
27
- "ruff ~= 0.4",
34
+ # Pinned exactly, not a range. `ruff ~= 0.4` means >=0.4,<1.0, so CI picked
35
+ # up whatever ruff had released that morning -- and ruff's formatter changes
36
+ # between minor versions. That turns `ruff format --check` into a job that
37
+ # can fail on a commit that touched nothing. Bump this deliberately.
38
+ "ruff == 0.16.4",
28
39
  "pytest ~= 8.2",
29
40
  "pytest-cov ~= 5.0",
30
41
  # Xarray integration testing
31
42
  "xarray",
32
43
  "numpy",
44
+ # Multi-node tests exercise a real PBSJob. Without this, CI silently skipped
45
+ # the multi-node path -- including the regression test for the 4x worker
46
+ # under-provisioning fixed in 2.2.0.
47
+ "dask-jobqueue>=0.8",
33
48
  ]
34
49
 
35
50
  [tool.setuptools.package-data]
@@ -281,7 +281,7 @@ except ImportError:
281
281
  )
282
282
 
283
283
 
284
- __version__ = "2.1.0"
284
+ __version__ = "2.2.0"
285
285
 
286
286
  __all__ = [
287
287
  # Core API — always available
@@ -46,11 +46,15 @@ from collections.abc import Callable, Sequence
46
46
  from dataclasses import dataclass, field
47
47
  from typing import TYPE_CHECKING, Any
48
48
 
49
+ from .logging import get_logger
50
+
49
51
  if TYPE_CHECKING:
50
52
  from dask.distributed import Client
51
53
 
52
54
  from .config import DaskSetupConfig
53
55
 
56
+ logger = get_logger("benchmark")
57
+
54
58
  __all__ = [
55
59
  "BenchmarkResult",
56
60
  "ScalingResult",
@@ -125,7 +129,10 @@ class BenchmarkResult:
125
129
  wall_time_std : float
126
130
  Standard deviation of per-repeat wall times (0.0 for a single repeat).
127
131
  peak_memory_gib : float
128
- Maximum in-memory data across all workers at the end of the run.
132
+ Highest total worker process memory (RSS summed across workers)
133
+ observed *during* the timed runs, sampled every 0.2 s. If in-flight
134
+ sampling was unavailable this falls back to a post-run reading and an
135
+ explanatory note is appended to ``errors``.
129
136
  spill_gib : float
130
137
  Total data written to disk spill storage during the run.
131
138
  n_tasks : int
@@ -461,6 +468,55 @@ class ChunkImpactResult:
461
468
  # ---------------------------------------------------------------------------
462
469
 
463
470
 
471
+ @contextlib.contextmanager
472
+ def _sample_peak_memory(client: Client) -> Any:
473
+ """Record peak cluster memory *while* the enclosed block runs.
474
+
475
+ Reading ``cluster_report(client).peak_memory_gib`` after ``.compute()``
476
+ returns does not measure a peak: by then the graph is done and the workers
477
+ have already released the data, so the figure is close to zero for exactly
478
+ the workloads whose memory use matters.
479
+
480
+ Yields a dict that is filled in on exit::
481
+
482
+ {"peak_gib": float, "sampled": bool}
483
+
484
+ ``sampled`` is ``False`` when no in-flight sampling was possible (older
485
+ ``distributed``, or a scheduler that refused the periodic callback), which
486
+ tells the caller to fall back to the post-run reading rather than report a
487
+ peak of 0.0.
488
+ """
489
+ result: dict[str, Any] = {"peak_gib": 0.0, "sampled": False}
490
+
491
+ label = "dask_setup_benchmark"
492
+ sampler = None
493
+ ctx = None
494
+ try:
495
+ from distributed.diagnostics import MemorySampler
496
+
497
+ sampler = MemorySampler()
498
+ # measure="process" is total RSS across workers -- the number that
499
+ # decides whether a job fits in its memory allocation.
500
+ ctx = sampler.sample(label, client=client, measure="process", interval=0.2)
501
+ ctx.__enter__()
502
+ except Exception as e: # pragma: no cover - depends on distributed version
503
+ logger.debug("In-flight memory sampling unavailable", error=str(e))
504
+ ctx = None
505
+
506
+ try:
507
+ yield result
508
+ finally:
509
+ if ctx is not None:
510
+ try:
511
+ ctx.__exit__(None, None, None)
512
+ samples = (sampler.samples or {}).get(label) or []
513
+ peak_bytes = max((float(b) for _t, b in samples), default=0.0)
514
+ result["peak_gib"] = peak_bytes / (1024**3)
515
+ result["sampled"] = bool(samples)
516
+ except Exception as e:
517
+ logger.debug("Could not read memory samples", error=str(e))
518
+
519
+
464
520
  def _measure_one(
465
521
  ds: Any,
466
522
  operation_fn: Callable[[Any], Any],
@@ -507,23 +563,24 @@ def _measure_one(
507
563
  with contextlib.suppress(Exception):
508
564
  n_workers = len(client.scheduler_info().get("workers", {}))
509
565
 
510
- # Optional warmup
566
+ # Optional warmup (not sampled -- it is not part of the measured run)
511
567
  if warmup:
512
568
  try:
513
569
  operation_fn(ds).compute()
514
570
  except Exception as e:
515
571
  errors.append(f"Warmup failed: {e}")
516
572
 
517
- # Timed runs
518
- for i in range(repeats):
519
- t0 = time.monotonic()
520
- try:
521
- operation_fn(ds).compute()
522
- except Exception as e:
523
- errors.append(f"Run {i + 1}/{repeats} failed: {e}")
524
- times.append(float("inf"))
525
- else:
526
- times.append(time.monotonic() - t0)
573
+ # Timed runs, with memory sampled while they are in flight
574
+ with _sample_peak_memory(client) as mem_samples:
575
+ for i in range(repeats):
576
+ t0 = time.monotonic()
577
+ try:
578
+ operation_fn(ds).compute()
579
+ except Exception as e:
580
+ errors.append(f"Run {i + 1}/{repeats} failed: {e}")
581
+ times.append(float("inf"))
582
+ else:
583
+ times.append(time.monotonic() - t0)
527
584
 
528
585
  # Filter infinities (failed runs)
529
586
  valid_times = [t for t in times if t != float("inf")]
@@ -535,8 +592,13 @@ def _measure_one(
535
592
  from .reporting import cluster_report
536
593
 
537
594
  report = cluster_report(client)
538
- peak_mem = report.peak_memory_gib
539
595
  spill = report.total_spill_gib
596
+ # Prefer the in-flight peak; the post-run reading is a floor at best.
597
+ if mem_samples["sampled"]:
598
+ peak_mem = mem_samples["peak_gib"]
599
+ else:
600
+ peak_mem = report.peak_memory_gib
601
+ errors.append("Peak memory sampled after the run; treat it as a lower bound")
540
602
  except Exception as e:
541
603
  errors.append(f"Metrics collection failed: {e}")
542
604
 
@@ -701,15 +763,16 @@ def scaling_analysis(
701
763
  repeats: int = 1,
702
764
  warmup: bool = False,
703
765
  fallback_on_detection_failure: bool = True,
766
+ mode: str = "local",
704
767
  plot: bool = False,
705
768
  verbose: bool = False,
706
769
  ) -> ScalingResult:
707
770
  """Measure parallel scaling across different worker counts.
708
771
 
709
- For each entry in *worker_counts*, a fresh cluster is created with
710
- ``max_workers`` set to that count and the operation is timed. The
711
- resulting :class:`ScalingResult` includes speedup and efficiency
712
- relative to the single-worker baseline.
772
+ For each entry in *worker_counts*, a fresh **local** cluster is created
773
+ with ``max_workers`` set to that count and the operation is timed. The
774
+ resulting :class:`ScalingResult` includes speedup and efficiency relative
775
+ to the single-worker baseline.
713
776
 
714
777
  Parameters
715
778
  ----------
@@ -730,6 +793,13 @@ def scaling_analysis(
730
793
  Run one un-timed warmup pass before timing.
731
794
  fallback_on_detection_failure:
732
795
  Passed to :func:`~dask_setup.client.setup_dask_client`.
796
+ mode:
797
+ Cluster mode passed to :func:`~dask_setup.client.setup_dask_client`.
798
+ Defaults to ``"local"`` so that the sweep always uses a
799
+ ``LocalCluster`` on the current node — this prevents auto-detection
800
+ from dispatching to PBS/SLURM and submitting child scheduler jobs,
801
+ which would time out and produce NaN results. Pass ``"auto"`` only
802
+ if you specifically want multi-node dispatch for each sweep step.
733
803
  plot:
734
804
  If ``True``, call :meth:`ScalingResult.plot` and display the figure
735
805
  (requires matplotlib).
@@ -748,11 +818,33 @@ def scaling_analysis(
748
818
  counts = list(worker_counts)
749
819
 
750
820
  if base_config is None:
751
- base_config = DaskSetupConfig(fallback_on_detection_failure=True)
821
+ # workload_type="cpu" is the only default under which this sweep means
822
+ # anything: decide_topology() pins n_workers=1 for "io" (and for "gpu"
823
+ # with no GPU present) regardless of max_workers, so an "io" sweep
824
+ # builds the same single-worker cluster at every point and the
825
+ # "scaling curve" is just timing noise.
826
+ base_config = DaskSetupConfig(workload_type="cpu", fallback_on_detection_failure=True)
827
+
828
+ # Use tqdm for progress if available (works in both terminals and Jupyter).
829
+ # tqdm.auto automatically picks the right bar type (notebook vs terminal).
830
+ try:
831
+ from tqdm.auto import tqdm as _tqdm
832
+
833
+ count_iter = _tqdm(counts, desc="scaling sweep", unit="config")
834
+ except ImportError:
835
+ count_iter = counts # type: ignore[assignment]
752
836
 
753
837
  raw_results: list[BenchmarkResult] = []
754
838
 
755
- for nw in counts:
839
+ for nw in count_iter:
840
+ # Show live status so notebook users know what's running right now.
841
+ # tqdm will show this as a postfix label; without tqdm it's a plain print.
842
+ _status = f"workers={nw} ({len(raw_results) + 1}/{len(counts)})"
843
+ try:
844
+ count_iter.set_postfix_str(_status) # type: ignore[union-attr]
845
+ except AttributeError:
846
+ print(f"[scaling_analysis] running {_status}…", flush=True)
847
+
756
848
  # Clone config with this worker count
757
849
  cfg_dict = base_config.to_dict()
758
850
  cfg_dict["max_workers"] = nw
@@ -766,8 +858,12 @@ def scaling_analysis(
766
858
 
767
859
  try:
768
860
  client, cluster, _tmp = setup_dask_client(
861
+ # cfg already carries max_workers=nw and adaptive=False; every
862
+ # field of a config= object is now honoured, so there is no
863
+ # need to re-pass them as explicit keyword arguments.
769
864
  config=cfg,
770
865
  fallback_on_detection_failure=fallback_on_detection_failure,
866
+ mode=mode,
771
867
  dashboard=False,
772
868
  )
773
869
  result = _measure_one(
@@ -803,15 +899,36 @@ def scaling_analysis(
803
899
  if not baseline_time or baseline_time != baseline_time: # nan check
804
900
  baseline_time = 1.0
805
901
 
902
+ # Efficiency is speedup relative to the *ratio* of workers, not to the
903
+ # absolute worker count. Dividing by nw made a sweep that starts anywhere
904
+ # other than 1 worker report a fraction of its true efficiency: a perfect
905
+ # (4, 8) sweep scored 0.25 at its own baseline.
906
+ actual_counts = [r.n_workers or nw for r, nw in zip(raw_results, counts, strict=False)]
907
+ baseline_workers = actual_counts[0] if actual_counts else 1
908
+
806
909
  speedups = []
807
910
  efficiencies = []
808
- for r, nw in zip(raw_results, counts, strict=False):
911
+ for r, nw in zip(raw_results, actual_counts, strict=False):
809
912
  if r.wall_time_seconds and r.wall_time_seconds == r.wall_time_seconds:
810
913
  sp = baseline_time / r.wall_time_seconds
811
914
  else:
812
915
  sp = float("nan")
813
916
  speedups.append(sp)
814
- efficiencies.append(sp / nw if nw > 0 else float("nan"))
917
+ worker_ratio = (nw / baseline_workers) if (nw > 0 and baseline_workers > 0) else 0.0
918
+ efficiencies.append(sp / worker_ratio if worker_ratio > 0 else float("nan"))
919
+
920
+ # A sweep in which the cluster never actually changed size produces a
921
+ # meaningless curve. Say so rather than letting the caller read noise as
922
+ # a scaling result.
923
+ distinct = {n for n in actual_counts if n > 0}
924
+ if len(counts) > 1 and len(distinct) == 1:
925
+ logger.warning(
926
+ "Scaling sweep ran at a constant worker count; the curve is not a scaling result",
927
+ requested=counts,
928
+ actual=actual_counts[0] if actual_counts else 0,
929
+ workload_type=base_config.workload_type,
930
+ hint="workload_type='io'/'gpu' pin n_workers=1; use workload_type='cpu' or 'mixed'",
931
+ )
815
932
 
816
933
  scaling = ScalingResult(
817
934
  results=raw_results,
@@ -823,10 +940,18 @@ def scaling_analysis(
823
940
  if plot:
824
941
  fig = scaling.plot()
825
942
  if fig is not None:
826
- with contextlib.suppress(Exception):
827
- import matplotlib.pyplot as plt
943
+ # In Jupyter, use IPython's display() so the figure renders inline
944
+ # in the cell output rather than via plt.show() which can misbehave
945
+ # in some notebook backends.
946
+ try:
947
+ from IPython.display import display as _ipy_display
948
+
949
+ _ipy_display(fig)
950
+ except ImportError:
951
+ with contextlib.suppress(Exception):
952
+ import matplotlib.pyplot as plt
828
953
 
829
- plt.show()
954
+ plt.show()
830
955
 
831
956
  return scaling
832
957
 
@@ -902,9 +1027,25 @@ def chunk_impact(
902
1027
  _generate_auto_chunks(dims) if auto_chunks and dims else [{}]
903
1028
  ) # single run with no rechunking
904
1029
 
1030
+ # Use tqdm for progress if available (works in both terminals and Jupyter).
1031
+ # tqdm.auto automatically picks the right bar type (notebook vs terminal).
1032
+ try:
1033
+ from tqdm.auto import tqdm as _tqdm
1034
+
1035
+ chunk_iter = _tqdm(chunk_sizes, desc="chunk sweep", unit="config")
1036
+ except ImportError:
1037
+ chunk_iter = chunk_sizes # type: ignore[assignment]
1038
+
905
1039
  raw_results: list[BenchmarkResult] = []
906
1040
 
907
- for cs in chunk_sizes:
1041
+ for cs in chunk_iter:
1042
+ # Show live status so notebook users know what's running right now.
1043
+ _status = f"{cs} ({len(raw_results) + 1}/{len(chunk_sizes)})"
1044
+ try:
1045
+ chunk_iter.set_postfix_str(_status) # type: ignore[union-attr]
1046
+ except AttributeError:
1047
+ print(f"[chunk_impact] running {_status}…", flush=True)
1048
+
908
1049
  errors: list[str] = []
909
1050
  ds_chunked = ds
910
1051
 
@@ -950,10 +1091,18 @@ def chunk_impact(
950
1091
  first_dim = next(iter(dims), None)
951
1092
  fig = impact.plot(dim=first_dim)
952
1093
  if fig is not None:
953
- with contextlib.suppress(Exception):
954
- import matplotlib.pyplot as plt
1094
+ # In Jupyter, use IPython's display() so the figure renders inline
1095
+ # in the cell output rather than via plt.show() which can misbehave
1096
+ # in some notebook backends.
1097
+ try:
1098
+ from IPython.display import display as _ipy_display
955
1099
 
956
- plt.show()
1100
+ _ipy_display(fig)
1101
+ except ImportError:
1102
+ with contextlib.suppress(Exception):
1103
+ import matplotlib.pyplot as plt
1104
+
1105
+ plt.show()
957
1106
 
958
1107
  return impact
959
1108
 
@@ -1142,19 +1291,24 @@ def run_synthetic_benchmark(
1142
1291
  if verbose:
1143
1292
  print(f"Cluster ready: {n_workers} workers. Running {operation!r} × {repeats} …")
1144
1293
 
1145
- for i in range(repeats):
1146
- t0 = time.monotonic()
1147
- try:
1148
- op_fn(arr).compute()
1149
- except Exception as e:
1150
- errors.append(f"Run {i + 1} failed: {e}")
1151
- times.append(float("inf"))
1152
- else:
1153
- times.append(time.monotonic() - t0)
1294
+ with _sample_peak_memory(client) as mem_samples:
1295
+ for i in range(repeats):
1296
+ t0 = time.monotonic()
1297
+ try:
1298
+ op_fn(arr).compute()
1299
+ except Exception as e:
1300
+ errors.append(f"Run {i + 1} failed: {e}")
1301
+ times.append(float("inf"))
1302
+ else:
1303
+ times.append(time.monotonic() - t0)
1154
1304
 
1155
1305
  report = cluster_report(client)
1156
- peak_mem = report.peak_memory_gib
1157
1306
  spill = report.total_spill_gib
1307
+ if mem_samples["sampled"]:
1308
+ peak_mem = mem_samples["peak_gib"]
1309
+ else:
1310
+ peak_mem = report.peak_memory_gib
1311
+ errors.append("Peak memory sampled after the run; treat it as a lower bound")
1158
1312
 
1159
1313
  except Exception as e:
1160
1314
  errors.append(f"Cluster error: {e}")
@@ -4,6 +4,7 @@ from __future__ import annotations
4
4
 
5
5
  import argparse
6
6
  import sys
7
+ from dataclasses import replace
7
8
  from typing import Any
8
9
 
9
10
  import yaml
@@ -161,10 +162,14 @@ def cmd_create_profile(args: argparse.Namespace) -> int:
161
162
  print(f" Base profile '{args.from_profile}' not found.", file=sys.stderr)
162
163
  return 1
163
164
 
164
- # Copy configuration and update name
165
- new_config = base_profile.config
166
- new_config.name = args.name
167
- new_config.description = f"Based on {args.from_profile}"
165
+ # Copy the configuration -- assigning base_profile.config directly
166
+ # and then setting .name on it mutates the source profile, which
167
+ # for a builtin means renaming it for the rest of the process.
168
+ new_config = replace(
169
+ base_profile.config,
170
+ name=args.name,
171
+ description=f"Based on {args.from_profile}",
172
+ )
168
173
 
169
174
  from .config import ConfigProfile
170
175