Skip to content
Merged
Show file tree
Hide file tree
Changes from 8 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