Skip to content
Open
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
125 changes: 124 additions & 1 deletion lib/ramble/ramble/expander.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@
import sys
from contextlib import contextmanager
from enum import Enum
from typing import Dict, FrozenSet, List, Optional, Union
from typing import Any, Dict, FrozenSet, List, Optional, Set, Union

import ramble.config
import ramble.error
Expand Down Expand Up @@ -223,11 +223,134 @@ def _maybe(expander, var_name, default=""):
}


def pow2_range(start, stop=None, inclusive=True):
"""Generate a sequence of numbers by doubling from start up to stop.

Args:
start (int): Starting value (if stop is provided), or stop value (if stop is None).
stop (int): Optional ending value. If None, start is treated as stop,
and sequence starts at 1.
inclusive (bool): Whether stop is inclusive. Defaults to True.

Returns:
list[int]: Sequence of doubling values / powers of 2.
"""
if stop is None:
start, stop = 1, start

start = int(start)
stop = int(stop)
if isinstance(inclusive, str):
inclusive = inclusive.lower() in ("true", "1", "yes")

if start <= 0 or stop < start:
return []

values = []
current = start
if inclusive:
while current <= stop:
values.append(current)
current *= 2
else:
while current < stop:
values.append(current)
current *= 2

return values


supported_list_function_pointers = {
"range": range,
"pow2_range": pow2_range,
"power2_range": pow2_range,
}


def is_dynamic_list_expression(val) -> bool:
"""Check if a variable value is an unexpanded call to a supported list function.

Returns True if val is a string matching a call to any function registered in
`supported_list_function_pointers`.
"""
if not isinstance(val, str):
return False
if not supported_list_function_pointers:
return False
func_names = "|".join(re.escape(fn) for fn in supported_list_function_pointers)
return bool(re.search(rf"\b(?:{func_names})\s*\(", val))


def extract_var_refs(expr) -> Set[str]:
"""Extract variable references enclosed in braces from an expression string."""
if not isinstance(expr, str):
return set()
refs = set(re.findall(r"\{\s*([a-zA-Z_][a-zA-Z0-9_]*)", expr))
in_brace = 0
cur_ident: List[str] = []
for ch in expr:
if ch == "{":
in_brace += 1
elif ch == "}":
if cur_ident:
refs.add("".join(cur_ident))
cur_ident = []
in_brace = max(0, in_brace - 1)
elif in_brace > 0:
if ch.isalnum() or ch == "_":
cur_ident.append(ch)
else:
if cur_ident:
refs.add("".join(cur_ident))
cur_ident = []
if cur_ident and in_brace > 0:
refs.add("".join(cur_ident))
return {r for r in refs if r and not r[0].isdigit()}


def find_dependent_vector_vars(
dynamic_list_vars: Dict[str, Any],
variables: Dict[str, Any],
expander: Optional["Expander"] = None,
) -> Set[str]:
"""Find all vector variables in `variables` that `dynamic_list_vars` transitively depend on."""
dependent_vector_vars = set()
visited = set()
queue = list(dynamic_list_vars.values())

if expander is not None:
for val in dynamic_list_vars.values():
saved_used = expander._used_variables.copy()
expander._used_variables = set()
try:
expander.expand_var(val)
except Exception:
pass
used = expander._used_variables
expander._used_variables = saved_used
for var_ref in used:
if var_ref in variables and var_ref not in dynamic_list_vars:
if isinstance(variables[var_ref], list):
dependent_vector_vars.add(var_ref)
elif var_ref not in visited:
visited.add(var_ref)
queue.append(str(variables[var_ref]))

while queue:
expr = queue.pop(0)
if not isinstance(expr, str):
continue
for var_ref in extract_var_refs(expr):
if var_ref in variables and var_ref not in dynamic_list_vars:
if isinstance(variables[var_ref], list):
dependent_vector_vars.add(var_ref)
elif var_ref not in visited:
visited.add(var_ref)
queue.append(str(variables[var_ref]))

return dependent_vector_vars


supported_modules = {
"math": math,
}
Expand Down
38 changes: 36 additions & 2 deletions lib/ramble/ramble/experiment_set.py
Original file line number Diff line number Diff line change
Expand Up @@ -285,6 +285,9 @@ def _prepare_experiment(

app_inst = self._setup_experiment_minimal(workload_template_name, variables, context)

if getattr(app_inst, "has_dynamic_range_variables", False):
return app_inst

final_wl_name = app_inst.expander.expand_var_name(
self.keywords.workload_name, allow_passthrough=False
)
Expand Down Expand Up @@ -345,6 +348,9 @@ def _get_used_variables(
):
app_inst = self._setup_experiment_minimal(workload_template_name, variables, context)

if getattr(app_inst, "has_dynamic_range_variables", False):
return set()

# The `_get_used_variables` is only called for the base experiment,
# so no need to consider repeat suffix.
exp_name = app_inst.expander.expand_var(exp_template_name, allow_passthrough=False)
Expand All @@ -369,6 +375,7 @@ def _process_render_object(
"""Helper to render a base and its repeated experiments, for parallel execution."""
experiment_vars, repeats = render_item
processed_experiments = []
dynamic_range_experiments = []
wl_stats = {}
# Expand and prepare base and repeated experiments
# TODO: Exploit the relationship between base and repeated experiments,
Expand All @@ -390,6 +397,10 @@ def _process_render_object(
cur_repeats,
)

if getattr(app_inst, "has_dynamic_range_variables", False):
dynamic_range_experiments.append(app_inst)
break

final_exp_name = app_inst.expander.expand_var_name(self.keywords.experiment_name)
final_exp_namespace = app_inst.expander.expand_var_name(
self.keywords.experiment_namespace
Expand Down Expand Up @@ -427,7 +438,7 @@ def _process_render_object(
if active:
app_inst.read_status()
processed_experiments.append((app_inst, final_exp_namespace, n == 0))
return processed_experiments, wl_stats
return processed_experiments, dynamic_range_experiments, wl_stats

def render_experiment_set(
self,
Expand Down Expand Up @@ -654,8 +665,10 @@ def _ingest_experiments(
results = list(executor.map(worker_func, render_list))

overall_wl_stats = {}
for processed_experiments, wl_stats in results:
all_dynamic_range_experiments = []
for processed_experiments, dynamic_range_experiments, wl_stats in results:
all_processed_experiments.extend(processed_experiments)
all_dynamic_range_experiments.extend(dynamic_range_experiments)
for wl_name, stats in wl_stats.items():
if wl_name not in overall_wl_stats:
overall_wl_stats[wl_name] = {
Expand All @@ -666,6 +679,27 @@ def _ingest_experiments(
overall_wl_stats[wl_name]["passed_global"] += stats["passed_global"]
overall_wl_stats[wl_name]["dropped_wl"] += stats["dropped_wl"]

if all_dynamic_range_experiments:
saved_contexts = {
self._contexts.application: self._context[self._contexts.application],
self._contexts.workload: self._context[self._contexts.workload],
self._contexts.experiment: self._context[self._contexts.experiment],
}
try:
for dyn_inst in all_dynamic_range_experiments:
range_rendered = dyn_inst.render_range_experiments(
saved_contexts[self._contexts.experiment],
warn_validation=warn_validation,
die_on_validate_error=die_on_validate_error,
chained=chained,
)
rendered_instances.extend(range_rendered)
for inst in range_rendered:
workload_names.add(inst.expander.workload_name)
finally:
for ctx_key, ctx_val in saved_contexts.items():
self._set_context(ctx_key, ctx_val)

# The results are now processed serially to update the experiment set state
for app_inst, final_exp_namespace, is_base_experiment in all_processed_experiments:
logger.debug(f" Final name: {final_exp_namespace}")
Expand Down
95 changes: 95 additions & 0 deletions lib/ramble/ramble/renderer.py
Original file line number Diff line number Diff line change
Expand Up @@ -468,6 +468,101 @@ def render_objects(self, render_group, exclude_where=None, ignore_used=True, fat
# Also expand all variables that generate lists
object_variables = self._expand_variables(variables, expander)

if render_group.object == "experiment":
dynamic_list_vars = {
k: v
for k, v in object_variables.items()
if ramble.expander.is_dynamic_list_expression(v)
}
if dynamic_list_vars:
dependent_vector_vars = ramble.expander.find_dependent_vector_vars(
dynamic_list_vars, object_variables, expander
)

if not dependent_vector_vars:
yield object_variables, ramble.repeats.Repeats()
return

# Expand only dependent_vector_vars (and any variables zipped with them)
# to produce separate seed instances for each vector value.
vars_to_scalarize = set(dependent_vector_vars)
if render_group.zips:
changed = True
while changed:
changed = False
for z_name, z_vars in render_group.zips.items():
if z_name in vars_to_scalarize or any(
v in vars_to_scalarize for v in z_vars
):
for v in z_vars:
if (
v not in dynamic_list_vars
and isinstance(object_variables.get(v), list)
and v not in vars_to_scalarize
):
vars_to_scalarize.add(v)
changed = True

if len(render_group.matrices) > 1:
if any(
any(v in vars_to_scalarize for v in mat) for mat in render_group.matrices
):
for mat in render_group.matrices:
for v in mat:
if v not in dynamic_list_vars and isinstance(
object_variables.get(v), list
):
vars_to_scalarize.add(v)

explicit_matrix_and_zip_vars = set()
for mat in render_group.matrices:
explicit_matrix_and_zip_vars.update(mat)
for z_name, z_vars in render_group.zips.items():
explicit_matrix_and_zip_vars.add(z_name)
explicit_matrix_and_zip_vars.update(z_vars)

if any(v not in explicit_matrix_and_zip_vars for v in vars_to_scalarize):
for k, v in object_variables.items():
if (
k not in dynamic_list_vars
and isinstance(v, list)
and k not in explicit_matrix_and_zip_vars
):
vars_to_scalarize.add(k)

sub_vars = {}
for k, v in object_variables.items():
if k in vars_to_scalarize or not isinstance(v, list):
if k not in dynamic_list_vars:
sub_vars[k] = v

sub_group = RenderGroup(render_group.object, render_group.action)
sub_group.variables = sub_vars

sub_zips = {}
for z_name, z_vars in render_group.zips.items():
filtered_z = [v for v in z_vars if v in vars_to_scalarize]
if len(filtered_z) > 1:
sub_zips[z_name] = filtered_z
sub_group.zips = sub_zips

sub_matrices = []
for mat in render_group.matrices:
filtered_mat = [v for v in mat if v in vars_to_scalarize or v in sub_zips]
if len(filtered_mat) >= 1:
sub_matrices.append(filtered_mat)
sub_group.matrices = sub_matrices
sub_group.n_repeats = 1
sub_group.used_variables = render_group.used_variables.union(vars_to_scalarize)

for rendered_sub_vars, _ in self.render_objects(
sub_group, ignore_used=False, fatal=fatal
):
seed_vars = object_variables.copy()
seed_vars.update(rendered_sub_vars)
yield seed_vars, ramble.repeats.Repeats()
return

# Expand zip and matrix members to allow indirections like
# ```
# variables:
Expand Down
29 changes: 29 additions & 0 deletions lib/ramble/ramble/test/expander.py
Original file line number Diff line number Diff line change
Expand Up @@ -417,3 +417,32 @@ class StatusEnum(str, Enum):

enum_expander = ramble.expander.Expander({"status": StatusEnum.SETUP}, None)
assert enum_expander.expand_var("{status}") == "SETUP"


def test_pow2_range_function():
from ramble.expander import pow2_range

assert pow2_range(8) == [1, 2, 4, 8]
assert pow2_range(1, 16) == [1, 2, 4, 8, 16]
assert pow2_range(2, 16) == [2, 4, 8, 16]
assert pow2_range(1, 15) == [1, 2, 4, 8]
assert pow2_range(1, 16, inclusive=False) == [1, 2, 4, 8]
assert pow2_range(1, 1) == [1]
assert pow2_range(1, 1, inclusive=False) == []
assert pow2_range(16, 4) == []
assert pow2_range(0) == []
assert pow2_range(-1, 8) == []
assert pow2_range(3, 24) == [3, 6, 12, 24]


def test_pow2_range_in_expander():
expander = ramble.expander.Expander({"max_nodes": "16"}, None)

assert expander.expand_lists("pow2_range(1, 16)") == [1, 2, 4, 8, 16]
assert expander.expand_lists("power2_range(1, 16)") == [1, 2, 4, 8, 16]

res = expander.expand_var("pow2_range(1, {max_nodes})", typed=True)
assert res == [1, 2, 4, 8, 16]

res_alias = expander.expand_var("power2_range(1, {max_nodes})", typed=True)
assert res_alias == [1, 2, 4, 8, 16]
Loading
Loading