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
29 changes: 29 additions & 0 deletions doc/USER_GUIDE.rst
Original file line number Diff line number Diff line change
Expand Up @@ -227,6 +227,35 @@ action, typically seeded by ``random_seed``.

Custom agents may extend the ``BaseAgentConfig`` and offer more parameters to configure.

Declarative DSE constraints
~~~~~~~~~~~~~~~~~~~~~~~~~~~

Use ``dse_constraints`` in a test TOML to reject invalid parameter combinations before they are executed. Declare
aliases once under ``variables``, then define one or more named Boolean expressions. Variable paths may select fields
from ``cmd_args``, ``extra_env_vars``, ``system``, or ``test_run``.

.. code-block:: toml

[dse_constraints.variables]
prefill_tp = "cmd_args.dynamo.prefill_worker.args.tensor_parallel_size"
prefill_pp = "cmd_args.dynamo.prefill_worker.args.pipeline_parallel_size"
decode_tp = "cmd_args.dynamo.decode_worker.args.tensor_parallel_size"
decode_pp = "cmd_args.dynamo.decode_worker.args.pipeline_parallel_size"
gpus_per_node = "system.gpus_per_node"

[dse_constraints.expressions]
prefill_not_larger = "prefill_tp <= decode_tp"
prefill_fits = "prefill_tp * prefill_pp <= gpus_per_node"
decode_fits = "decode_tp * decode_pp <= gpus_per_node"

Expressions support numeric arithmetic, comparisons, membership checks, and the ``and``, ``or``, and ``not``
operators. Function calls and private attribute access are not allowed. A malformed expression or unresolved variable
path fails the DSE run with an evaluation error instead of silently accepting or rejecting the configuration.

The expression key is used as the constraint name in rejection logs. Constraints run in declaration order and stop at
the first failure, before the workload's Python ``constraint_check`` method. Workload-specific checks remain available
for validation that is too complex to express declaratively.

DSE parameter exclusions
~~~~~~~~~~~~~~~~~~~~~~~~

Expand Down
2 changes: 1 addition & 1 deletion src/cloudai/configurator/cloudai_gym.py
Original file line number Diff line number Diff line change
Expand Up @@ -152,7 +152,7 @@ def step(self, action: Any) -> Tuple[list, float, bool, dict]:
}
reward = cast(float, cached_result["reward"])
else:
if not self.test_run.test.constraint_check(self.test_run, self.runner.system):
if not self.test_run.test.check_constraints(self.test_run, self.runner.system):
logging.info("Constraint check failed. Skipping step.")
return [-1.0], self.rewards.constraint_failure, True, info

Expand Down
3 changes: 3 additions & 0 deletions src/cloudai/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,7 @@
)
from .configurator.grid_search import GridSearchAgent
from .configurator.gymnasium_adapter import GymnasiumAdapter
from .models.dse_constraint import ConstraintEvaluationError, DSEConstraints
from .models.workload import CmdArgs, NsysConfiguration, PredictorConfig, TestDefinition
from .parser import Parser
from .reporter import JUnitReporter, PerTestReporter, StatusReporter, TarballReporter
Expand All @@ -87,6 +88,8 @@
"CmdArgs",
"CommandGenStrategy",
"ConfigPaths",
"ConstraintEvaluationError",
"DSEConstraints",
"DockerImage",
"Encoding",
"File",
Expand Down
192 changes: 192 additions & 0 deletions src/cloudai/models/dse_constraint.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,192 @@
# SPDX-FileCopyrightText: NVIDIA CORPORATION & AFFILIATES
# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Safe declarative constraints for DSE configurations."""

from __future__ import annotations

import ast
import operator
from collections.abc import Mapping
from typing import Any, Callable

from pydantic import BaseModel, ConfigDict, Field, model_validator
from typing_extensions import Self


class ConstraintEvaluationError(ValueError):
"""Raised when a declarative DSE constraint cannot be evaluated."""


_ROOT_NAMES = {"cmd_args", "extra_env_vars", "system", "test_run"}
_BINARY_OPERATORS: dict[type[ast.operator], Callable[[Any, Any], Any]] = {
ast.Add: operator.add,
ast.Sub: operator.sub,
ast.Mult: operator.mul,
ast.Div: operator.truediv,
ast.FloorDiv: operator.floordiv,
ast.Mod: operator.mod,
}
_COMPARISON_OPERATORS: dict[type[ast.cmpop], Callable[[Any, Any], bool]] = {
ast.Eq: operator.eq,
ast.NotEq: operator.ne,
ast.Lt: operator.lt,
ast.LtE: operator.le,
ast.Gt: operator.gt,
ast.GtE: operator.ge,
ast.In: lambda left, right: left in right,
ast.NotIn: lambda left, right: left not in right,
}
_ALLOWED_AST_NODES = (
ast.Expression,
ast.BoolOp,
ast.And,
ast.Or,
ast.BinOp,
*_BINARY_OPERATORS,
ast.UnaryOp,
ast.Not,
ast.UAdd,
ast.USub,
ast.Compare,
*_COMPARISON_OPERATORS,
ast.Name,
ast.Attribute,
ast.Subscript,
ast.Constant,
ast.List,
ast.Tuple,
ast.Set,
ast.Load,
)


def _parse_expression(expression: str) -> ast.Expression:
try:
tree = ast.parse(expression, mode="eval")
except SyntaxError as e:
raise ValueError(f"Invalid constraint expression: {e.msg}") from e

for node in ast.walk(tree):
if not isinstance(node, _ALLOWED_AST_NODES):
raise ValueError(f"Constraint expression does not allow {type(node).__name__}")
if isinstance(node, ast.Attribute) and node.attr.startswith("_"):
raise ValueError("Constraint expression cannot access private attributes")
if isinstance(node, ast.Constant) and not isinstance(node.value, (bool, int, float, str, type(None))):
raise ValueError(f"Constraint expression does not allow {type(node.value).__name__} literals")
return tree


def _resolve_path(context: Mapping[str, Any], path: str) -> Any:
value: Any = context
for component in path.split("."):
if not isinstance(value, Mapping) or component not in value:
raise ConstraintEvaluationError(f"Cannot resolve constraint variable path '{path}'")
value = value[component]
return value


def _resolve_member(value: Any, key: Any) -> Any:
if not isinstance(value, Mapping) or key not in value:
raise ConstraintEvaluationError(f"Cannot resolve constraint member {key!r}")
return value[key]


def _evaluate_node(node: ast.AST, context: Mapping[str, Any]) -> Any: # noqa: C901
if isinstance(node, ast.Expression):
return _evaluate_node(node.body, context)
if isinstance(node, ast.Constant):
return node.value
if isinstance(node, ast.Name):
if node.id not in context:
raise ConstraintEvaluationError(f"Unknown constraint variable '{node.id}'")
return context[node.id]
if isinstance(node, ast.Attribute):
return _resolve_member(_evaluate_node(node.value, context), node.attr)
if isinstance(node, ast.Subscript):
return _resolve_member(_evaluate_node(node.value, context), _evaluate_node(node.slice, context))
if isinstance(node, ast.List):
return [_evaluate_node(element, context) for element in node.elts]
if isinstance(node, ast.Tuple):
return tuple(_evaluate_node(element, context) for element in node.elts)
if isinstance(node, ast.Set):
return {_evaluate_node(element, context) for element in node.elts}
if isinstance(node, ast.BoolOp):
if isinstance(node.op, ast.And):
return all(_evaluate_node(value, context) for value in node.values)
return any(_evaluate_node(value, context) for value in node.values)
if isinstance(node, ast.UnaryOp):
value = _evaluate_node(node.operand, context)
if isinstance(node.op, ast.Not):
return not value
if not isinstance(value, (int, float)) or isinstance(value, bool):
raise ConstraintEvaluationError("Unary arithmetic requires a numeric operand")
return +value if isinstance(node.op, ast.UAdd) else -value
if isinstance(node, ast.BinOp):
left = _evaluate_node(node.left, context)
right = _evaluate_node(node.right, context)
if any(not isinstance(value, (int, float)) or isinstance(value, bool) for value in (left, right)):
raise ConstraintEvaluationError("Constraint arithmetic requires numeric operands")
return _BINARY_OPERATORS[type(node.op)](left, right)
if isinstance(node, ast.Compare):
left = _evaluate_node(node.left, context)
for op, comparator in zip(node.ops, node.comparators, strict=True):
right = _evaluate_node(comparator, context)
if not _COMPARISON_OPERATORS[type(op)](left, right):
return False
left = right
return True
raise ConstraintEvaluationError(f"Unsupported constraint expression element: {type(node).__name__}")


class DSEConstraints(BaseModel):
"""Shared variable bindings and named Boolean constraints for DSE candidates."""

model_config = ConfigDict(extra="forbid")

variables: dict[str, str] = Field(default_factory=dict)
expressions: dict[str, str] = Field(min_length=1)

@model_validator(mode="after")
def validate_constraints(self) -> Self:
for alias, path in self.variables.items():
if not alias.isidentifier() or alias.startswith("_"):
raise ValueError(f"Invalid constraint variable name: {alias!r}")
if alias in _ROOT_NAMES:
raise ValueError(f"Constraint variable name is reserved: {alias!r}")
root, separator, remainder = path.partition(".")
if root not in _ROOT_NAMES or not separator or not remainder or ".." in path:
raise ValueError(
"Constraint variable path must start with one of "
f"{sorted(_ROOT_NAMES)} and select a field: {path!r}"
)

for name, expression in self.expressions.items():
if not name.strip():
raise ValueError("Constraint names cannot be empty")
tree = _parse_expression(expression)
referenced_names = {node.id for node in ast.walk(tree) if isinstance(node, ast.Name)}
unknown_names = referenced_names - _ROOT_NAMES - self.variables.keys()
if unknown_names:
raise ValueError(f"Constraint '{name}' uses unknown variables: {', '.join(sorted(unknown_names))}")
return self

def evaluate(self, context: Mapping[str, Any]) -> tuple[bool, str | None, str | None]:
"""Evaluate constraints in declaration order and return the first failure."""
evaluation_context = dict(context)
try:
evaluation_context.update({alias: _resolve_path(context, path) for alias, path in self.variables.items()})
except ConstraintEvaluationError as e:
raise ConstraintEvaluationError(f"Failed to resolve DSE constraint variables: {e}") from e

for name, expression in self.expressions.items():
try:
result = _evaluate_node(_parse_expression(expression), evaluation_context)
except (ArithmeticError, ConstraintEvaluationError, TypeError, ValueError) as e:
raise ConstraintEvaluationError(f"Failed to evaluate DSE constraint '{name}': {e}") from e
if not isinstance(result, bool):
raise ConstraintEvaluationError(f"DSE constraint '{name}' did not evaluate to a Boolean")
if not result:
return False, name, expression
return True, None, None
32 changes: 32 additions & 0 deletions src/cloudai/models/workload.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.

import logging
from abc import ABC
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Union
Expand All @@ -25,6 +26,7 @@
from cloudai.core import GitRepo, Installable, JobStatusResult, PythonExecutable, Registry, System, TestRun

from ..configurator.env_params import EnvParamSpec
from .dse_constraint import DSEConstraints


class CmdArgs(BaseModel):
Expand Down Expand Up @@ -123,6 +125,10 @@ class TestDefinition(BaseModel, ABC):
agent_metrics: list[str] = Field(default=["default"])
agent_reward_function: str = "inverse"
agent_config: dict[str, Any] | None = Field(default=None, description="Agent configuration.")
dse_constraints: DSEConstraints | None = Field(
default=None,
description="Declarative constraints evaluated before a DSE configuration is executed.",
)
env_params: dict[str, EnvParamSpec] = Field(
default_factory=dict,
description=(
Expand Down Expand Up @@ -155,6 +161,32 @@ def installables(self) -> list[Installable]:
def constraint_check(self, tr: TestRun, system: Optional[System]) -> bool:
return True

def check_constraints(self, tr: TestRun, system: Optional[System]) -> bool:
"""Evaluate declarative constraints followed by workload-specific constraints."""
if self.dse_constraints:
context = {
"cmd_args": tr.test.cmd_args.model_dump(mode="python"),
"extra_env_vars": tr.test.extra_env_vars,
"system": system.model_dump(mode="python") if system is not None else {},
"test_run": {
"name": tr.name,
"num_nodes": tr.num_nodes,
"nodes": tr.nodes,
"iterations": tr.iterations,
"current_iteration": tr.current_iteration,
"step": tr.step,
},
}
accepted, failed_name, failed_expression = self.dse_constraints.evaluate(context)
if not accepted:
logging.info(
"DSE constraint '%s' rejected the configuration: %s",
failed_name,
failed_expression,
)
return False
return self.constraint_check(tr, system)

def is_env_sampled(self, cmd_args_path: str) -> bool:
"""Whether a cmd_args field is env-sampled (env draws it per trial, not the agent)."""
return cmd_args_path in self.env_params
Expand Down
4 changes: 2 additions & 2 deletions src/cloudai/systems/slurm/single_sbatch_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -132,7 +132,7 @@ def unroll_dse(self, tr: TestRun) -> Generator[TestRun, None, None]:
next_tr.step = idx
next_tr.output_path = self.get_job_output_path(next_tr)

if next_tr.test.constraint_check(next_tr, self.system):
if next_tr.test.check_constraints(next_tr, self.system):
yield next_tr

def get_global_env_vars(self) -> str:
Expand Down Expand Up @@ -225,7 +225,7 @@ def handle_dse(self):
next_tr.step = idx
next_tr.output_path = self.get_job_output_path(next_tr)

if not next_tr.test.constraint_check(next_tr, self.system):
if not next_tr.test.check_constraints(next_tr, self.system):
continue

gym.test_run = next_tr
Expand Down
Loading