[AMD][CI] CI: fix AMD 2-GPU multimodal-gen partition-count abort (#34328)
Co-authored-by: Cursor Agent <cursoragent@cursor.com> Co-authored-by: quitenode <quitenode@users.noreply.github.com>
This commit is contained in:
co-authored by
Cursor Agent
quitenode
parent
273d978bed
commit
bbda7f32b1
@@ -30,3 +30,17 @@ def partition_items_by_lpt(
|
||||
partition_sums[min_idx] += item.est_time
|
||||
|
||||
return partitions
|
||||
|
||||
|
||||
def assign_partition(
|
||||
items: list[PartitionItem], partition_id: int, num_partitions: int
|
||||
) -> list[PartitionItem]:
|
||||
"""Return the LPT slice of ``items`` owned by ``partition_id``.
|
||||
|
||||
The LPT pass is deterministic, so shards that each call this with the same
|
||||
item list cover the list exactly once between them.
|
||||
"""
|
||||
partitions = partition_items_by_lpt(items, num_partitions)
|
||||
if partition_id < 0 or partition_id >= len(partitions):
|
||||
return []
|
||||
return partitions[partition_id]
|
||||
|
||||
@@ -20,10 +20,7 @@ import tabulate
|
||||
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.test.partitioning import (
|
||||
PartitionItem,
|
||||
partition_items_by_lpt,
|
||||
)
|
||||
from sglang.multimodal_gen.test.partitioning import PartitionItem, assign_partition
|
||||
from sglang.multimodal_gen.test.runner.pytest_runner import (
|
||||
partition_items_by_index,
|
||||
run_pytest,
|
||||
@@ -114,23 +111,6 @@ def validate_standalone_file_est_times() -> dict[str, list[str]]:
|
||||
return missing_by_suite
|
||||
|
||||
|
||||
def auto_partition(
|
||||
cases: list[DiffusionTestCase], rank: int, size: int
|
||||
) -> list[DiffusionTestCase]:
|
||||
if not cases or size <= 0:
|
||||
return []
|
||||
|
||||
case_by_id = {case.id: case for case in cases}
|
||||
items = [
|
||||
PartitionItem(kind="case", item_id=case.id, est_time=get_case_est_time(case.id))
|
||||
for case in cases
|
||||
]
|
||||
partitions = partition_items_by_lpt(items, size)
|
||||
if rank >= len(partitions):
|
||||
return []
|
||||
return [case_by_id[item.item_id] for item in partitions[rank]]
|
||||
|
||||
|
||||
def get_suite_files_rel(suite: str, parametrized_only: bool = False) -> list[str]:
|
||||
if parametrized_only and suite in PARAMETRIZED_CASE_GROUPS:
|
||||
return [filename for filename, _ in PARAMETRIZED_CASE_GROUPS[suite]]
|
||||
@@ -183,6 +163,39 @@ def parse_partition_plan(
|
||||
)
|
||||
|
||||
|
||||
def build_local_partition_assignment(
|
||||
suite: str,
|
||||
partition_id: int,
|
||||
total_partitions: int,
|
||||
) -> PartitionAssignment:
|
||||
"""Assign this shard's work when CI did not precompute a partition plan.
|
||||
|
||||
Lanes with a hardcoded ``--total-partitions`` (the AMD ones) cannot give
|
||||
every standalone file a shard of its own, so standalone files are LPT
|
||||
balanced together with the parametrized cases instead.
|
||||
"""
|
||||
items = [
|
||||
PartitionItem(kind="case", item_id=case.id, est_time=get_case_est_time(case.id))
|
||||
for case in _get_dynamic_suite_cases(suite)
|
||||
]
|
||||
for standalone_file in STANDALONE_FILES.get(suite, []):
|
||||
items.append(
|
||||
PartitionItem(
|
||||
kind="standalone",
|
||||
item_id=standalone_file,
|
||||
est_time=get_standalone_file_est_time(suite, standalone_file)[0],
|
||||
)
|
||||
)
|
||||
|
||||
my_items = assign_partition(items, partition_id, total_partitions)
|
||||
return PartitionAssignment(
|
||||
case_ids=[item.item_id for item in my_items if item.kind == "case"],
|
||||
standalone_files=[
|
||||
item.item_id for item in my_items if item.kind == "standalone"
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def _merge_execution_results(
|
||||
executed_cases: list[str],
|
||||
case_results: dict[str, str],
|
||||
@@ -467,19 +480,6 @@ def _get_parametrized_files_for_case_ids(
|
||||
return files
|
||||
|
||||
|
||||
def _get_standalone_file(target_dir: Path, suite: str, index: int) -> str | None:
|
||||
standalone_files = STANDALONE_FILES.get(suite, [])
|
||||
if index < 0 or index >= len(standalone_files):
|
||||
return None
|
||||
file_path = target_dir / standalone_files[index]
|
||||
if file_path.exists():
|
||||
return str(file_path)
|
||||
logger.warning(
|
||||
"Standalone test file %s not found in %s", standalone_files[index], target_dir
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
def _run_dynamic_suite(args, target_dir: Path) -> int:
|
||||
if args.partition_plan_json:
|
||||
assignment = parse_partition_plan(
|
||||
@@ -488,251 +488,139 @@ def _run_dynamic_suite(args, target_dir: Path) -> int:
|
||||
total_partitions=args.total_partitions,
|
||||
plan_json=args.partition_plan_json,
|
||||
)
|
||||
|
||||
rows = [[args.suite, f"{args.partition_id + 1}/{args.total_partitions}"]]
|
||||
print(tabulate.tabulate(rows, headers=["Suite", "Partition"], tablefmt="psql"))
|
||||
|
||||
total_est_time = 0.0
|
||||
executed_cases: list[str] = []
|
||||
case_results: dict[str, str] = {}
|
||||
missing_standalone_estimates: list[str] = []
|
||||
standalone_measurements: list[dict] = []
|
||||
overall_exit_code = 0
|
||||
|
||||
if assignment.case_ids:
|
||||
case_id_set = set(assignment.case_ids)
|
||||
total_est_time += sum(
|
||||
get_case_est_time(case_id) for case_id in assignment.case_ids
|
||||
)
|
||||
suite_files = _get_parametrized_files_for_case_ids(
|
||||
args.suite, case_id_set, target_dir
|
||||
)
|
||||
if not suite_files:
|
||||
print(
|
||||
f"No valid parametrized test files found for suite '{args.suite}'."
|
||||
)
|
||||
return 0
|
||||
|
||||
partition_filter = " or ".join(
|
||||
f"[{case_id}]" for case_id in assignment.case_ids
|
||||
)
|
||||
filter_expr = (
|
||||
f"({partition_filter}) and ({args.filter})"
|
||||
if args.filter
|
||||
else partition_filter
|
||||
)
|
||||
|
||||
print(
|
||||
f"Running {len(assignment.case_ids)} parametrized cases with estimated total "
|
||||
f"{sum(get_case_est_time(case_id) for case_id in assignment.case_ids):.1f}s:"
|
||||
)
|
||||
for case_id in assignment.case_ids:
|
||||
print(f" - case: {case_id} ({get_case_est_time(case_id):.1f}s)")
|
||||
print(f"Test files: {[Path(f).name for f in suite_files]}")
|
||||
print(f"Filter expression: {filter_expr}")
|
||||
|
||||
junit_xml_path = str(
|
||||
target_dir / f"junit_results_{args.suite}_{args.partition_id}.xml"
|
||||
)
|
||||
exit_code, new_executed_cases, new_case_results = run_pytest(
|
||||
suite_files,
|
||||
filter_expr=filter_expr,
|
||||
junit_xml_path=junit_xml_path,
|
||||
)
|
||||
_merge_execution_results(
|
||||
executed_cases, case_results, new_executed_cases, new_case_results
|
||||
)
|
||||
if exit_code != 0 and overall_exit_code == 0:
|
||||
overall_exit_code = exit_code
|
||||
if exit_code != 0 and not args.continue_on_error:
|
||||
write_execution_report(
|
||||
suite=args.suite,
|
||||
partition_id=args.partition_id,
|
||||
total_partitions=args.total_partitions,
|
||||
executed_cases=executed_cases,
|
||||
is_standalone=False,
|
||||
standalone_file=None,
|
||||
case_results=case_results,
|
||||
missing_standalone_estimates=missing_standalone_estimates,
|
||||
standalone_measurements=standalone_measurements,
|
||||
)
|
||||
return overall_exit_code
|
||||
|
||||
if assignment.standalone_files:
|
||||
standalone_estimate = sum(
|
||||
get_standalone_file_est_time(args.suite, standalone_file)[0]
|
||||
for standalone_file in assignment.standalone_files
|
||||
)
|
||||
total_est_time += standalone_estimate
|
||||
print(
|
||||
f"Running {len(assignment.standalone_files)} standalone file(s) with estimated total "
|
||||
f"{standalone_estimate:.1f}s:"
|
||||
)
|
||||
for standalone_file in assignment.standalone_files:
|
||||
est_time, used_fallback_estimate = get_standalone_file_est_time(
|
||||
args.suite, standalone_file
|
||||
)
|
||||
fallback_suffix = (
|
||||
f", fallback estimate {DEFAULT_STANDALONE_EST_TIME_SECONDS:.1f}s"
|
||||
if used_fallback_estimate
|
||||
else ""
|
||||
)
|
||||
print(
|
||||
f" - standalone: {standalone_file} "
|
||||
f"({est_time:.1f}s{fallback_suffix})"
|
||||
)
|
||||
|
||||
for standalone_file in assignment.standalone_files:
|
||||
exit_code, new_executed_cases, new_case_results, measurement = (
|
||||
_run_standalone_file(
|
||||
args.suite,
|
||||
standalone_file,
|
||||
target_dir,
|
||||
extra_filter=args.filter,
|
||||
)
|
||||
)
|
||||
if measurement["used_fallback_estimate"]:
|
||||
missing_standalone_estimates.append(standalone_file)
|
||||
standalone_measurements.append(measurement)
|
||||
_merge_execution_results(
|
||||
executed_cases,
|
||||
case_results,
|
||||
new_executed_cases,
|
||||
new_case_results,
|
||||
)
|
||||
if exit_code != 0 and overall_exit_code == 0:
|
||||
overall_exit_code = exit_code
|
||||
if exit_code != 0 and not args.continue_on_error:
|
||||
break
|
||||
|
||||
print(f"Partition estimated total time: {total_est_time:.1f}s")
|
||||
write_execution_report(
|
||||
else:
|
||||
assignment = build_local_partition_assignment(
|
||||
suite=args.suite,
|
||||
partition_id=args.partition_id,
|
||||
total_partitions=args.total_partitions,
|
||||
executed_cases=executed_cases,
|
||||
is_standalone=False,
|
||||
standalone_file=None,
|
||||
case_results=case_results,
|
||||
missing_standalone_estimates=missing_standalone_estimates,
|
||||
standalone_measurements=standalone_measurements,
|
||||
)
|
||||
return overall_exit_code
|
||||
return _run_partition_assignment(args, target_dir, assignment)
|
||||
|
||||
all_cases = _get_dynamic_suite_cases(args.suite)
|
||||
standalone_files = STANDALONE_FILES.get(args.suite, [])
|
||||
parametrized_partitions = args.total_partitions - len(standalone_files)
|
||||
|
||||
if parametrized_partitions < 0:
|
||||
print(
|
||||
f"Error: total_partitions ({args.total_partitions}) must be >= "
|
||||
f"standalone files ({len(standalone_files)})"
|
||||
def _run_partition_assignment(
|
||||
args, target_dir: Path, assignment: PartitionAssignment
|
||||
) -> int:
|
||||
rows = [[args.suite, f"{args.partition_id + 1}/{args.total_partitions}"]]
|
||||
print(tabulate.tabulate(rows, headers=["Suite", "Partition"], tablefmt="psql"))
|
||||
|
||||
total_est_time = 0.0
|
||||
executed_cases: list[str] = []
|
||||
case_results: dict[str, str] = {}
|
||||
missing_standalone_estimates: list[str] = []
|
||||
standalone_measurements: list[dict] = []
|
||||
overall_exit_code = 0
|
||||
|
||||
if assignment.case_ids:
|
||||
case_id_set = set(assignment.case_ids)
|
||||
total_est_time += sum(
|
||||
get_case_est_time(case_id) for case_id in assignment.case_ids
|
||||
)
|
||||
return 1
|
||||
|
||||
if args.partition_id < parametrized_partitions:
|
||||
if not all_cases:
|
||||
print(f"No cases found for suite '{args.suite}'.")
|
||||
return 0
|
||||
|
||||
my_cases = auto_partition(all_cases, args.partition_id, parametrized_partitions)
|
||||
if not my_cases:
|
||||
print(
|
||||
f"No cases assigned to partition {args.partition_id}. Exiting success."
|
||||
)
|
||||
write_execution_report(
|
||||
suite=args.suite,
|
||||
partition_id=args.partition_id,
|
||||
total_partitions=args.total_partitions,
|
||||
executed_cases=[],
|
||||
is_standalone=False,
|
||||
missing_standalone_estimates=[],
|
||||
standalone_measurements=[],
|
||||
)
|
||||
return 0
|
||||
|
||||
case_ids = [case.id for case in my_cases]
|
||||
case_id_set = set(case_ids)
|
||||
total_est_time = sum(get_case_est_time(case.id) for case in my_cases)
|
||||
suite_files = _get_parametrized_files_for_case_ids(
|
||||
args.suite, case_id_set, target_dir
|
||||
)
|
||||
|
||||
if not suite_files:
|
||||
print(f"No valid parametrized test files found for suite '{args.suite}'.")
|
||||
return 0
|
||||
|
||||
partition_filter = " or ".join(f"[{case_id}]" for case_id in case_ids)
|
||||
partition_filter = " or ".join(
|
||||
f"[{case_id}]" for case_id in assignment.case_ids
|
||||
)
|
||||
filter_expr = (
|
||||
f"({partition_filter}) and ({args.filter})"
|
||||
if args.filter
|
||||
else partition_filter
|
||||
)
|
||||
|
||||
rows = [[args.suite, f"{args.partition_id + 1}/{args.total_partitions}"]]
|
||||
print(tabulate.tabulate(rows, headers=["Suite", "Partition"], tablefmt="psql"))
|
||||
print(
|
||||
f"Running {len(my_cases)} cases with estimated total "
|
||||
f"{total_est_time:.1f}s:"
|
||||
f"Running {len(assignment.case_ids)} parametrized cases with estimated total "
|
||||
f"{sum(get_case_est_time(case_id) for case_id in assignment.case_ids):.1f}s:"
|
||||
)
|
||||
for case in my_cases:
|
||||
print(f" - {case.id} ({get_case_est_time(case.id):.1f}s)")
|
||||
for case_id in assignment.case_ids:
|
||||
print(f" - case: {case_id} ({get_case_est_time(case_id):.1f}s)")
|
||||
print(f"Test files: {[Path(f).name for f in suite_files]}")
|
||||
print(f"Filter expression: {filter_expr}")
|
||||
|
||||
junit_xml_path = str(
|
||||
target_dir / f"junit_results_{args.suite}_{args.partition_id}.xml"
|
||||
)
|
||||
exit_code, executed_cases, case_results = run_pytest(
|
||||
exit_code, new_executed_cases, new_case_results = run_pytest(
|
||||
suite_files,
|
||||
filter_expr=filter_expr,
|
||||
junit_xml_path=junit_xml_path,
|
||||
)
|
||||
write_execution_report(
|
||||
suite=args.suite,
|
||||
partition_id=args.partition_id,
|
||||
total_partitions=args.total_partitions,
|
||||
executed_cases=executed_cases,
|
||||
is_standalone=False,
|
||||
case_results=case_results,
|
||||
missing_standalone_estimates=[],
|
||||
standalone_measurements=[],
|
||||
_merge_execution_results(
|
||||
executed_cases, case_results, new_executed_cases, new_case_results
|
||||
)
|
||||
return exit_code
|
||||
# A failing case must not swallow this shard's standalone files: they
|
||||
# are separate pytest runs, and they only share a shard because the
|
||||
# shard count is fixed. --continue-on-error still decides whether a
|
||||
# failing standalone file stops the ones queued behind it.
|
||||
if exit_code != 0 and overall_exit_code == 0:
|
||||
overall_exit_code = exit_code
|
||||
|
||||
standalone_idx = args.partition_id - parametrized_partitions
|
||||
if standalone_idx >= len(standalone_files):
|
||||
if assignment.standalone_files:
|
||||
standalone_estimate = sum(
|
||||
get_standalone_file_est_time(args.suite, standalone_file)[0]
|
||||
for standalone_file in assignment.standalone_files
|
||||
)
|
||||
total_est_time += standalone_estimate
|
||||
print(
|
||||
f"ERROR: Standalone partition index {standalone_idx} exceeds available "
|
||||
f"standalone files ({len(standalone_files)}) for suite '{args.suite}'."
|
||||
f"Running {len(assignment.standalone_files)} standalone file(s) with estimated total "
|
||||
f"{standalone_estimate:.1f}s:"
|
||||
)
|
||||
return 1
|
||||
for standalone_file in assignment.standalone_files:
|
||||
est_time, used_fallback_estimate = get_standalone_file_est_time(
|
||||
args.suite, standalone_file
|
||||
)
|
||||
fallback_suffix = (
|
||||
f", fallback estimate {DEFAULT_STANDALONE_EST_TIME_SECONDS:.1f}s"
|
||||
if used_fallback_estimate
|
||||
else ""
|
||||
)
|
||||
print(
|
||||
f" - standalone: {standalone_file} "
|
||||
f"({est_time:.1f}s{fallback_suffix})"
|
||||
)
|
||||
|
||||
standalone_rel = standalone_files[standalone_idx]
|
||||
print(
|
||||
f"Suite: {args.suite} | Partition: {args.partition_id + 1}/{args.total_partitions} (standalone)"
|
||||
)
|
||||
print(f"Running standalone test file: {Path(standalone_rel).name}")
|
||||
exit_code, executed_cases, case_results, measurement = _run_standalone_file(
|
||||
args.suite,
|
||||
standalone_rel,
|
||||
target_dir,
|
||||
extra_filter=args.filter,
|
||||
)
|
||||
for standalone_file in assignment.standalone_files:
|
||||
exit_code, new_executed_cases, new_case_results, measurement = (
|
||||
_run_standalone_file(
|
||||
args.suite,
|
||||
standalone_file,
|
||||
target_dir,
|
||||
extra_filter=args.filter,
|
||||
)
|
||||
)
|
||||
if measurement["used_fallback_estimate"]:
|
||||
missing_standalone_estimates.append(standalone_file)
|
||||
standalone_measurements.append(measurement)
|
||||
_merge_execution_results(
|
||||
executed_cases,
|
||||
case_results,
|
||||
new_executed_cases,
|
||||
new_case_results,
|
||||
)
|
||||
if exit_code != 0 and overall_exit_code == 0:
|
||||
overall_exit_code = exit_code
|
||||
if exit_code != 0 and not args.continue_on_error:
|
||||
break
|
||||
|
||||
if not assignment.case_ids and not assignment.standalone_files:
|
||||
print(f"No work assigned to partition {args.partition_id}. Exiting success.")
|
||||
|
||||
print(f"Partition estimated total time: {total_est_time:.1f}s")
|
||||
write_execution_report(
|
||||
suite=args.suite,
|
||||
partition_id=args.partition_id,
|
||||
total_partitions=args.total_partitions,
|
||||
executed_cases=executed_cases,
|
||||
is_standalone=True,
|
||||
standalone_file=standalone_rel,
|
||||
is_standalone=False,
|
||||
standalone_file=None,
|
||||
case_results=case_results,
|
||||
missing_standalone_estimates=(
|
||||
[standalone_rel] if measurement["used_fallback_estimate"] else []
|
||||
),
|
||||
standalone_measurements=[measurement],
|
||||
missing_standalone_estimates=missing_standalone_estimates,
|
||||
standalone_measurements=standalone_measurements,
|
||||
)
|
||||
return exit_code
|
||||
return overall_exit_code
|
||||
|
||||
|
||||
def main():
|
||||
|
||||
@@ -17,14 +17,16 @@ import sys
|
||||
from pathlib import Path
|
||||
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.test.partitioning import (
|
||||
PartitionItem,
|
||||
partition_items_by_lpt,
|
||||
)
|
||||
from sglang.multimodal_gen.test.run_suite import (
|
||||
SUITES,
|
||||
PartitionItem,
|
||||
_maybe_pin_update_weights_model_pair,
|
||||
get_case_est_time,
|
||||
get_suite_files_rel,
|
||||
parse_partition_plan,
|
||||
partition_items_by_lpt,
|
||||
)
|
||||
from sglang.multimodal_gen.test.runner.pytest_runner import (
|
||||
collect_test_items,
|
||||
|
||||
@@ -0,0 +1,110 @@
|
||||
"""Shard assignment invariants for the diffusion suites.
|
||||
|
||||
Lanes that hardcode ``--total-partitions`` (AMD) must still schedule the whole
|
||||
suite when standalone files outnumber the shards, and the shards must agree on
|
||||
who runs what without talking to each other.
|
||||
"""
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from sglang.multimodal_gen.test import run_suite
|
||||
from sglang.multimodal_gen.test.partitioning import PartitionItem, assign_partition
|
||||
from sglang.multimodal_gen.test.run_suite import (
|
||||
PartitionAssignment,
|
||||
build_local_partition_assignment,
|
||||
)
|
||||
from sglang.multimodal_gen.test.server.gpu_cases import (
|
||||
PARAMETRIZED_CASE_GROUPS,
|
||||
STANDALONE_FILES,
|
||||
)
|
||||
|
||||
|
||||
def _items(*est_times: float) -> list[PartitionItem]:
|
||||
return [
|
||||
PartitionItem(kind="case", item_id=f"case-{idx}", est_time=est_time)
|
||||
for idx, est_time in enumerate(est_times)
|
||||
]
|
||||
|
||||
|
||||
def _expected_work(suite: str) -> tuple[list[str], list[str]]:
|
||||
case_ids = [
|
||||
case.id
|
||||
for _, case_group in PARAMETRIZED_CASE_GROUPS[suite]
|
||||
for case in case_group
|
||||
]
|
||||
return case_ids, list(STANDALONE_FILES.get(suite, []))
|
||||
|
||||
|
||||
def test_assign_partition_covers_every_item_once():
|
||||
items = _items(300.0, 120.0, 600.0, 60.0, 180.0)
|
||||
assigned = [
|
||||
item.item_id for rank in range(3) for item in assign_partition(items, rank, 3)
|
||||
]
|
||||
assert sorted(assigned) == sorted(item.item_id for item in items)
|
||||
|
||||
|
||||
def test_assign_partition_is_empty_outside_the_shard_range():
|
||||
items = _items(300.0, 120.0)
|
||||
assert assign_partition(items, 5, 3) == []
|
||||
assert assign_partition(items, -1, 3) == []
|
||||
assert assign_partition(items, 0, 0) == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize("suite", sorted(PARAMETRIZED_CASE_GROUPS))
|
||||
@pytest.mark.parametrize("total_partitions", [1, 2, 3, 4, 8])
|
||||
def test_suite_is_fully_scheduled_for_any_shard_count(suite, total_partitions):
|
||||
expected_case_ids, expected_standalone_files = _expected_work(suite)
|
||||
|
||||
scheduled_case_ids: list[str] = []
|
||||
scheduled_standalone_files: list[str] = []
|
||||
for partition_id in range(total_partitions):
|
||||
assignment = build_local_partition_assignment(
|
||||
suite=suite,
|
||||
partition_id=partition_id,
|
||||
total_partitions=total_partitions,
|
||||
)
|
||||
scheduled_case_ids.extend(assignment.case_ids)
|
||||
scheduled_standalone_files.extend(assignment.standalone_files)
|
||||
|
||||
# More standalone files than shards used to abort the run; they now share
|
||||
# shards with the parametrized cases instead.
|
||||
assert sorted(scheduled_case_ids) == sorted(expected_case_ids)
|
||||
assert sorted(scheduled_standalone_files) == sorted(expected_standalone_files)
|
||||
|
||||
|
||||
def test_failing_cases_do_not_skip_the_shards_standalone_files(monkeypatch, tmp_path):
|
||||
"""Standalone files used to own a shard, so cases could not block them."""
|
||||
standalone_rel = "../single_test_file/test_disagg_server.py"
|
||||
executed_standalone = []
|
||||
|
||||
def fake_run_standalone_file(suite, rel, target_dir, extra_filter=None):
|
||||
executed_standalone.append(rel)
|
||||
key = f"standalone:{rel}"
|
||||
return 0, [key], {key: "pass"}, {"used_fallback_estimate": False}
|
||||
|
||||
monkeypatch.setattr(
|
||||
run_suite, "run_pytest", lambda *a, **k: (1, ["a_case"], {"a_case": "fail"})
|
||||
)
|
||||
monkeypatch.setattr(run_suite, "_run_standalone_file", fake_run_standalone_file)
|
||||
monkeypatch.setattr(
|
||||
run_suite, "_get_parametrized_files_for_case_ids", lambda *a, **k: ["a_file.py"]
|
||||
)
|
||||
monkeypatch.setattr(run_suite, "write_execution_report", lambda **kwargs: "")
|
||||
|
||||
args = SimpleNamespace(
|
||||
suite="2-gpu",
|
||||
partition_id=0,
|
||||
total_partitions=2,
|
||||
filter=None,
|
||||
continue_on_error=False,
|
||||
)
|
||||
exit_code = run_suite._run_partition_assignment(
|
||||
args,
|
||||
tmp_path,
|
||||
PartitionAssignment(case_ids=["a_case"], standalone_files=[standalone_rel]),
|
||||
)
|
||||
|
||||
assert executed_standalone == [standalone_rel]
|
||||
assert exit_code == 1
|
||||
Reference in New Issue
Block a user