Skip to content

Commit 47d3251

Browse files
Simplify benchmark success smoke checks
1 parent 4968f55 commit 47d3251

1 file changed

Lines changed: 9 additions & 14 deletions

File tree

scripts/benchmarks/test/test_benchmark_smoke.py

Lines changed: 9 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -43,14 +43,13 @@ def _run(command: list[str]) -> None:
4343
"max_iterations",
4444
"formatter",
4545
"expect_reward_series",
46-
"expect_success_rate",
4746
"expect_checkpoint",
4847
),
4948
[
50-
("rl_games", 512, 20, "schema", True, True, False),
51-
("rsl_rl", 16, 20, "schema,omniperf", True, True, False),
52-
("sb3", 16, 70, "schema", False, True, True),
53-
("skrl", 16, 20, "schema", True, True, False),
49+
("rl_games", 512, 20, "schema", True, False),
50+
("rsl_rl", 16, 20, "schema,omniperf", True, False),
51+
("sb3", 16, 70, "schema", False, True),
52+
("skrl", 16, 20, "schema", True, False),
5453
],
5554
)
5655
def test_training_and_play_write_bundles(
@@ -61,7 +60,6 @@ def test_training_and_play_write_bundles(
6160
max_iterations: int,
6261
formatter: str,
6362
expect_reward_series: bool,
64-
expect_success_rate: bool,
6563
expect_checkpoint: bool,
6664
):
6765
"""Each RL library trains and plays a policy with benchmark output."""
@@ -105,14 +103,11 @@ def test_training_and_play_write_bundles(
105103
assert training_data["learning"]["reward"]["final_ema"] is not None
106104
if expect_reward_series:
107105
assert len(training_data["learning"]["reward"]["series_per_iter"]) >= 1
108-
if expect_success_rate:
109-
assert training_data["success_rate"] is not None
110-
success_curve = training_data["learning"]["success_rate"]
111-
assert success_curve is not None
112-
assert success_curve["series_per_iter"]
113-
assert success_curve["final_raw"] == pytest.approx(success_curve["series_per_iter"][-1])
114-
else:
115-
assert training_data["learning"]["success_rate"] is None
106+
assert training_data["success_rate"] is not None
107+
success_curve = training_data["learning"]["success_rate"]
108+
assert success_curve is not None
109+
assert success_curve["series_per_iter"]
110+
assert success_curve["final_raw"] == pytest.approx(success_curve["series_per_iter"][-1])
116111
if expect_checkpoint:
117112
assert Path(training_data["checkpoint_path"]).is_file()
118113

0 commit comments

Comments
 (0)