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
28 changes: 28 additions & 0 deletions go/job.go
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,8 @@ type Job struct {

artifactsSource string
artifactsDestination string
maxRunAttempts int64
currentRunAttempt int64
}

// JobDict is the JSON-friendly public representation of a job.
Expand Down Expand Up @@ -103,6 +105,7 @@ type jobOptions struct {
artifactsSource string
artifactsDest string
maxRuntime int
maxRunAttempts int64
scratchDisks []ScratchDisk
}

Expand Down Expand Up @@ -154,6 +157,9 @@ type JobOptions struct {
ArtifactsDestination string
// MaxRuntime limits the job runtime in seconds.
MaxRuntime int
// MaxRunAttempts is the max number of run attempts for this job.
// 0 means unset (backend default). 1 means a single attempt (no retries).
MaxRunAttempts int64
// ScratchDisks mounts temporary scratch storage for studio-backed jobs.
ScratchDisks []ScratchDisk
}
Expand Down Expand Up @@ -306,6 +312,24 @@ func (j *Job) ArtifactsDestination() string {
return j.artifactsDestination
}

// MaxRunAttempts returns the max number of run attempts for this job.
// 0 means unset.
func (j *Job) MaxRunAttempts() int64 {
if j == nil {
return 0
}
return j.maxRunAttempts
}

// CurrentRunAttempt returns the current run attempt for this job.
// 0 means unset.
func (j *Job) CurrentRunAttempt() int64 {
if j == nil {
return 0
}
return j.currentRunAttempt
}

// GetJob returns an existing job by name or ID.
func GetJob(name string, opts ...JobOptions) (*Job, error) {
resolved := applyJobOptions(opts...)
Expand Down Expand Up @@ -640,6 +664,7 @@ func applyJobOptions(opts ...JobOptions) jobOptions {
resolved.artifactsSource = opts[0].ArtifactsSource
resolved.artifactsDest = opts[0].ArtifactsDestination
resolved.maxRuntime = opts[0].MaxRuntime
resolved.maxRunAttempts = opts[0].MaxRunAttempts
resolved.scratchDisks = opts[0].ScratchDisks
}
return resolved
Expand Down Expand Up @@ -735,6 +760,8 @@ func jobFromModel(model *models.V1Job, opts jobOptions) *Job {
result.studioID = model.Spec.CloudspaceID
result.artifactsSource = model.Spec.ArtifactsSource
result.artifactsDestination = model.Spec.ArtifactsDestination
result.maxRunAttempts = model.Spec.MaxRunAttempts
result.currentRunAttempt = model.Spec.CurrentRunAttempt
}
return result
}
Expand Down Expand Up @@ -767,6 +794,7 @@ func jobSpec(machine, command string, opts jobOptions) *models.V1JobSpec {
ImageSecretRef: opts.imageCredentials,
InstanceName: machine,
PathMappings: jobPathMappings(opts),
MaxRunAttempts: opts.maxRunAttempts,
RequestedRunDurationSeconds: maxRuntime(opts.maxRuntime),
Spot: opts.interruptible,
Volumes: scratchVolumes(opts.scratchDisks),
Expand Down
36 changes: 36 additions & 0 deletions go/job_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,33 @@ func TestJobExposesStartAndStopTimes(t *testing.T) {
"StoppedAt = %v, want zero time", existing.StoppedAt())
}

func TestJobExposesRunAttemptFields(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(map[string]any{
"id": "job-retry",
"name": "train",
"projectId": "project-1",
"state": "pending",
"spec": map[string]any{
"maxRunAttempts": 3,
"currentRunAttempt": 2,
"parentJobId": "job-parent",
},
})
}))
defer server.Close()
t.Setenv("LIGHTNING_CLOUD_URL", server.URL)

existing, err := lit.GetJob("train", lit.JobOptions{Teamspace: mustTeamspace(t, "project-1", "default", "alice")})
require.NoErrorf(t, err,
"GetJob returned error")
assert.Falsef(t, existing.MaxRunAttempts() != 3,
"MaxRunAttempts = %d, want 3", existing.MaxRunAttempts())
assert.Falsef(t, existing.CurrentRunAttempt() != 2,
"CurrentRunAttempt = %d, want 2", existing.CurrentRunAttempt())
}

func TestJobGetWithIDUsesSimpleStruct(t *testing.T) {
j, err := lit.GetJob("train", lit.JobOptions{ID: "job-1", Teamspace: mustTeamspace(t, "project-1", "")})
require.NoErrorf(t, err,
Expand Down Expand Up @@ -288,6 +315,7 @@ func TestJobRunMapsAdvancedV2Options(t *testing.T) {
ImageClusterCredentials bool `json:"imageClusterCredentials"`
ImageSecretRef string `json:"imageSecretRef"`
RequestedRunDurationSeconds string `json:"requestedRunDurationSeconds"`
MaxRunAttempts int64 `json:"maxRunAttempts"`
PathMappings []struct {
ContainerPath string `json:"containerPath"`
ConnectionName string `json:"connectionName"`
Expand All @@ -308,6 +336,8 @@ func TestJobRunMapsAdvancedV2Options(t *testing.T) {
"imageSecretRef = %q, want docker-secret", body.Spec.ImageSecretRef)
assert.Falsef(t, body.Spec.RequestedRunDurationSeconds != "7200",
"requestedRunDurationSeconds = %q, want 7200", body.Spec.RequestedRunDurationSeconds)
assert.Falsef(t, body.Spec.MaxRunAttempts != 3,
"maxRunAttempts = %d, want 3", body.Spec.MaxRunAttempts)
assert.Falsef(t, len(body.Spec.PathMappings) != 2,
"pathMappings length = %d, want 2", len(body.Spec.PathMappings))

Expand All @@ -330,6 +360,9 @@ func TestJobRunMapsAdvancedV2Options(t *testing.T) {
"image": "registry.example/train:latest",
"artifactsSource": "/outputs",
"artifactsDestination": "efs:data:outputs/run-1",
"maxRunAttempts": 3,
"currentRunAttempt": 1,
"parentJobId": "",
},
})
}))
Expand All @@ -349,6 +382,7 @@ func TestJobRunMapsAdvancedV2Options(t *testing.T) {
ArtifactsSource: "/outputs",
ArtifactsDestination: "efs:data:outputs/run-1",
MaxRuntime: 7200,
MaxRunAttempts: 3,
PathMappings: []lit.JobPathMapping{{
ContainerPath: "/data",
ConnectionName: "dataset",
Expand All @@ -364,6 +398,8 @@ func TestJobRunMapsAdvancedV2Options(t *testing.T) {
if got, want := created.ArtifactPath(), "/teamspace/efs_connections/data/outputs/run-1"; got != want {
assert.Fail(t, fmt.Sprintf("created artifact path = %q, want %q", got, want))
}
assert.Falsef(t, created.MaxRunAttempts() != 3 || created.CurrentRunAttempt() != 1,
"unexpected run-attempt fields: %d %d", created.MaxRunAttempts(), created.CurrentRunAttempt())
}

func TestJobRunMapsScratchDisksForStudioJobs(t *testing.T) {
Expand Down
13 changes: 11 additions & 2 deletions python/lightning_sdk/api/job_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
from contextlib import suppress
from datetime import datetime, timezone
from functools import lru_cache
from typing import TYPE_CHECKING, Dict, Iterator, List, Optional, Union
from typing import TYPE_CHECKING, Any, Dict, Iterator, List, Optional, Union
from urllib.error import URLError
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
from urllib.request import urlopen
Expand Down Expand Up @@ -147,6 +147,7 @@ def submit_job(
entrypoint: Optional[str],
path_mappings: Optional[Dict[str, str]],
max_runtime: Optional[int] = None,
max_run_attempts: Optional[int] = None,
reuse_snapshot: bool = True,
scratch_disks: Optional[Dict[str, int]] = None,
placement_group_id: Optional[str] = None,
Expand All @@ -171,6 +172,8 @@ def submit_job(
max_runtime: DWS (Dynamic Workload Scheduler) reservation duration in seconds
(e.g. some top-end GCP GPUs). Has no effect on non-DWS or interruptible
(spot) machines. ``None`` means no reservation is requested.
max_run_attempts: Max number of run attempts for this job. ``None`` or ``0`` means
unset (backend default). ``1`` means a single attempt (no retries).
reuse_snapshot: Whether to reuse the Studio's existing filesystem snapshot.
scratch_disks: Optional mapping of scratch-disk mount paths to their sizes in GiB.
placement_group_id: Optional placement group identifier for colocating the job.
Expand Down Expand Up @@ -203,6 +206,7 @@ def submit_job(
entrypoint=entrypoint,
path_mappings=path_mappings,
max_runtime=max_runtime,
max_run_attempts=max_run_attempts,
reuse_snapshot=reuse_snapshot,
scratch_disks=sanitized_scratch_disks,
placement_group_id=placement_group_id,
Expand All @@ -227,6 +231,7 @@ def _create_job_body(
path_mappings: Optional[Dict[str, str]],
reuse_snapshot: bool,
max_runtime: Optional[int] = None,
max_run_attempts: Optional[int] = None,
machine_image_version: Optional[str] = None,
scratch_disks: Optional[Dict[str, int]] = None,
placement_group_id: Optional[str] = None,
Expand All @@ -250,6 +255,8 @@ def _create_job_body(
max_runtime: DWS (Dynamic Workload Scheduler) reservation duration in seconds
(e.g. some top-end GCP GPUs). Has no effect on non-DWS or interruptible
(spot) machines. ``None`` means no reservation is requested.
max_run_attempts: Max number of run attempts for this job. ``None`` or ``0`` means
unset (backend default). ``1`` means a single attempt (no retries).
machine_image_version: Pinned machine-image version string, or ``None`` for the default.
scratch_disks: Optional mapping of scratch-disk mount paths to their sizes in GiB.
placement_group_id: Optional placement group identifier for colocating the job.
Expand All @@ -269,9 +276,11 @@ def _create_job_body(
path_mappings_list = resolve_path_mappings(mappings=path_mappings or {})

# need to go via kwargs for typing compatibility since autogenerated apis accept None but aren't typed with None
optional_spec_kwargs = {}
optional_spec_kwargs: Dict[str, Any] = {}
if max_runtime:
optional_spec_kwargs["requested_run_duration_seconds"] = str(max_runtime)
if max_run_attempts:
optional_spec_kwargs["max_run_attempts"] = max_run_attempts

# don't do default dicts, as they'll be mutable. Create a fresh one here
scratch_disks = scratch_disks or {}
Expand Down
28 changes: 28 additions & 0 deletions python/lightning_sdk/job.py
Original file line number Diff line number Diff line change
Expand Up @@ -224,6 +224,7 @@ def run(
entrypoint: Optional[str] = None,
path_mappings: Optional[Dict[str, str]] = None,
max_runtime: Optional[int] = None,
max_run_attempts: Optional[int] = None,
reuse_snapshot: bool = True,
scratch_disks: Optional[Dict[str, int]] = None,
placement_group_id: Optional[str] = None,
Expand Down Expand Up @@ -263,6 +264,9 @@ def run(
max_runtime: DWS (Dynamic Workload Scheduler) reservation duration in seconds
(e.g. some top-end GCP GPUs). Has no effect on non-DWS or interruptible
(spot) machines. ``None`` means no reservation is requested.
max_run_attempts: Max number of run attempts for this job. ``None`` or ``0`` means
unset (backend default). ``1`` means a single attempt (no retries).
``N > 1`` allows up to ``N`` attempts.
reuse_snapshot: Whether to reuse a Studio snapshot when multiple jobs for the same Studio are
submitted. Turning this off may result in longer startup times. Defaults to True.
scratch_disks: Optional mapping of scratch-disk mount paths to their sizes in GiB.
Expand All @@ -288,6 +292,8 @@ def run(
raise ValueError("A job needs to run on at least one machine")
if num_machines > 1 and scratch_disks:
raise ValueError("scratch_disks are not supported for multi-machine jobs")
if num_machines > 1 and max_run_attempts:
raise ValueError("max_run_attempts is not supported for multi-machine jobs")

if image is None:
if not isinstance(studio, Studio):
Expand Down Expand Up @@ -365,6 +371,7 @@ def run(
entrypoint=entrypoint,
path_mappings=path_mappings,
max_runtime=max_runtime,
max_run_attempts=max_run_attempts,
reuse_snapshot=reuse_snapshot,
scratch_disks=scratch_disks,
placement_group_id=placement_group_id,
Expand All @@ -388,6 +395,7 @@ def _submit(
entrypoint: Optional[str] = None,
path_mappings: Optional[Dict[str, str]] = None,
max_runtime: Optional[int] = None,
max_run_attempts: Optional[int] = None,
reuse_snapshot: bool = True,
scratch_disks: Optional[Dict[str, int]] = None,
placement_group_id: Optional[str] = None,
Expand All @@ -397,6 +405,8 @@ def _submit(
raise ValueError("A job needs to run on at least one machine")
if num_machines > 1 and scratch_disks:
raise ValueError("scratch_disks are not supported for multi-machine jobs")
if num_machines > 1 and max_run_attempts:
raise ValueError("max_run_attempts is not supported for multi-machine jobs")

if studio is not None:
studio_id = studio._studio.id
Expand Down Expand Up @@ -450,6 +460,9 @@ def _submit(
raise ValueError("scratch_disk path cannot contain '..'")

self._num_machines = num_machines
extra_submit_kwargs: Dict[str, Any] = {}
if num_machines <= 1:
extra_submit_kwargs["max_run_attempts"] = max_run_attempts
submitted = self._job_api.submit_job(
name=self.name,
command=command,
Expand All @@ -469,6 +482,7 @@ def _submit(
placement_group_id=placement_group_id,
num_machines=num_machines,
scratch_disks=scratch_disks,
**extra_submit_kwargs,
)
if num_machines <= 1 and submitted.name != self._name:
warnings.warn(
Expand Down Expand Up @@ -590,6 +604,20 @@ def private_ip_address(self) -> Optional[str]:
def placement_group_id(self) -> Optional[str]:
return self._guaranteed_job.spec.placement_group_id

@property
def max_run_attempts(self) -> Optional[int]:
"""Max number of run attempts for this job, or ``None`` if unset."""
spec = getattr(self._guaranteed_job, "spec", None)
value = getattr(spec, "max_run_attempts", None)
return value or None

@property
def current_run_attempt(self) -> Optional[int]:
"""Current run attempt for this job, or ``None`` if unset."""
spec = getattr(self._guaranteed_job, "spec", None)
value = getattr(spec, "current_run_attempt", None)
return value or None

@property
def rank(self) -> Optional[int]:
if self.is_multi_machine:
Expand Down
10 changes: 10 additions & 0 deletions python/lightning_sdk/pipeline/steps.py
Original file line number Diff line number Diff line change
Expand Up @@ -205,6 +205,7 @@ def __init__(
entrypoint: str = "sh -c",
path_mappings: Optional[Dict[str, str]] = None,
max_runtime: Optional[int] = None,
max_run_attempts: Optional[int] = None,
wait_for: Union[str, List[str], None] = DEFAULT,
reuse_snapshot: bool = True,
scratch_disks: Optional[Dict[str, int]] = None,
Expand All @@ -230,6 +231,9 @@ def __init__(
entrypoint: Container entrypoint. Defaults to ``sh -c``.
path_mappings: Mappings from container paths to data-connection paths.
max_runtime: Maximum runtime in seconds.
max_run_attempts: Max number of run attempts for this job. ``None`` or ``0`` means
unset (backend default). ``1`` means a single attempt (no retries).
``N > 1`` allows up to ``N`` attempts. Only supported for single-machine jobs.
wait_for: Names of steps that must complete before this step starts.
reuse_snapshot: Whether to reuse a studio snapshot across jobs. Defaults to True.
scratch_disks: Extra volumes to mount under ``/teamspace/scratch``.
Expand All @@ -239,6 +243,8 @@ def __init__(
"""
if num_machines < 1:
raise ValueError("A job needs to run on at least one machine")
if num_machines > 1 and max_run_attempts:
raise ValueError("max_run_attempts is not supported for multi-machine jobs")
self.name = name
self.machine = machine or Machine.CPU
self.command = command
Expand All @@ -256,6 +262,7 @@ def __init__(
self.entrypoint = entrypoint
self.path_mappings = path_mappings
self.max_runtime = max_runtime
self.max_run_attempts = max_run_attempts
self.wait_for = wait_for
self.reuse_snapshot = reuse_snapshot
self.scratch_disks = scratch_disks
Expand Down Expand Up @@ -297,6 +304,8 @@ def to_proto(
if self.num_machines > 1:
if self.scratch_disks:
raise ValueError("scratch_disks are not supported for multi-machine jobs")
if self.max_run_attempts:
raise ValueError("max_run_attempts is not supported for multi-machine jobs")
body = MMTApiV2._create_mmt_body(
name=cast(str, self.name),
num_machines=self.num_machines,
Expand Down Expand Up @@ -337,6 +346,7 @@ def to_proto(
entrypoint=self.entrypoint,
path_mappings=self.path_mappings,
max_runtime=self.max_runtime,
max_run_attempts=self.max_run_attempts,
machine_image_version=machine_image_version,
reuse_snapshot=self.reuse_snapshot,
scratch_disks=self.scratch_disks,
Expand Down
4 changes: 4 additions & 0 deletions python/lightning_sdk/studio.py
Original file line number Diff line number Diff line change
Expand Up @@ -781,6 +781,7 @@ def run_job(
env: Optional[Dict[str, str]] = None,
interruptible: bool = False,
reuse_snapshot: bool = True,
max_run_attempts: Optional[int] = None,
) -> "Job":
"""Run async workloads using the compute environment from your studio.

Expand All @@ -792,6 +793,8 @@ def run_job(
interruptible: Whether the job should run on interruptible instances. They are cheaper but can be preempted.
reuse_snapshot: Whether the job should reuse a Studio snapshot when multiple jobs for the same Studio are
submitted. Turning this off may result in longer job startup times. Defaults to True.
max_run_attempts: Max number of run attempts for this job. ``None`` or ``0`` means
unset (backend default). ``1`` means a single attempt (no retries).

Returns:
Job: The submitted :class:`Job` instance.
Expand All @@ -809,6 +812,7 @@ def run_job(
env=env,
interruptible=interruptible,
reuse_snapshot=reuse_snapshot,
max_run_attempts=max_run_attempts,
)

def run_mmt(
Expand Down
Loading
Loading