Skip to content
Draft
Show file tree
Hide file tree
Changes from 2 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
56 changes: 39 additions & 17 deletions src/together/lib/cli/api/beta/clusters/create.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,11 +22,11 @@
Optional[Literal["RESERVED", "ON_DEMAND", "SCHEDULED_CAPACITY"]],
Parameter(help="Billing type to use for the cluster"),
]
NvidiaDriverVersionParameter = Annotated[Optional[str], Parameter(help="Nvidia driver version to use for the cluster")]
CudaVersionParameter = Annotated[Optional[str], Parameter(help="CUDA version to use for the cluster")]
DriverParameter = Annotated[Optional[str], Parameter(help="Legacy NVIDIA driver selector; pair with --cuda-version")]
CudaVersionParameter = Annotated[Optional[str], Parameter(help="Legacy CUDA selector; pair with --driver")]
OSParameter = Annotated[Optional[str], Parameter(help="Operating system for NVIDIA version selection")]
NvidiaVersionIDParameter = Annotated[
Optional[str], Parameter(help="NVIDIA version catalog ID to use directly for the cluster")
Optional[str], Parameter(help="Canonical NVIDIA version catalog ID to use for the cluster")
]
DurationDaysParameter = Annotated[
Optional[int], Parameter(help="Duration in days to keep the cluster running for reserved clusters")
Expand Down Expand Up @@ -151,24 +151,46 @@ async def _set_nvidia_version_params(
os_name: str | None,
) -> None:
semantic_version_given = any(value is not None for value in (nvidia_driver_version, cuda_version, os_name))
if nvidia_version_id and semantic_version_given:
raise TogetherError("Use either --nvidia-version-id or --nvidia-driver-version/--cuda-version/--os, not both.")

has_driver = nvidia_driver_version is not None
has_cuda = cuda_version is not None
if not nvidia_version_id and (has_driver != has_cuda or (os_name is not None and not has_driver)):
raise TogetherError("--nvidia-driver-version and --cuda-version must be provided together; --os requires both.")
if has_driver != has_cuda or (os_name is not None and not has_driver):
raise TogetherError("--driver and --cuda-version must be provided together; --os requires both.")

if nvidia_version_id:
params["nvidia_version_id"] = nvidia_version_id
params.pop("nvidia_driver_version", None)
params.pop("cuda_version", None)
return
if not semantic_version_given:
params.pop("nvidia_driver_version", None)
params.pop("cuda_version", None)
return

if not interactive and not has_driver:
raise TogetherError(
"Use --nvidia-version-id or provide --nvidia-driver-version and --cuda-version in non-interactive mode."
if os_name is None:
return

region = params.get("region")
if not region:
raise TogetherError("--region is required when selecting an NVIDIA version.")

if catalog is None:
catalog = await config.client.beta.clusters.list_regions()
selected = _resolve_nvidia_version(
catalog,
region=region,
nvidia_driver_version=nvidia_driver_version,
cuda_version=cuda_version,
os_name=os_name,
)
if selected.id and selected.id != nvidia_version_id:
raise TogetherError(
f"--nvidia-version-id {nvidia_version_id!r} does not match "
f"--driver/--cuda-version/--os selection {selected.id!r}."
)
return

if not has_driver:
if not interactive:
params.pop("nvidia_driver_version", None)
params.pop("cuda_version", None)
return

if has_driver and os_name is None:
return
Expand Down Expand Up @@ -204,7 +226,7 @@ async def create(
num_gpus: NumGpusParameter = None,
region: RegionParameter = None,
billing_type: BillingTypeParameter = None,
nvidia_driver_version: NvidiaDriverVersionParameter = None,
driver: DriverParameter = None,
cuda_version: CudaVersionParameter = None,
os: OSParameter = None,
nvidia_version_id: NvidiaVersionIDParameter = None,
Expand Down Expand Up @@ -234,7 +256,7 @@ async def create(
num_gpus=num_gpus,
region=region,
billing_type=billing_type,
nvidia_driver_version=nvidia_driver_version,
nvidia_driver_version=driver,
cuda_version=cuda_version,
duration_days=duration_days,
gpu_type=gpu_type,
Expand Down Expand Up @@ -318,7 +340,7 @@ async def create(
catalog=catalog,
interactive=interactive,
nvidia_version_id=nvidia_version_id,
nvidia_driver_version=nvidia_driver_version,
nvidia_driver_version=driver,
cuda_version=cuda_version,
os_name=os,
)
Expand Down
5 changes: 2 additions & 3 deletions src/together/lib/cli/utils/_help_examples.py
Original file line number Diff line number Diff line change
Expand Up @@ -520,7 +520,7 @@
[primary]tg beta clusters create --non-interactive \\
--name my-cluster --cluster-type KUBERNETES --gpu-type H100_SXM \\
--region us-central-8 --num-gpus 8 --billing-type ON_DEMAND \\
--nvidia-driver-version 565 --cuda-version 12.6 --volume <volume-id>[/primary]
--nvidia-version-id nvidia-565-22 --volume <volume-id>[/primary]

[dim]-[/dim] Update or delete a cluster:
[primary]tg beta clusters update <cluster-id> --num-gpus 16 --cluster-type KUBERNETES[/primary]
Expand All @@ -546,8 +546,7 @@
--region us-central-8 \\
--num-gpus 8 \\
--billing-type ON_DEMAND \\
--nvidia-driver-version 565 \\
--cuda-version 12.6 \\
--nvidia-version-id nvidia-565-22 \\
--volume <volume-id>[/primary]
"""

Expand Down
76 changes: 64 additions & 12 deletions tests/cli/test_beta_clusters.py
Original file line number Diff line number Diff line change
Expand Up @@ -666,10 +666,9 @@ class TestBetaClustersNvidiaVersionSelection:
[
("595", None, None, "must be provided together"),
(None, None, "ubuntu-24.04", "--os requires both"),
(None, None, None, "Use --nvidia-version-id"),
],
)
async def test_non_interactive_selection_requires_complete_selector(
async def test_non_interactive_selection_rejects_incomplete_selector(
self,
nvidia_driver_version: str | None,
cuda_version: str | None,
Expand All @@ -688,6 +687,25 @@ async def test_non_interactive_selection_requires_complete_selector(
os_name=os_name,
)

@pytest.mark.asyncio
async def test_non_interactive_selection_can_omit_selector(self) -> None:
params: dict[str, Any] = {}

await create_cli._set_nvidia_version_params(
config=cast(Any, None),
params=params,
catalog=None,
interactive=False,
nvidia_version_id=None,
nvidia_driver_version=None,
cuda_version=None,
os_name=None,
)

assert "nvidia_version_id" not in params
assert "nvidia_driver_version" not in params
assert "cuda_version" not in params

@pytest.mark.asyncio
async def test_interactive_explicit_legacy_pair_passes_through(self) -> None:
params: dict[str, Any] = {
Expand Down Expand Up @@ -861,7 +879,7 @@ def test_retrieve_json(self, respx_mock: MockRouter, cli_runner: CliRunner) -> N


class TestBetaClustersCreate:
def test_invalid_nvidia_selector_is_json_in_json_mode(self, cli_runner: CliRunner) -> None:
def test_incomplete_nvidia_selector_is_json_in_json_mode(self, cli_runner: CliRunner) -> None:
result = cli_runner.invoke(
[
"beta",
Expand All @@ -872,12 +890,8 @@ def test_invalid_nvidia_selector_is_json_in_json_mode(self, cli_runner: CliRunne
"KUBERNETES",
"--gpu-type",
"H100_SXM",
"--nvidia-version-id",
"nvidia-595-24",
"--nvidia-driver-version",
"--driver",
"595",
"--cuda-version",
"13.2",
"--region",
"us-central-8",
"--num-gpus",
Expand All @@ -891,9 +905,42 @@ def test_invalid_nvidia_selector_is_json_in_json_mode(self, cli_runner: CliRunne

assert result.exit_code == 1
assert json.loads(result.output) == {
"error": "Use either --nvidia-version-id or --nvidia-driver-version/--cuda-version/--os, not both."
"error": "--driver and --cuda-version must be provided together; --os requires both."
}

@pytest.mark.respx(base_url=base_url)
def test_create_non_interactive_can_omit_nvidia_selector(
self, respx_mock: MockRouter, cli_runner: CliRunner
) -> None:
created = _cluster_body("new-id", "default-selector")
route = respx_mock.post("/compute/clusters").mock(return_value=httpx.Response(200, json=created))
result = cli_runner.invoke(
[
"beta",
"clusters",
"create",
"--non-interactive",
"--cluster-type",
"KUBERNETES",
"--gpu-type",
"H100_SXM",
"--region",
"us-central-8",
"--num-gpus",
"8",
"--billing-type",
"ON_DEMAND",
"--name",
"default-selector",
],
)

assert result.exit_code == 0, result.output
body = json.loads(cast(Call, route.calls[0]).request.content.decode())
assert "nvidia_version_id" not in body
assert "nvidia_driver_version" not in body
assert "cuda_version" not in body

@pytest.mark.respx(base_url=base_url)
def test_create_non_interactive_posts_expected_body(self, respx_mock: MockRouter, cli_runner: CliRunner) -> None:
created = _cluster_body("new-id", "together-py-testing-suite")
Expand All @@ -908,7 +955,7 @@ def test_create_non_interactive_posts_expected_body(self, respx_mock: MockRouter
"KUBERNETES",
"--gpu-type",
"H100_SXM",
"--nvidia-driver-version",
"--driver",
"565",
"--cuda-version",
"12.6",
Expand Down Expand Up @@ -987,7 +1034,7 @@ def test_create_semantic_nvidia_selection_posts_resolved_id(
"KUBERNETES",
"--gpu-type",
"H100_SXM",
"--nvidia-driver-version",
"--driver",
"595",
"--cuda-version",
"13.2",
Expand Down Expand Up @@ -1024,10 +1071,12 @@ def test_create_accepts_new_cluster_params(self, respx_mock: MockRouter, cli_run
"SLURM",
"--gpu-type",
"H100_SXM",
"--nvidia-driver-version",
"--driver",
"565",
"--cuda-version",
"12.6",
"--nvidia-version-id",
"nvidia-565-22",
"--region",
"us-central-8",
"--num-gpus",
Expand Down Expand Up @@ -1091,6 +1140,9 @@ def test_create_accepts_new_cluster_params(self, respx_mock: MockRouter, cli_run
assert body["reservation_end_time"] == "2026-06-02T00:00:00Z"
assert body["slurm_image"] == "slurm:latest"
assert body["slurm_shm_size_gib"] == 32
assert body["nvidia_driver_version"] == "565"
assert body["cuda_version"] == "12.6"
assert body["nvidia_version_id"] == "nvidia-565-22"
assert result.exit_code == 0


Expand Down
2 changes: 1 addition & 1 deletion tests/cli/test_json_mode_pipeable_to_jq.py
Original file line number Diff line number Diff line change
Expand Up @@ -140,7 +140,7 @@ def test_beta_clusters_json_mode(self) -> None:
beta_clusters = JSONValidator(("beta", "clusters"))
beta_clusters.run_and_assert(
"create --non-interactive --cluster-type KUBERNETES --gpu-type H100_SXM "
"--nvidia-driver-version 565 --cuda-version 12.6 --region us-central-8 --num-gpus 8 "
"--driver 565 --cuda-version 12.6 --region us-central-8 --num-gpus 8 "
"--billing-type ON_DEMAND --name together-py-testing-suite --volume 123"
)
beta_clusters.run_and_assert("delete cluster-123")
Expand Down
Loading