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.
- slurm_workflows/__init__.py +1 -0
- slurm_workflows/ds_service.py +44 -0
- slurm_workflows/run_jupyter.py +57 -0
- slurm_workflows/slurm_pilot_executor.py +384 -0
- slurm_workflows/slurm_pilot_worker.py +168 -0
- slurm_workflows/slurm_utils.py +141 -0
- slurm_workflows/templates/__init__.py +143 -0
- slurm_workflows/templates/run_jupyter.jinja +18 -0
- slurm_workflows/templates/slurm_pilot.jinja +25 -0
- slurm_workflows/templates/slurm_utils.jinja +10 -0
- slurm_workflows/utils.py +72 -0
- slurm_workflows-1.0.0.dist-info/METADATA +315 -0
- slurm_workflows-1.0.0.dist-info/RECORD +17 -0
- slurm_workflows-1.0.0.dist-info/WHEEL +5 -0
- slurm_workflows-1.0.0.dist-info/entry_points.txt +3 -0
- slurm_workflows-1.0.0.dist-info/licenses/LICENSE +19 -0
- slurm_workflows-1.0.0.dist-info/top_level.txt +1 -0
|
@@ -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()
|