slurm-workflows 1.0.0__py3-none-any.whl

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.
@@ -0,0 +1 @@
1
+ from .slurm_pilot_executor import SlurmPilotExecutor, check_for_error
@@ -0,0 +1,44 @@
1
+ """Run ds-service."""
2
+
3
+ import shlex
4
+ import subprocess
5
+
6
+ from .utils import (
7
+ Closeable,
8
+ terminate_gracefully,
9
+ ignoring_sigint,
10
+ )
11
+
12
+
13
+ class DsService(Closeable):
14
+ """Run ds-service server locally"""
15
+
16
+ def __init__(
17
+ self,
18
+ host: str = "0.0.0.0",
19
+ port: int = 5051,
20
+ server_exe: str = "ds-service",
21
+ ):
22
+ self.address = f"{host}:{port}"
23
+ self.server_exe = server_exe
24
+ self._proc: subprocess.Popen | None = None
25
+
26
+ def start(self):
27
+ cmd = self.server_exe + f" --address {self.address}"
28
+ cmd = shlex.split(cmd)
29
+ print("Starting server ...")
30
+ print("executing: ", " ".join(cmd))
31
+
32
+ assert self._proc is None
33
+ with ignoring_sigint():
34
+ self._proc = subprocess.Popen(
35
+ cmd,
36
+ stdin=subprocess.DEVNULL,
37
+ )
38
+
39
+ print(f"server address: {self.address}")
40
+
41
+ def close(self):
42
+ if self._proc is not None:
43
+ terminate_gracefully(self._proc)
44
+ self._proc = None
@@ -0,0 +1,57 @@
1
+ """Start a Jupyter Lab instance."""
2
+
3
+ import sys
4
+ import subprocess
5
+ from datetime import datetime
6
+ from pathlib import Path
7
+
8
+ import click
9
+ import platformdirs
10
+
11
+ from .templates import render_template
12
+ from .slurm_utils import submit_sbatch_job
13
+
14
+ JUPYTER_EXE = "jupyter"
15
+
16
+
17
+ @click.command()
18
+ @click.option(
19
+ "--setup-script",
20
+ type=click.Path(exists=True, file_okay=True, dir_okay=False, path_type=Path),
21
+ required=True,
22
+ help="Path to setup script.",
23
+ )
24
+ @click.argument("sbatch-args", nargs=-1)
25
+ def run_jupyter(
26
+ sbatch_args: list[str],
27
+ setup_script: Path,
28
+ ):
29
+ """Start a Jupyter Lab instance."""
30
+ print("Sbatch args: ", " ".join(sbatch_args))
31
+
32
+ name = "jupyter"
33
+
34
+ script = render_template(
35
+ "run_jupyter:script_template",
36
+ setup_script=setup_script,
37
+ jupyter_executable=JUPYTER_EXE,
38
+ )
39
+
40
+ now = datetime.now().isoformat()
41
+ work_dir = platformdirs.user_cache_path(appname=f"run-jupyter") / now
42
+ work_dir.mkdir(parents=True)
43
+
44
+ try:
45
+ job = submit_sbatch_job(
46
+ name=name, sbatch_args=sbatch_args, script=script, work_dir=work_dir
47
+ )
48
+
49
+ print(f"Job ID: {job.job_id}")
50
+ print(f"Output file: {job.output_file!s}")
51
+ except subprocess.CalledProcessError as cp:
52
+ print(f"Failed to submit job: {cp.returncode}")
53
+ if cp.stdout.strip():
54
+ print(cp.stdout)
55
+ if cp.stderr.strip():
56
+ print(cp.stderr)
57
+ sys.exit(1)
@@ -0,0 +1,384 @@
1
+ """Pilot workers for slurm."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import time
6
+ import json
7
+ import pickle
8
+ import logging
9
+ import subprocess
10
+ from pathlib import Path
11
+ from datetime import datetime
12
+ from dataclasses import dataclass, field
13
+ from typing import Callable, Iterable, Any, cast
14
+
15
+ import platformdirs
16
+ import cloudpickle
17
+ from typeguard import typechecked
18
+ from tqdm import tqdm
19
+ from ds_service_client import Client, TaskState
20
+
21
+ from .slurm_utils import (
22
+ get_running_jobids,
23
+ cancel_jobs,
24
+ submit_sbatch_job,
25
+ SlurmJob,
26
+ )
27
+
28
+ from .utils import (
29
+ RemoteExecutionError,
30
+ gen_random_string,
31
+ LOG_FORMAT,
32
+ LOG_LEVEL,
33
+ )
34
+
35
+ from .templates import render_template
36
+
37
+ NoOutput = object()
38
+
39
+ POLL_INTERVAL_S: float = 0.1
40
+
41
+
42
+ @dataclass
43
+ class Task:
44
+ task_id: str
45
+ queue: list[str]
46
+ priority: float
47
+ function: Callable | str
48
+ input: tuple
49
+ output: Any
50
+
51
+
52
+ @dataclass
53
+ class WorkerGroup:
54
+ name: str
55
+ sbatch_args: list[str]
56
+ is_batch_worker: bool
57
+ worker_exe: str
58
+ actor_class_name: str
59
+ setup_script: str
60
+ python_paths: list[str]
61
+ workers: dict[str, SlurmJob] = field(default_factory=dict, compare=False)
62
+ next_worker_index: int = field(default=0, compare=False)
63
+
64
+
65
+ class SlurmPilotExecutor:
66
+ def __init__(
67
+ self,
68
+ server_address: str,
69
+ work_dir: Path | str | None = None,
70
+ ):
71
+ self.executor_id = gen_random_string()
72
+ self.next_task_index = 0
73
+
74
+ self.server_address = server_address
75
+ self.client = Client(server_address)
76
+
77
+ if work_dir is None:
78
+ now = datetime.now().isoformat()
79
+ work_dir = platformdirs.user_cache_path(appname=f"slurm-pilot") / now
80
+ self.work_dir = Path(work_dir)
81
+ self.work_dir.mkdir(parents=True, exist_ok=True)
82
+
83
+ self.logger = logging.getLogger("pilot_coordinator")
84
+ self.logger.setLevel(LOG_LEVEL)
85
+ handler = logging.FileHandler(self.work_dir / "coordinator.log", delay=True)
86
+ handler.setLevel(LOG_LEVEL)
87
+ formatter = logging.Formatter(LOG_FORMAT)
88
+ handler.setFormatter(formatter)
89
+ self.logger.addHandler(handler)
90
+
91
+ self.groups: dict[str, WorkerGroup] = {}
92
+
93
+ @typechecked
94
+ def define_worker(
95
+ self,
96
+ name: str,
97
+ sbatch_args: list[str],
98
+ setup_script: str,
99
+ worker_exe: str = "slurm-pilot-worker",
100
+ is_batch_worker: bool = False,
101
+ actor_class_name: str | None = None,
102
+ python_paths: list[str | Path] | None = None,
103
+ add_cwd_to_python_path: bool = True,
104
+ ) -> None:
105
+ python_str_paths: list[str] = []
106
+ if python_paths is not None:
107
+ for path in python_paths:
108
+ python_str_paths.append(str(path))
109
+ if add_cwd_to_python_path:
110
+ python_str_paths.append(str(Path.cwd()))
111
+
112
+ if actor_class_name is None:
113
+ actor_class_name = ""
114
+
115
+ group = WorkerGroup(
116
+ name=name,
117
+ sbatch_args=sbatch_args,
118
+ worker_exe=worker_exe,
119
+ is_batch_worker=is_batch_worker,
120
+ actor_class_name=actor_class_name,
121
+ setup_script=setup_script,
122
+ python_paths=python_str_paths,
123
+ )
124
+
125
+ if group.name in self.groups:
126
+ assert self.groups[group.name] == group
127
+ else:
128
+ self.groups[group.name] = group
129
+
130
+ def _add_worker(self, group: WorkerGroup) -> None:
131
+ worker_index = group.next_worker_index
132
+ group.next_worker_index += 1
133
+ name = f"slurm_pilot_worker.{group.name}.{worker_index}"
134
+
135
+ worker_script = render_template(
136
+ "slurm_pilot:worker_script",
137
+ group=group.name,
138
+ name=name,
139
+ server_address=self.server_address,
140
+ worker_exe=group.worker_exe,
141
+ work_dir=str(self.work_dir),
142
+ python_paths_json=json.dumps(group.python_paths),
143
+ setup_script=group.setup_script,
144
+ actor_class_name=group.actor_class_name,
145
+ )
146
+ worker_script_path = self.work_dir / f"{name}.sh"
147
+ worker_script_path.write_text(worker_script)
148
+ worker_script_path.chmod(0o755)
149
+
150
+ worker_sbatch_script = render_template(
151
+ "slurm_pilot:worker_sbatch_script",
152
+ is_batch_worker=group.is_batch_worker,
153
+ worker_script_path=worker_script_path,
154
+ )
155
+
156
+ self.logger.info("Starting worker %s", name)
157
+ try:
158
+ slurm_job = submit_sbatch_job(
159
+ name=name,
160
+ sbatch_args=group.sbatch_args,
161
+ script=worker_sbatch_script,
162
+ work_dir=self.work_dir,
163
+ )
164
+ group.workers[name] = slurm_job
165
+ except subprocess.CalledProcessError as cp:
166
+ print(f"Failed to submit slurm job: returncode={cp.returncode}")
167
+ if cp.stdout.strip():
168
+ print(cp.stdout)
169
+ if cp.stderr.strip():
170
+ print(cp.stderr)
171
+ raise cp
172
+
173
+ @typechecked
174
+ def scale_workers(self, name: str, count: int) -> None:
175
+ assert name in self.groups, "Unknown worker type"
176
+
177
+ group = self.groups[name]
178
+ if len(group.workers) < count:
179
+ to_hire = count - len(group.workers)
180
+ for _ in range(to_hire):
181
+ self._add_worker(group)
182
+
183
+ if len(group.workers) > count:
184
+ to_retire = len(group.workers) - count
185
+
186
+ try:
187
+ running_jobids = get_running_jobids()
188
+ except subprocess.CalledProcessError as cp:
189
+ print(
190
+ f"Failed to get running slurm job ids: returncode={cp.returncode}"
191
+ )
192
+ if cp.stdout.strip():
193
+ print(cp.stdout)
194
+ if cp.stderr.strip():
195
+ print(cp.stderr)
196
+ raise RuntimeError("Failed to get running slurm job ids")
197
+ except Exception:
198
+ raise RuntimeError("Failed to get running slurm job ids")
199
+
200
+ to_cancel_jobids = []
201
+ for _ in range(to_retire):
202
+ _, worker = group.workers.popitem()
203
+ self.logger.info("Canceling worker: %s", worker.name)
204
+ if worker.job_id in running_jobids:
205
+ to_cancel_jobids.append(worker.job_id)
206
+
207
+ if not to_cancel_jobids:
208
+ return
209
+
210
+ try:
211
+ cancel_jobs(to_cancel_jobids)
212
+ except subprocess.CalledProcessError as cp:
213
+ print(f"Failed to cancel slurm jobs: returncode={cp.returncode}")
214
+ if cp.stdout.strip():
215
+ print(cp.stdout)
216
+ if cp.stderr.strip():
217
+ print(cp.stderr)
218
+ raise RuntimeError("Failed to cancel slurm jobs")
219
+ except Exception:
220
+ raise RuntimeError("Failed to cancel slurm jobs")
221
+
222
+ def _submit(
223
+ self,
224
+ queue: list[str],
225
+ fn: Callable | str,
226
+ *args,
227
+ **kwargs,
228
+ ) -> Task:
229
+ priority = time.perf_counter()
230
+ function_bytes = cloudpickle.dumps(fn, protocol=pickle.HIGHEST_PROTOCOL)
231
+ input_bytes = cloudpickle.dumps(
232
+ (args, kwargs), protocol=pickle.HIGHEST_PROTOCOL
233
+ )
234
+
235
+ task_id = f"{self.executor_id}:{self.next_task_index}"
236
+ self.next_task_index += 1
237
+
238
+ task = Task(
239
+ task_id=task_id,
240
+ queue=queue,
241
+ priority=priority,
242
+ function=fn,
243
+ input=(args, kwargs),
244
+ output=NoOutput,
245
+ )
246
+
247
+ self.client.task_add(
248
+ task_id=task_id,
249
+ queue=queue,
250
+ priority=priority,
251
+ function=function_bytes,
252
+ input=input_bytes,
253
+ )
254
+ return task
255
+
256
+ @typechecked
257
+ def submit(
258
+ self, queue: str | list[str], fn: Callable | str, *args, **kwargs
259
+ ) -> Task:
260
+ if isinstance(queue, str):
261
+ queue = [queue]
262
+
263
+ return self._submit(queue, fn, *args, **kwargs)
264
+
265
+ def _as_completed(self, tasks: list[Task]) -> Iterable[Task]:
266
+ pending: list[Task] = []
267
+ for task in tasks:
268
+ if task.output is NoOutput:
269
+ pending.append(task)
270
+ else:
271
+ yield task
272
+
273
+ while pending:
274
+ # Status for every pending task comes back in a single request,
275
+ # in the same order as the ids we sent.
276
+ states = self.client.task_get_status([t.task_id for t in pending])
277
+ states = cast(list[Task], states)
278
+
279
+ next_pending: list[Task] = []
280
+ completed = 0
281
+ for task, state in zip(pending, states):
282
+ if state == TaskState.Complete:
283
+ output = self.client.task_get_output(task.task_id)
284
+ task.output = cloudpickle.loads(output)
285
+ completed += 1
286
+ yield task
287
+ elif state == TaskState.Undefined:
288
+ raise RuntimeError(
289
+ f"Task {task.task_id} is unknown to the task queue server"
290
+ )
291
+ else:
292
+ next_pending.append(task)
293
+
294
+ pending = next_pending
295
+ if pending and not completed:
296
+ time.sleep(POLL_INTERVAL_S)
297
+
298
+ @typechecked
299
+ def as_completed(
300
+ self, tasks: Iterable[Task], desc: str | None = None, unit: str = "task"
301
+ ) -> Iterable[Task]:
302
+ tasks = list(tasks)
303
+ iterable = self._as_completed(tasks)
304
+ iterable = tqdm(iterable, total=len(tasks), desc=desc, unit=unit)
305
+ return iterable
306
+
307
+ @typechecked
308
+ def wait(
309
+ self, tasks: Iterable[Task], desc: str | None = None, unit: str = "task"
310
+ ) -> None:
311
+ for _ in self.as_completed(tasks, desc, unit):
312
+ pass
313
+
314
+ def num_groups(self):
315
+ return len(self.groups)
316
+
317
+ def num_workers(self, detail: bool = False):
318
+ if detail:
319
+ return {g.name: len(g.workers) for g in self.groups.values()}
320
+ else:
321
+ return sum(len(g.workers) for g in self.groups.values())
322
+
323
+ def _cleanup_all_workers(self):
324
+ try:
325
+ job_ids = get_running_jobids()
326
+ except subprocess.CalledProcessError as cp:
327
+ print(f"Failed to get running slurm job ids: returncode={cp.returncode}")
328
+ if cp.stdout.strip():
329
+ print(cp.stdout)
330
+ if cp.stderr.strip():
331
+ print(cp.stderr)
332
+ job_ids = None
333
+ except Exception:
334
+ self.logger.exception("Failed to get running slurm job ids")
335
+ job_ids = None
336
+
337
+ if job_ids is None:
338
+ return
339
+
340
+ to_cancel_jobids = []
341
+ for group in self.groups.values():
342
+ for worker in group.workers.values():
343
+ if worker.job_id in job_ids:
344
+ to_cancel_jobids.append(worker.job_id)
345
+
346
+ if not to_cancel_jobids:
347
+ return
348
+
349
+ try:
350
+ cancel_jobs(to_cancel_jobids)
351
+ except subprocess.CalledProcessError as cp:
352
+ print(f"Failed to cancel slurm jobs: returncode={cp.returncode}")
353
+ if cp.stdout.strip():
354
+ print(cp.stdout)
355
+ if cp.stderr.strip():
356
+ print(cp.stderr)
357
+ except Exception:
358
+ self.logger.exception("Failed to cancel slurm jobs")
359
+
360
+ def close(self):
361
+ self._cleanup_all_workers()
362
+ for group in self.groups.values():
363
+ group.workers.clear()
364
+
365
+ self.client.close()
366
+
367
+ def stop(self):
368
+ self._cleanup_all_workers()
369
+ for group in self.groups.values():
370
+ group.workers.clear()
371
+
372
+
373
+ def check_for_error(tasks: list[Task], verbose: bool = True) -> list[Task]:
374
+ ret = []
375
+ for task in tasks:
376
+ if isinstance(task.output, RemoteExecutionError):
377
+ ret.append(task)
378
+
379
+ if verbose:
380
+ print(f"task_id={task.task_id}")
381
+ print(f" error={task.output.error}")
382
+ print(f" error_id={task.output.error_id}")
383
+
384
+ return ret
@@ -0,0 +1,168 @@
1
+ """Pilot workers for Slurm pilot."""
2
+
3
+ import os
4
+ import sys
5
+ import json
6
+ import time
7
+ import pickle
8
+ import socket
9
+ import logging
10
+ import importlib
11
+ from pathlib import Path
12
+ from typing import Any
13
+
14
+ import click
15
+ import cloudpickle
16
+ from ds_service_client import Client
17
+
18
+ from .utils import gen_error_id, RemoteExecutionError, LOG_FORMAT, LOG_LEVEL
19
+
20
+ NEXT_TASK_RETRY_TIME_S: float = 0.1
21
+
22
+
23
+ class PilotWorkerProcess:
24
+ def __init__(
25
+ self,
26
+ group: str,
27
+ name: str,
28
+ actor_class_name: str,
29
+ server_address: str,
30
+ work_dir: Path,
31
+ slurm_job_id: int,
32
+ hostname: str,
33
+ pid: int,
34
+ ):
35
+ self.group = group
36
+ self.server_address = server_address
37
+ self.work_dir = work_dir
38
+
39
+ self.worker_id = "%s:%s:%s:%s:%s" % (group, name, slurm_job_id, hostname, pid)
40
+ self.logger = logging.getLogger("worker_process")
41
+ self.client = Client(self.server_address)
42
+
43
+ self.actor_instance: Any | None
44
+ if actor_class_name == "":
45
+ self.actor_instance = None
46
+ else:
47
+ class_name_parts = actor_class_name.split(".")
48
+ module_name = ".".join(class_name_parts[:-1])
49
+ class_name = class_name_parts[-1]
50
+
51
+ module = importlib.import_module(module_name)
52
+ klass = getattr(module, class_name)
53
+ self.actor_instance = klass()
54
+
55
+ def close(self):
56
+ self.client.close()
57
+ if self.actor_instance is not None:
58
+ if hasattr(self.actor_instance, "close"):
59
+ self.actor_instance.close()
60
+ self.actor_instance = None
61
+
62
+ def main(self):
63
+ self.logger.info("Starting worker: %s" % self.worker_id)
64
+
65
+ while True:
66
+ try:
67
+ task = self.client.task_get(self.worker_id, self.group)
68
+ try:
69
+ self.logger.info(
70
+ "task_id=%s: Deserializing function and inputs ...",
71
+ task.task_id,
72
+ )
73
+ function = cloudpickle.loads(task.function)
74
+ if self.actor_instance is not None:
75
+ function = getattr(self.actor_instance, function)
76
+ args, kwargs = cloudpickle.loads(task.input)
77
+
78
+ self.logger.info("task_id=%s: Executing ...", task.task_id)
79
+ retval = function(*args, **kwargs)
80
+
81
+ self.logger.info("task_id=%s: Serializng output ...", task.task_id)
82
+ output = cloudpickle.dumps(retval, protocol=pickle.HIGHEST_PROTOCOL)
83
+
84
+ self.client.task_done(task.task_id, output)
85
+ except Exception as e:
86
+ eid = gen_error_id()
87
+ self.logger.exception(
88
+ "Error executing %s: %s: %s", task.task_id, eid, e
89
+ )
90
+
91
+ retval = RemoteExecutionError(error=str(e), error_id=eid)
92
+ output = cloudpickle.dumps(retval, protocol=pickle.HIGHEST_PROTOCOL)
93
+ self.client.task_done(task.task_id, output)
94
+ except TimeoutError:
95
+ time.sleep(NEXT_TASK_RETRY_TIME_S)
96
+ except Exception:
97
+ self.logger.exception("Unexpected exception")
98
+
99
+
100
+ @click.command()
101
+ @click.option("--group", type=str, required=True, help="Worker group.")
102
+ @click.option("--name", type=str, required=True, help="Worker job name.")
103
+ @click.option(
104
+ "--actor-class-name",
105
+ type=str,
106
+ required=True,
107
+ help="Name for actor class in DS server store.",
108
+ )
109
+ @click.option("--server-address", type=str, required=True, help="Pilot server address.")
110
+ @click.option(
111
+ "--work-dir",
112
+ type=click.Path(exists=True, file_okay=False, dir_okay=True, path_type=Path),
113
+ required=True,
114
+ help="Work directory.",
115
+ )
116
+ @click.option(
117
+ "--python-paths-json",
118
+ type=str,
119
+ required=True,
120
+ help="JSON encoded Python paths.",
121
+ )
122
+ def slurm_pilot_worker(
123
+ group: str,
124
+ name: str,
125
+ actor_class_name: str,
126
+ server_address: str,
127
+ work_dir: Path,
128
+ python_paths_json: str,
129
+ ):
130
+ """Start a slurm pilot worker."""
131
+ slurm_job_id = int(os.environ.get("SLURM_JOB_ID", -1))
132
+ hostname = socket.gethostname()
133
+ pid = os.getpid()
134
+ log_file = work_dir / f"{name}-{slurm_job_id}-{hostname}-{pid}.log"
135
+
136
+ print(f"Redirecting standard output and standard error to {log_file}")
137
+ sys.stdout.flush()
138
+ sys.stderr.flush()
139
+
140
+ with open(log_file, "wt") as fout:
141
+ sys.stdout = fout
142
+ sys.stderr = fout
143
+
144
+ logging.basicConfig(stream=fout, format=LOG_FORMAT, level=LOG_LEVEL)
145
+
146
+ os.environ["PILOT_WORKER_NAME"] = name
147
+ os.environ["PILOT_WORKER_GROUP"] = group
148
+ os.environ["DS_SERVER_ADDRESS"] = server_address
149
+
150
+ python_paths: list[str] = json.loads(python_paths_json)
151
+ for path in python_paths:
152
+ sys.path.insert(0, path)
153
+
154
+ worker = PilotWorkerProcess(
155
+ group=group,
156
+ name=name,
157
+ actor_class_name=actor_class_name,
158
+ server_address=server_address,
159
+ work_dir=work_dir,
160
+ slurm_job_id=slurm_job_id,
161
+ hostname=hostname,
162
+ pid=pid,
163
+ )
164
+
165
+ try:
166
+ worker.main()
167
+ finally:
168
+ worker.close()