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
1 change: 1 addition & 0 deletions packages/data-designer-slurm/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@ bump = true
[tool.hatch.metadata.hooks.uv-dynamic-versioning]
dependencies = [
"data-designer=={{ version }}",
"packaging>=25,<27",
"pydantic>=2.9.2,<3",
]

Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Immutable benchmark records for Data Designer Slurm."""

from __future__ import annotations

from data_designer.slurm.benchmark.records import (
BenchmarkCaseResult,
BenchmarkChildRun,
BenchmarkManifest,
BenchmarkOutcome,
BenchmarkRecommendation,
BenchmarkRecommendationKind,
BenchmarkReport,
)

__all__ = [
"BenchmarkCaseResult",
"BenchmarkChildRun",
"BenchmarkManifest",
"BenchmarkOutcome",
"BenchmarkRecommendation",
"BenchmarkRecommendationKind",
"BenchmarkReport",
]
Original file line number Diff line number Diff line change
@@ -0,0 +1,162 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

from __future__ import annotations

from datetime import datetime, timedelta
from enum import Enum
from typing import Annotated

from pydantic import (
Field,
NonNegativeFloat,
NonNegativeInt,
PositiveInt,
StringConstraints,
field_validator,
model_validator,
)

from data_designer.slurm.contracts import ArtifactReference, ContractRecord, ContractValue, Identifier


class BenchmarkChildRun(ContractValue):
case_id: Identifier
child_run_id: Identifier
child_authored_config: ArtifactReference

@model_validator(mode="after")
def validate_authored_config(self) -> BenchmarkChildRun:
expected_suffix = f"/runs/{self.child_run_id}/authored-config.json"
if not self.child_authored_config.path.endswith(expected_suffix):
raise ValueError("child authored config path must match the child run identity")
return self


class BenchmarkManifest(ContractRecord):
"""Stable mapping from benchmark cases to ordinary child runs."""

benchmark_id: Identifier
benchmark_config: ArtifactReference
children: tuple[BenchmarkChildRun, ...] = Field(min_length=1)

@model_validator(mode="after")
def validate_children(self) -> BenchmarkManifest:
case_ids = tuple(child.case_id for child in self.children)
run_ids = tuple(child.child_run_id for child in self.children)
if len(case_ids) != len(set(case_ids)):
raise ValueError("benchmark case IDs must be unique")
if len(run_ids) != len(set(run_ids)):
raise ValueError("benchmark child run IDs must be unique")
return self


class BenchmarkOutcome(str, Enum):
PENDING = "pending"
ACCOUNTING_LAG = "accounting_lag"
SUCCEEDED = "succeeded"
FAILED = "failed"
INCOMPLETE = "incomplete"


class BenchmarkCaseResult(ContractValue):
case_id: Identifier
child_run_id: Identifier
outcome: BenchmarkOutcome
topology_digest: Annotated[str, StringConstraints(pattern=r"^[0-9a-f]{64}$")]
requested_records: PositiveInt
actual_records: NonNegativeInt | None = None
boot_seconds: NonNegativeFloat | None = None
generation_seconds: NonNegativeFloat | None = None
wall_seconds: NonNegativeFloat | None = None
rows_per_second: NonNegativeFloat | None = None
request_count: NonNegativeInt | None = None
token_count: NonNegativeInt | None = None
gpus_per_job: PositiveInt
nodes_per_job: PositiveInt
gpu_hours_per_job: NonNegativeFloat | None = None
total_gpu_hours: NonNegativeFloat | None = None
target_jobs: PositiveInt | None = None
feasible: bool | None = None

@model_validator(mode="after")
def validate_metrics(self) -> BenchmarkCaseResult:
if self.actual_records is not None and self.actual_records > self.requested_records:
raise ValueError("benchmark actual_records must not exceed requested_records")
required = (
self.actual_records,
self.boot_seconds,
self.generation_seconds,
self.wall_seconds,
self.rows_per_second,
self.gpu_hours_per_job,
self.total_gpu_hours,
self.target_jobs,
self.feasible,
)
if self.outcome is BenchmarkOutcome.SUCCEEDED and any(value is None for value in required):
raise ValueError("successful benchmark cases require complete timing and feasibility metrics")
if self.outcome is BenchmarkOutcome.SUCCEEDED:
if self.actual_records != self.requested_records:
raise ValueError("successful benchmark cases require the requested record count")
if self.generation_seconds == 0 or self.wall_seconds == 0 or self.rows_per_second == 0:
raise ValueError("successful benchmark generation, wall time, and throughput must be positive")
return self


class BenchmarkRecommendationKind(str, Enum):
PARETO = "pareto"
MINIMUM_JOBS = "minimum_jobs"
MINIMUM_GPU_HOURS = "minimum_gpu_hours"


class BenchmarkRecommendation(ContractValue):
kind: BenchmarkRecommendationKind
case_id: Identifier


class BenchmarkReport(ContractRecord):
"""Atomic point-in-time benchmark analysis output."""

benchmark_id: Identifier
analysis_id: Identifier
benchmark_manifest: ArtifactReference
created_at: datetime
cases: tuple[BenchmarkCaseResult, ...] = Field(min_length=1)
recommendations: tuple[BenchmarkRecommendation, ...] = ()

@field_validator("created_at")
@classmethod
def validate_created_at(cls, value: datetime) -> datetime:
if value.tzinfo is None or value.utcoffset() != timedelta(0):
raise ValueError("created_at must be timezone-aware UTC")
return value

@model_validator(mode="after")
def validate_report(self) -> BenchmarkReport:
case_ids = tuple(case.case_id for case in self.cases)
child_run_ids = tuple(case.child_run_id for case in self.cases)
if len(case_ids) != len(set(case_ids)):
raise ValueError("benchmark report case IDs must be unique")
if len(child_run_ids) != len(set(child_run_ids)):
raise ValueError("benchmark report child run IDs must be unique")
unknown = {recommendation.case_id for recommendation in self.recommendations}.difference(case_ids)
if unknown:
raise ValueError(f"recommendations reference unknown cases: {', '.join(sorted(unknown))}")
recommendable = {
case.case_id for case in self.cases if case.outcome is BenchmarkOutcome.SUCCEEDED and case.feasible is True
}
identities: set[tuple[BenchmarkRecommendationKind, str]] = set()
singleton_kinds: set[BenchmarkRecommendationKind] = set()
for recommendation in self.recommendations:
if recommendation.case_id not in recommendable:
raise ValueError("benchmark recommendations must reference successful feasible cases")
identity = (recommendation.kind, recommendation.case_id)
if identity in identities:
raise ValueError("benchmark recommendations must be unique")
identities.add(identity)
if recommendation.kind is not BenchmarkRecommendationKind.PARETO:
if recommendation.kind in singleton_kinds:
raise ValueError("minimum benchmark recommendation kinds must be unique")
singleton_kinds.add(recommendation.kind)
return self
Original file line number Diff line number Diff line change
@@ -0,0 +1,10 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Semantic client records shared with Slurm state consumers."""

from __future__ import annotations

from data_designer.slurm.client.records import ClientOutcome, ClientResult

__all__ = ["ClientOutcome", "ClientResult"]
Original file line number Diff line number Diff line change
@@ -0,0 +1,108 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

from __future__ import annotations

from datetime import datetime, timedelta
from enum import Enum
from typing import Annotated, Literal

from pydantic import NonNegativeInt, PositiveInt, StringConstraints, field_validator, model_validator

from data_designer.slurm.contracts import (
ArtifactReference,
AttemptId,
ContractRecord,
Identifier,
ShardId,
validate_absolute_path,
)


class ClientOutcome(str, Enum):
COMPLETE = "complete"
PARTIAL = "partial"
FAILED = "failed"


class ClientResult(ContractRecord):
"""Semantic Data Designer outcome independent of engine-internal result types."""

run_id: Identifier
shard_id: ShardId
attempt_id: AttemptId
completed_at: datetime
requested_records: PositiveInt
actual_records: NonNegativeInt | None
outcome: ClientOutcome
dataset_path: str | None = None
early_shutdown: bool | None = None
requested_resume_mode: Literal["never", "always", "if_possible"]
effective_resume_mode: Literal["never", "always"] | None = None
candidate_output_manifest: ArtifactReference | None = None
error_code: Identifier | None = None
redacted_message: Annotated[str, StringConstraints(max_length=512)] | None = None

@field_validator("completed_at")
@classmethod
def validate_completed_at(cls, value: datetime) -> datetime:
if value.tzinfo is None or value.utcoffset() != timedelta(0):
raise ValueError("completed_at must be timezone-aware UTC")
return value

@field_validator("dataset_path")
@classmethod
def validate_dataset_path(cls, value: str | None) -> str | None:
return None if value is None else validate_absolute_path(value)

@field_validator("redacted_message")
@classmethod
def validate_message(cls, value: str | None) -> str | None:
if value is not None and any(ord(character) < 32 or ord(character) == 127 for character in value):
raise ValueError("redacted_message must not contain control characters")
return value

@model_validator(mode="after")
def validate_outcome(self) -> ClientResult:
if self.actual_records is not None and self.actual_records > self.requested_records:
raise ValueError("actual_records must not exceed requested_records")
if self.requested_resume_mode != "if_possible" and self.effective_resume_mode not in {
None,
self.requested_resume_mode,
}:
raise ValueError("effective resume mode must match a fixed requested mode")
if self.outcome is not ClientOutcome.FAILED:
if self.early_shutdown is None or self.effective_resume_mode is None:
raise ValueError("non-failed client results require resume and early-shutdown facts")
if self.outcome is ClientOutcome.COMPLETE:
if self.actual_records != self.requested_records:
raise ValueError("complete client results require the requested record count")
if self.early_shutdown:
raise ValueError("complete client results cannot report early shutdown")
self._require_success_artifacts()
elif self.outcome is ClientOutcome.PARTIAL:
if self.actual_records is None or self.actual_records >= self.requested_records:
raise ValueError("partial client results require fewer than the requested record count")
self._require_success_artifacts()
else:
if self.candidate_output_manifest is not None:
raise ValueError("failed client results cannot reference a candidate output manifest")
if self.error_code is None:
raise ValueError("failed client results require error_code")
return self

def _require_success_artifacts(self) -> None:
if self.dataset_path is None or self.candidate_output_manifest is None:
raise ValueError("successful client results require dataset and candidate manifest paths")
if self.error_code is not None or self.redacted_message is not None:
raise ValueError("successful client results cannot contain failure details")
shard_root = f"/runs/{self.run_id}/shards/{self.shard_id}"
if self.effective_resume_mode == "never":
expected_dataset = f"{shard_root}/attempts/{self.attempt_id}/dataset"
else:
expected_dataset = f"{shard_root}/dataset"
if not self.dataset_path.endswith(expected_dataset):
raise ValueError("dataset path must match the run, shard, attempt, and resume policy")
expected_manifest = f"{shard_root}/attempts/{self.attempt_id}/output-manifest.json"
if not self.candidate_output_manifest.path.endswith(expected_manifest):
raise ValueError("candidate output reference must match the run, shard, and attempt")
Loading
Loading