Skip to content
Merged
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
66 changes: 10 additions & 56 deletions inferencex-e2e/infx/tests/matrix/test_generate_sweep_configs.py
Original file line number Diff line number Diff line change
Expand Up @@ -1068,9 +1068,15 @@ def test_sweep_expands_each_sequence_length_across_concurrencies(self, sample_si
sample_single_node_config,
sample_runner_config
)
assert [(row["isl"], row["osl"], row["conc"]) for row in result] == [
(isl, osl, conc)
for isl, osl in [(1024, 1024), (8192, 1024)]
assert [
(row["isl"], row["osl"], row["conc"], row["exp-name"], row["max-model-len"])
for row in result
] == [
(isl, osl, conc, name, context)
for isl, osl, name, context in [
(1024, 1024, "dsr1_1k1k", 2304),
(8192, 1024, "dsr1_8k1k", 9472),
]
for conc in [4, 8, 16, 32, 64]
]

Expand Down Expand Up @@ -1305,27 +1311,6 @@ def test_step_size(self, sample_single_node_config, sample_runner_config, full_s
assert 16 in conc_values
assert 64 in conc_values

def test_exp_name_format(self, sample_single_node_config, sample_runner_config, full_sweep_args_single_node):
full_sweep_args_single_node.seq_lens = ["1k1k"]
result = generate_full_sweep(
full_sweep_args_single_node,
sample_single_node_config,
sample_runner_config
)
assert all(entry["exp-name"] == "dsr1_1k1k" for entry in result)

def test_max_model_len_calculation(self, sample_single_node_config, sample_runner_config, full_sweep_args_single_node):
"""max-model-len should be isl + osl + 256."""
result = generate_full_sweep(
full_sweep_args_single_node,
sample_single_node_config,
sample_runner_config
)
assert {
(entry["isl"], entry["osl"], entry["max-model-len"])
for entry in result
} == {(1024, 1024, 2304), (8192, 1024, 9472)}

def test_runner_node_filter(self, sample_single_node_config, sample_runner_config, full_sweep_args_single_node):
"""Runner node filter should expand entries to individual matching nodes."""
full_sweep_args_single_node.runner_type = ["mi300x"]
Expand Down Expand Up @@ -1353,20 +1338,6 @@ def test_runner_node_filter_no_match(self, sample_single_node_config, sample_run
)
assert len(result) == 0

def test_runner_node_filter_without_runner_type(self, sample_single_node_config, sample_runner_config, full_sweep_args_single_node):
"""Runner node filter should work without explicit runner type (uses config's runner)."""
full_sweep_args_single_node.runner_node_filter = "amd"
full_sweep_args_single_node.seq_lens = ["1k1k"]
full_sweep_args_single_node.max_conc = 4
result = generate_full_sweep(
full_sweep_args_single_node,
sample_single_node_config,
sample_runner_config
)
# Config has runner=mi300x, filter "amd" matches mi300x-amd_0 and mi300x-amd_1
assert len(result) == 2
assert all("amd" in entry["runner"] for entry in result)



class TestGenerateFullSweepMultiNode:
Expand All @@ -1379,6 +1350,7 @@ def test_multinode_entry_structure(self, sample_multinode_config, sample_runner_
sample_runner_config
)
entry = result[0]
assert entry["conc"] == [2150]
assert entry["prefill"]["num-worker"] == 5
assert entry["decode"]["num-worker"] == 1
assert entry["disagg"] is True
Expand Down Expand Up @@ -1418,24 +1390,6 @@ def test_multinode_parallelism_fields(self, sample_multinode_config, sample_runn
entry["decode"]["pcp-size"],
) == (2, 4, 1)

def test_multinode_conc_as_list(self, sample_multinode_config, sample_runner_config, full_sweep_args_multi_node):
"""Multinode conc should be passed as list."""
result = generate_full_sweep(
full_sweep_args_multi_node,
sample_multinode_config,
sample_runner_config
)
entry = result[0]
assert entry["conc"] == [2150]

def test_single_node_flag_skips_multinode(self, sample_multinode_config, sample_runner_config, full_sweep_args_single_node):
result = generate_full_sweep(
full_sweep_args_single_node,
sample_multinode_config,
sample_runner_config
)
assert len(result) == 0

def test_runner_node_filter_multinode(self, sample_runner_config, full_sweep_args_multi_node):
# Create a multinode config with h200 runner (which has 4 nodes)
config = {
Expand Down
Loading