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
27 changes: 25 additions & 2 deletions src/cloudai/workloads/ai_dynamo/ai_dynamo.py
Original file line number Diff line number Diff line change
Expand Up @@ -116,8 +116,9 @@ class WorkerConfig(BaseModel):

model_config = ConfigDict(extra="forbid", populate_by_name=True)

cmd: str
worker_initialized_regex: str = Field(
cmd: str | None = None
worker_initialized_regex: str | None = Field(
default=None,
validation_alias=AliasChoices("worker-initialized-regex", "worker_initialized_regex"),
serialization_alias="worker-initialized-regex",
)
Expand All @@ -140,6 +141,28 @@ class WorkerConfig(BaseModel):
validation_alias=AliasChoices("extra-args", "extra_args"),
)

@property
def is_enabled(self) -> bool:
"""Return whether any configured trial launches this worker."""
node_counts = self.num_nodes if isinstance(self.num_nodes, list) else [self.num_nodes]
return not node_counts or any(count != 0 for count in node_counts)
Comment thread
coderabbitai[bot] marked this conversation as resolved.

@model_validator(mode="after")
def validate_enabled_worker_fields(self) -> "WorkerConfig":
"""Require launch fields unless every configured worker node count is zero."""
if not self.is_enabled:
return self

missing_fields = []
if not self.cmd:
missing_fields.append("cmd")
if not self.worker_initialized_regex:
missing_fields.append("worker-initialized-regex")
if missing_fields:
raise ValueError(f"{', '.join(missing_fields)} must be set when num-nodes is non-zero")

return self


class DCGMExporter(BaseModel):
"""Optional DCGM exporter launch configuration."""
Expand Down
26 changes: 16 additions & 10 deletions src/cloudai/workloads/ai_dynamo/kubernetes_json_gen_strategy.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,39 +48,45 @@ def gen_frontend_dict(self, cni_networks: list[str] | None = None) -> dict[str,

def gen_decode_dict(self, cni_networks: list[str] | None = None) -> dict[str, Any]:
tdef = cast(AIDynamoTestDefinition, self.test_run.test)
decode_worker = tdef.cmd_args.dynamo.decode_worker
if not decode_worker.is_enabled:
raise ValueError("Decode worker must be enabled for Kubernetes deployments.")
assert decode_worker.cmd is not None

decode_cfg = self._get_base_service_dict(cni_networks)
decode_cfg["extraPodSpec"]["mainContainer"]["command"] = tdef.cmd_args.dynamo.decode_worker.cmd.split()
decode_cfg["extraPodSpec"]["mainContainer"]["command"] = decode_worker.cmd.split()

args = ["--model", tdef.cmd_args.dynamo.model]
if tdef.cmd_args.dynamo.prefill_worker:
if tdef.cmd_args.dynamo.prefill_worker.is_enabled:
Comment thread
jj10306 marked this conversation as resolved.
decode_cfg["subComponentType"] = "decode-worker"
args.append("--is-decode-worker")
args.extend(self._args_from_worker_config(tdef.cmd_args.dynamo.decode_worker))
args.extend(self._args_from_worker_config(decode_worker))

decode_cfg["extraPodSpec"]["mainContainer"]["args"] = args

self._set_multinode_if_needed(decode_cfg, tdef.cmd_args.dynamo.decode_worker)
self._set_multinode_if_needed(decode_cfg, decode_worker)
Comment thread
coderabbitai[bot] marked this conversation as resolved.

return decode_cfg

def gen_prefill_dict(self, cni_networks: list[str] | None = None) -> dict[str, Any]:
tdef = cast(AIDynamoTestDefinition, self.test_run.test)
if not tdef.cmd_args.dynamo.prefill_worker:
raise ValueError("Prefill worker configuration is not defined in the test definition.")
prefill_worker = tdef.cmd_args.dynamo.prefill_worker
if not prefill_worker.is_enabled:
raise ValueError("Prefill worker is disabled in the test definition.")
assert prefill_worker.cmd is not None

prefill_cfg = self._get_base_service_dict(cni_networks)
prefill_cfg["subComponentType"] = "prefill"
prefill_cfg["extraPodSpec"]["mainContainer"]["command"] = tdef.cmd_args.dynamo.prefill_worker.cmd.split()
prefill_cfg["extraPodSpec"]["mainContainer"]["command"] = prefill_worker.cmd.split()

prefill_cfg["extraPodSpec"]["mainContainer"]["args"] = [
"--model",
tdef.cmd_args.dynamo.model,
"--is-prefill-worker",
*self._args_from_worker_config(tdef.cmd_args.dynamo.prefill_worker),
*self._args_from_worker_config(prefill_worker),
]

self._set_multinode_if_needed(prefill_cfg, tdef.cmd_args.dynamo.prefill_worker)
self._set_multinode_if_needed(prefill_cfg, prefill_worker)

return prefill_cfg

Expand All @@ -100,7 +106,7 @@ def gen_json(self) -> Dict[Any, Any]:
},
},
}
if td.cmd_args.dynamo.prefill_worker:
if td.cmd_args.dynamo.prefill_worker.is_enabled:
deployment["spec"]["services"]["prefill"] = self.gen_prefill_dict(cni_networks)

with (self.test_run.output_path / self.DEPLOYMENT_FILE_NAME).open("w") as f:
Expand Down
36 changes: 36 additions & 0 deletions tests/workloads/ai_dynamo/test_command_gen_strategy_slurm.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,29 @@
)


def test_worker_config_allows_omitting_launch_fields_when_num_nodes_is_zero() -> None:
worker = WorkerConfig.model_validate({"num-nodes": 0})

assert worker.num_nodes == 0
assert worker.cmd is None
assert worker.worker_initialized_regex is None


@pytest.mark.parametrize("num_nodes", [1, [0, 1]])
def test_worker_config_requires_launch_fields_when_any_num_nodes_is_nonzero(num_nodes: int | list[int]) -> None:
with pytest.raises(ValueError, match="cmd, worker-initialized-regex must be set when num-nodes is non-zero"):
WorkerConfig.model_validate({"num-nodes": num_nodes})


def test_ai_dynamo_args_accepts_disabled_prefill_worker_without_launch_fields() -> None:
args = AIDynamoArgs.model_validate({"prefill_worker": {"num-nodes": 0}})

assert args.prefill_worker.num_nodes == 0
assert args.prefill_worker.cmd is None
assert args.prefill_worker.worker_initialized_regex is None
assert args.decode_worker.cmd == "python3 -m dynamo.vllm"


@pytest.fixture
def cmd_args() -> AIDynamoCmdArgs:
return AIDynamoCmdArgs(
Expand Down Expand Up @@ -120,6 +143,19 @@ def strategy(slurm_system: SlurmSystem, test_run: TestRun) -> AIDynamoSlurmComma
return AIDynamoSlurmCommandGenStrategy(slurm_system, test_run)


def test_gen_script_args_omits_launch_fields_for_disabled_prefill_worker(
strategy: AIDynamoSlurmCommandGenStrategy,
) -> None:
td = cast(AIDynamoTestDefinition, strategy.test_run.test)
td.cmd_args.dynamo.prefill_worker = WorkerConfig(num_nodes=0)

args = strategy._gen_script_args(td)

assert '--prefill-num-nodes "0"' in args
assert not any(arg.startswith("--prefill-cmd ") for arg in args)
assert not any(arg.startswith("--prefill-worker-initialized-regex ") for arg in args)


def test_container_mounts(strategy: AIDynamoSlurmCommandGenStrategy, test_run: TestRun) -> None:
mounts = strategy._container_mounts()

Expand Down
25 changes: 20 additions & 5 deletions tests/workloads/ai_dynamo/test_json_gen_strategy_kubernetes.py
Original file line number Diff line number Diff line change
Expand Up @@ -99,7 +99,7 @@ def test_gen_decode(json_gen: AIDynamoKubernetesJsonGenStrategy) -> None:
assert decode.get("replicas") == 1

args = ["--model", tdef.cmd_args.dynamo.model]
if tdef.cmd_args.dynamo.prefill_worker:
if tdef.cmd_args.dynamo.prefill_worker.is_enabled:
assert decode.get("subComponentType") == "decode-worker"
args.append("--is-decode-worker")

Expand All @@ -111,6 +111,7 @@ def test_gen_decode(json_gen: AIDynamoKubernetesJsonGenStrategy) -> None:
main_container = decode.get("extraPodSpec", {}).get("mainContainer", {})
assert main_container.get("image") == tdef.cmd_args.docker_image_url
assert main_container.get("workingDir") == tdef.cmd_args.dynamo.workspace_path
assert tdef.cmd_args.dynamo.decode_worker.cmd is not None
assert main_container.get("command") == tdef.cmd_args.dynamo.decode_worker.cmd.split()
assert main_container.get("args") == args

Expand All @@ -136,8 +137,8 @@ def test_gen_prefill(json_gen: AIDynamoKubernetesJsonGenStrategy) -> None:
system = cast(KubernetesSystem, json_gen.system)
tdef = cast(AIDynamoTestDefinition, json_gen.test_run.test)

if not tdef.cmd_args.dynamo.prefill_worker:
with pytest.raises(ValueError, match=r"Prefill worker configuration is not defined in the test definition."):
if not tdef.cmd_args.dynamo.prefill_worker.is_enabled:
with pytest.raises(ValueError, match=r"Prefill worker is disabled in the test definition."):
json_gen.gen_prefill_dict()
return

Expand All @@ -156,6 +157,7 @@ def test_gen_prefill(json_gen: AIDynamoKubernetesJsonGenStrategy) -> None:
main_container = prefill.get("extraPodSpec", {}).get("mainContainer", {})
assert main_container.get("image") == tdef.cmd_args.docker_image_url
assert main_container.get("workingDir") == tdef.cmd_args.dynamo.workspace_path
assert tdef.cmd_args.dynamo.prefill_worker.cmd is not None
assert main_container.get("command") == tdef.cmd_args.dynamo.prefill_worker.cmd.split()
assert main_container.get("args") == args

Expand All @@ -166,7 +168,7 @@ def test_gen_prefill(json_gen: AIDynamoKubernetesJsonGenStrategy) -> None:
@pytest.mark.parametrize("num_nodes", [1, 2, 4])
def test_gen_prefill_num_nodes(num_nodes: int, json_gen: AIDynamoKubernetesJsonGenStrategy) -> None:
tdef = cast(AIDynamoTestDefinition, json_gen.test_run.test)
if not tdef.cmd_args.dynamo.prefill_worker:
if not tdef.cmd_args.dynamo.prefill_worker.is_enabled:
pytest.skip("Prefill worker configuration is not defined in the test definition.")

tdef.cmd_args.dynamo.prefill_worker.num_nodes = num_nodes
Expand Down Expand Up @@ -203,7 +205,7 @@ def test_gen_json(json_gen: AIDynamoKubernetesJsonGenStrategy) -> None:
assert deployment.get("kind") == "DynamoGraphDeployment"
assert deployment.get("metadata", {}).get("name") == k8s_system.default_namespace

if tdef.cmd_args.dynamo.prefill_worker:
if tdef.cmd_args.dynamo.prefill_worker.is_enabled:
assert "prefill" in deployment.get("spec", {}).get("services", {})
else:
assert "spec" in deployment
Expand All @@ -215,6 +217,19 @@ def test_gen_json(json_gen: AIDynamoKubernetesJsonGenStrategy) -> None:
assert content == deployment


def test_gen_json_omits_disabled_prefill_worker(json_gen: AIDynamoKubernetesJsonGenStrategy) -> None:
tdef = cast(AIDynamoTestDefinition, json_gen.test_run.test)
tdef.cmd_args.dynamo.prefill_worker = WorkerConfig(num_nodes=0)
json_gen.test_run.output_path.mkdir(parents=True, exist_ok=True)

deployment = json_gen.gen_json()

services = deployment["spec"]["services"]
assert "prefill" not in services
assert "subComponentType" not in services["decode"]
assert "--is-decode-worker" not in services["decode"]["extraPodSpec"]["mainContainer"]["args"]


class TestDynamoCniNetworking:
CNI_NETS: ClassVar[list[str]] = ["default/nic0-rail0-plane0", "default/nic0-rail0-plane1"]

Expand Down
Loading