Skip to content
Draft
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
60 changes: 42 additions & 18 deletions src/together/lib/cli/api/beta/clusters/create.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,12 +32,14 @@
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")]
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")
NvidiaDriverVersionParameter = Annotated[
Optional[str], Parameter(help="Legacy NVIDIA driver selector; pair with --cuda-version")
]
CudaVersionParameter = Annotated[
Optional[str], Parameter(help="Legacy CUDA selector; pair with --nvidia-driver-version")
]
OSParameter = Annotated[Optional[str], Parameter(help="Operating system for NVIDIA version selection")]
DriverParameter = Annotated[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 @@ -119,7 +121,7 @@ def _resolve_nvidia_version(
)
if len(matches) > 1:
choices = "; ".join(_format_nvidia_version(version) for version in matches)
guidance = "Use --nvidia-version-id." if os_name else "Add --os or use --nvidia-version-id."
guidance = "Use --driver." if os_name else "Add --os or use --driver."
raise TogetherError(
f"Multiple NVIDIA versions match {requested} in region '{region}'. {guidance} Matches: {choices}"
)
Expand Down Expand Up @@ -158,24 +160,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)):
if 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 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"--driver {nvidia_version_id!r} does not match "
f"--nvidia-driver-version/--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 @@ -210,7 +234,7 @@ async def create(
nvidia_driver_version: NvidiaDriverVersionParameter = None,
cuda_version: CudaVersionParameter = None,
os: OSParameter = None,
nvidia_version_id: NvidiaVersionIDParameter = None,
driver: DriverParameter = None,
duration_days: DurationDaysParameter = None,
gpu_type: GpuTypeParameter = None,
cluster_type: ClusterTypeParameter = None,
Expand Down Expand Up @@ -318,7 +342,7 @@ async def create(
params=params,
catalog=catalog,
interactive=interactive,
nvidia_version_id=nvidia_version_id,
nvidia_version_id=driver,
nvidia_driver_version=nvidia_driver_version,
cuda_version=cuda_version,
os_name=os,
Expand Down
4 changes: 2 additions & 2 deletions src/together/lib/cli/utils/_help_examples.py
Original file line number Diff line number Diff line change
Expand Up @@ -541,7 +541,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-version-id <nvidia-version-id> --volume <volume-id>[/primary]
--driver 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 @@ -567,7 +567,7 @@
--region us-central-8 \\
--num-gpus 8 \\
--billing-type ON_DEMAND \\
--nvidia-version-id <nvidia-version-id> \\
--driver nvidia-565-22 \\
--volume <volume-id>[/primary]
"""

Expand Down
74 changes: 63 additions & 11 deletions tests/cli/test_beta_clusters.py
Original file line number Diff line number Diff line change
Expand Up @@ -675,10 +675,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 @@ -697,6 +696,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 @@ -772,7 +790,7 @@ def test_semantic_selection_uses_os_to_disambiguate_duplicate_cuda_rows(self) ->
assert selected.id == "nvidia-595-22"

def test_semantic_selection_requires_disambiguation_for_duplicate_cuda_rows(self) -> None:
with pytest.raises(TogetherError, match="Add --os or use --nvidia-version-id"):
with pytest.raises(TogetherError, match="Add --os or use --driver"):
create_cli._resolve_nvidia_version(
ClusterListRegionsResponse(**_REGIONS_BODY),
region="us-central-8",
Expand Down Expand Up @@ -810,7 +828,7 @@ def test_semantic_selection_with_duplicate_os_recommends_id_only(self) -> None:
)
)

with pytest.raises(TogetherError, match=r"Use --nvidia-version-id\. Matches"):
with pytest.raises(TogetherError, match=r"Use --driver\. Matches"):
create_cli._resolve_nvidia_version(
catalog,
region="us-central-8",
Expand Down Expand Up @@ -856,7 +874,7 @@ def test_create_help_mentions_b300_gpu_type(self, cli_runner: CliRunner) -> None
assert "B300_SXM" in result.output
assert result.exit_code == 0

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 @@ -867,12 +885,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",
"595",
"--cuda-version",
"13.2",
"--region",
"us-central-8",
"--num-gpus",
Expand All @@ -886,9 +900,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": "--nvidia-driver-version 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 Down Expand Up @@ -945,7 +992,7 @@ def test_create_direct_nvidia_version_id_posts_id(self, respx_mock: MockRouter,
"KUBERNETES",
"--gpu-type",
"H100_SXM",
"--nvidia-version-id",
"--driver",
"nvidia-595-24",
"--region",
"us-central-8",
Expand Down Expand Up @@ -1023,6 +1070,8 @@ def test_create_accepts_new_cluster_params(self, respx_mock: MockRouter, cli_run
"565",
"--cuda-version",
"12.6",
"--driver",
"nvidia-565-22",
"--region",
"us-central-8",
"--num-gpus",
Expand Down Expand Up @@ -1087,6 +1136,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
Loading