"""Backend that submits workers as a Slurm job array."""
from __future__ import annotations
import logging
import shlex
import shutil
import subprocess
import time
from datetime import datetime
from pathlib import Path
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from .protocol import DispatchMode
from stilt.project import project_slug
from stilt.store import is_uri
logger = logging.getLogger(__name__)
# Number of tries for a scheduler query that times out. A busy controller can
# be slow to answer squeue or sacct, and one slow answer should not end the wait.
_POLL_RETRIES = 5
def _run_scheduler_query(
cmd: list[str], *, timeout: int = 30
) -> subprocess.CompletedProcess:
"""Run a Slurm query command, retrying when it times out."""
last_exc: subprocess.TimeoutExpired | None = None
for _ in range(_POLL_RETRIES):
try:
return subprocess.run(cmd, capture_output=True, text=True, timeout=timeout)
except subprocess.TimeoutExpired as exc:
last_exc = exc
time.sleep(5)
assert last_exc is not None
raise last_exc
def _write_chunks(
chunk_dir: Path,
sim_ids: list[str],
*,
n_workers: int,
) -> int:
"""
Split receptor ids round-robin into one chunk file per array task.
Returns the number of chunk files written.
"""
if not sim_ids:
return 0
chunk_dir.mkdir(parents=True, exist_ok=True)
n_chunks = max(1, min(n_workers, len(sim_ids)))
for idx in range(n_chunks):
chunk = sim_ids[idx::n_chunks]
(chunk_dir / f"task_{idx}.txt").write_text(
"\n".join(chunk) + "\n", encoding="utf-8"
)
return n_chunks
class SlurmHandle:
"""Handle to a Slurm job array submitted with ``sbatch``."""
def __init__(
self,
job_id: str,
*,
chunk_dir: Path | None = None,
) -> None:
self._job_id = job_id
self._chunk_dir = chunk_dir
self._completed = False
@property
def job_id(self) -> str:
"""Job id reported by ``sbatch``."""
return self._job_id
@property
def detached(self) -> bool:
"""Always True, since Slurm jobs run on after this process exits."""
return True
def wait(self) -> None:
"""
Block until the job leaves the Slurm queue.
Polls ``squeue`` every 30 s, then checks the final state with
``sacct`` and raises ``RuntimeError`` if any task failed, was
cancelled, or timed out. The chunk files are deleted once the job has
left the queue.
"""
if self._completed:
return
# Chunk files are deleted only after the job has left the queue,
# whether or not it succeeded. If wait() exits early (a query that
# keeps timing out, an interrupt), tasks still to run need them.
job_left_queue = False
try:
while True:
result = _run_scheduler_query(
["squeue", "--job", self._job_id, "--noheader"]
)
if result.returncode != 0:
raise RuntimeError(result.stderr.strip() or "squeue failed")
if not result.stdout.strip():
job_left_queue = True
break
time.sleep(30)
status = _run_scheduler_query(
[
"sacct",
"--jobs",
self._job_id,
"--noheader",
"--parsable2",
"--format=State",
]
)
if status.returncode != 0:
raise RuntimeError(status.stderr.strip() or "sacct failed")
states = {
line.strip().split("|")[0]
for line in status.stdout.splitlines()
if line.strip()
}
if any(
state.startswith(prefix)
for state in states
for prefix in ("FAILED", "CANCELLED", "TIMEOUT")
):
raise RuntimeError(
f"Slurm job {self._job_id} finished unsuccessfully: {sorted(states)}"
)
self._completed = True
finally:
# Once the job has left the queue no task will read the chunks
# again. Otherwise they are left in place.
if job_left_queue and self._chunk_dir is not None:
shutil.rmtree(self._chunk_dir, ignore_errors=True)
[docs]
class SlurmExecutor:
"""
Run receptors as a Slurm job array submitted with ``sbatch``.
:meth:`start` splits the receptor ids into one chunk file per array task
under ``<project>/chunks/<batch>/``, writes a submission script under
``<project>/slurm/``, and submits it. Each task runs
``stilt push-worker`` on its chunk. The project must be a local
directory.
Parameters
----------
n_workers : int
Number of array tasks.
cpus_per_task : int, default 1
CPUs per array task. With more than one, each task runs its
receptors in a process pool of that size.
array_parallelism : int, optional
Maximum number of array tasks running at once (the ``%N`` suffix
of ``--array``).
setup : list of str, optional
Shell commands to run before the worker, such as loading modules.
**kwargs
Other ``sbatch`` options, written as ``#SBATCH --key=value``.
Underscores in keys become hyphens, and ``True`` writes a bare flag.
"""
dispatch: DispatchMode = "push"
def __init__(
self,
n_workers: int,
cpus_per_task: int = 1,
array_parallelism: int | None = None,
setup: list[str] | None = None,
**kwargs: Any,
) -> None:
self._n_workers = n_workers
self._cpus_per_task = cpus_per_task
self._array_parallelism = array_parallelism
self._setup: list[str] = setup or []
self._kwargs = kwargs
@property
def n_workers(self) -> int:
"""Number of array tasks."""
return self._n_workers
[docs]
@classmethod
def from_config(cls, config: dict[str, Any]) -> SlurmExecutor:
"""Return an executor for a config's ``execution`` settings, which must set ``n_workers``."""
cfg = dict(config)
cfg.pop("backend", None)
n_workers = cfg.pop("n_workers", None)
if n_workers is None:
raise ValueError(
"SlurmExecutor requires explicit 'n_workers' in execution config."
)
cpus_per_task = cfg.pop("cpus_per_task", cfg.pop("cpus-per-task", 1))
array_parallelism = cfg.pop("array_parallelism", None)
setup = cfg.pop("setup", None)
if isinstance(setup, str):
setup = [setup]
return cls(
n_workers=n_workers,
cpus_per_task=cpus_per_task,
array_parallelism=array_parallelism,
setup=setup,
**cfg,
)
def _resolved_slurm_kwargs(self, project: str) -> dict[str, Any]:
"""Return the ``sbatch`` options, with a default job name."""
kwargs = dict(self._kwargs)
kwargs.setdefault("job_name", f"pystilt-{project_slug(project)}")
return kwargs
def _render_sbatch_directives(self, n_workers: int, *, project: str) -> str:
"""Return the ``#SBATCH`` lines of a submission script."""
lines: list[str] = []
array_spec = f"0-{n_workers - 1}"
if self._array_parallelism is not None:
array_spec += f"%{self._array_parallelism}"
lines.append(f"#SBATCH --array={array_spec}")
if self._cpus_per_task > 1:
lines.append(f"#SBATCH --cpus-per-task={self._cpus_per_task}")
for key, value in self._resolved_slurm_kwargs(project).items():
flag = key.replace("_", "-")
if isinstance(value, bool):
if value:
lines.append(f"#SBATCH --{flag}")
else:
lines.append(f"#SBATCH --{flag}={value}")
return "\n".join(lines)
[docs]
def start(
self,
pending: list[str],
*,
project: str,
compute_root: str | None = None,
skip_existing: bool | None = None,
) -> SlurmHandle:
"""Write the chunk files and submission script, submit it, and return a handle."""
if is_uri(project):
raise ValueError("Slurm push dispatch requires a local project root.")
project_dir = Path(project)
batch_id = datetime.now().strftime("%Y%m%d_%H%M%S")
chunk_dir = project_dir / "chunks" / batch_id
n_written = _write_chunks(chunk_dir, pending, n_workers=self._n_workers)
if not n_written:
return SlurmHandle("none")
slurm_dir = project_dir / "slurm"
# One directory per submission, so a later array never overwrites these.
logs_dir = slurm_dir / "logs" / batch_id
logs_dir.mkdir(parents=True, exist_ok=True)
script_path = slurm_dir / f"submit_{batch_id}.sh"
directives = self._render_sbatch_directives(n_written, project=project)
cpus_flag = f" --cpus {self._cpus_per_task}" if self._cpus_per_task > 1 else ""
compute_flag = (
f" --compute-root {shlex.quote(compute_root)}"
if compute_root is not None
else ""
)
skip_flag = " --no-skip" if skip_existing is False else ""
script_lines = [
"#!/bin/bash",
directives,
f"#SBATCH --output={logs_dir}/%a.out",
f"#SBATCH --error={logs_dir}/%a.err",
"",
*self._setup,
*([""] if self._setup else []),
f"CHUNK_PATH={shlex.quote(str(chunk_dir))}/task_${{SLURM_ARRAY_TASK_ID}}.txt",
(
f"stilt push-worker {shlex.quote(project)}"
' --chunk "$CHUNK_PATH"'
f"{cpus_flag}{compute_flag}{skip_flag}"
),
]
script_path.write_text("\n".join(script_lines) + "\n")
script_path.chmod(0o755)
result = subprocess.run(
["sbatch", str(script_path)],
capture_output=True,
text=True,
timeout=60,
)
if result.returncode != 0:
raise RuntimeError(
f"sbatch failed (exit {result.returncode}):\n"
f" script: {script_path}\n"
f" stdout: {result.stdout.strip()}\n"
f" stderr: {result.stderr.strip()}"
)
job_id = result.stdout.strip().split()[-1]
logger.info(f"Submitted job: {job_id}")
return SlurmHandle(job_id, chunk_dir=chunk_dir)