Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
52 changes: 52 additions & 0 deletions packages/data-designer-slurm/src/data_designer/slurm/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,11 +14,13 @@

from data_designer.slurm.config import ImageBuildRequest, SlurmConfigLoadError, load_run_config
from data_designer.slurm.contracts import canonical_json
from data_designer.slurm.images.records import validate_oci_source_for_lifecycle
from data_designer.slurm.services import (
SlurmServiceError,
SlurmServiceErrorCode,
SlurmServiceOperation,
create_slurm_image_service,
create_slurm_profile_service,
create_slurm_run_service,
)

Expand All @@ -37,7 +39,9 @@
no_args_is_help=True,
)
image_app = typer.Typer(help="Manage verified Slurm images", no_args_is_help=True)
profile_app = typer.Typer(help="Initialize and validate Slurm profiles", no_args_is_help=True)
app.add_typer(image_app, name="image")
app.add_typer(profile_app, name="profile")


@app.callback()
Expand Down Expand Up @@ -94,6 +98,46 @@ def cancel_command(
_emit_result(result)


@profile_app.command("init")
def profile_init_command(
workspace_root: Path = typer.Option(..., "--workspace-root", file_okay=False),
image_build_partition: str = typer.Option(..., "--image-build-partition"),
profile_file: Path | None = typer.Option(None, "--profile-file", dir_okay=False),
cluster: str = typer.Option("default", "--cluster"),
account: str | None = typer.Option(None, "--account"),
partition: str | None = typer.Option(None, "--partition"),
host_pattern: list[str] | None = typer.Option(None, "--host-pattern"),
) -> None:
"""Create a safe portable starter profile without overwriting."""
operation = SlurmServiceOperation.INIT_PROFILE
result = _invoke(
operation,
lambda: create_slurm_profile_service(profile_file=profile_file).initialize(
workspace_root=workspace_root,
image_build_partition=image_build_partition,
cluster=cluster,
account=account,
partition=partition,
host_patterns=tuple(host_pattern or ()),
),
)
_emit_result(result)


@profile_app.command("validate")
def profile_validate_command(
profile_file: Path | None = typer.Option(None, "--profile-file", dir_okay=False),
cluster: str | None = typer.Option(None, "--cluster"),
) -> None:
"""Validate strict loading, cluster selection, workspace, and Slurm facts."""
operation = SlurmServiceOperation.VALIDATE_PROFILE
result = _invoke(
operation,
lambda: create_slurm_profile_service(profile_file=profile_file, cluster=cluster).validate(),
)
_emit_result(result)


@image_app.command("add")
def image_add_command(
source: str = typer.Argument(...),
Expand All @@ -113,6 +157,14 @@ def add() -> BaseModel:
operation,
"OCI image source must be digest-qualified as name@sha256:<digest>",
)
try:
validate_oci_source_for_lifecycle(source)
except ValueError:
raise SlurmServiceError(
SlurmServiceErrorCode.INVALID_REQUEST,
operation,
"OCI image source must be a credential-free registry reference without a scheme",
) from None
request = ImageBuildRequest(name=name or _derive_image_name(source), kind=kind, source=source)
return create_slurm_image_service(profile_file=profile_file, cluster=cluster).add(request, replace=replace)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
)
from data_designer.slurm.launcher.parsing import (
parse_accounting,
parse_default_partition_gpu_counts,
parse_gpu_counts,
parse_named_jobs,
parse_queue,
Expand Down Expand Up @@ -224,12 +225,13 @@ def release(self, job_id: int) -> None:
self._run((self._executables.scontrol, "release", _format_job_id(job_id)))

def query_gpu_counts(self, *, partition: Identifier | None = None) -> tuple[int, ...]:
"""Return configured GPU counts reported for eligible node groups."""
command = [self._executables.sinfo, "--noheader", "--format=%G"]
if partition is not None:
if type(partition) is not str or _IDENTIFIER_PATTERN.fullmatch(partition) is None:
raise ValueError("Slurm partition must be a valid identifier")
command.append(f"--partition={partition}")
"""Return configured GPU counts for the requested or default partition."""
if partition is None:
command = (self._executables.sinfo, "--noheader", "--format=%P|%G")
return parse_default_partition_gpu_counts(self._run(command))
if type(partition) is not str or _IDENTIFIER_PATTERN.fullmatch(partition) is None:
raise ValueError("Slurm partition must be a valid identifier")
command = (self._executables.sinfo, "--noheader", "--format=%G", f"--partition={partition}")
return parse_gpu_counts(self._run(command))

def _run(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -146,6 +146,27 @@ def parse_gpu_counts(output: str) -> tuple[int, ...]:
return tuple(counts)


def parse_default_partition_gpu_counts(output: str) -> tuple[int, ...]:
"""Parse GPU counts from ``sinfo --format=%P|%G`` default-partition rows."""
default_partition: str | None = None
resources: list[str] = []
for line_number, line in _collect_nonempty_lines(output):
fields = tuple(field.strip() for field in line.split("|"))
if len(fields) != 2 or not all(fields):
raise SlurmCommandOutputError(f"sinfo line {line_number} must contain a partition and resources")
partition, gres = fields
if not partition.endswith("*"):
continue
partition = partition.removesuffix("*")
if _CLUSTER_NAME_PATTERN.fullmatch(partition) is None:
raise SlurmCommandOutputError(f"sinfo line {line_number} contains an invalid default partition")
if default_partition is not None and partition != default_partition:
raise SlurmCommandOutputError("sinfo returned multiple default partitions")
default_partition = partition
resources.append(gres)
return parse_gpu_counts("\n".join(resources))


def parse_state(value: str) -> SchedulerState:
"""Normalize one Slurm long state spelling without guessing unknown states."""
normalized = value.strip().upper().removesuffix("+")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,13 @@
SlurmServiceOperation,
)
from data_designer.slurm.services.images import SlurmImageManager, SlurmImageResolver, SlurmImageService
from data_designer.slurm.services.profiles import (
SlurmProfileInitialization,
SlurmProfileMatch,
SlurmProfileService,
SlurmProfileValidation,
create_slurm_profile_service,
)
from data_designer.slurm.services.results import (
SlurmPersistedAttemptStatus,
SlurmPersistedRunStatus,
Expand Down Expand Up @@ -47,6 +54,10 @@
"SlurmPersistedAttemptStatus",
"SlurmPersistedRunStatus",
"SlurmPersistedShardStatus",
"SlurmProfileInitialization",
"SlurmProfileMatch",
"SlurmProfileService",
"SlurmProfileValidation",
"SlurmRunArtifactPublisher",
"SlurmRunBackend",
"SlurmRunCancellation",
Expand All @@ -57,6 +68,7 @@
"SlurmServiceErrorCode",
"SlurmServiceOperation",
"create_slurm_image_service",
"create_slurm_profile_service",
"create_slurm_run_service",
]

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,8 @@ class SlurmServiceOperation(str, Enum):
EXECUTE_RUN = "execute_run"
STATUS_RUN = "status_run"
CANCEL_RUN = "cancel_run"
INIT_PROFILE = "init_profile"
VALIDATE_PROFILE = "validate_profile"
RESOLVE_IMAGE = "resolve_image"
ADD_IMAGE = "add_image"
LIST_IMAGES = "list_images"
Expand Down
Loading
Loading