From 62f2f35406548eb1a137f1e29189c7473e346ec5 Mon Sep 17 00:00:00 2001 From: Gotham-Zolio <18781106300@163.com> Date: Sun, 13 Sep 2026 12:58:12 -0400 Subject: [PATCH] fix: report the rollout's timings when the server ends the run `rollout()` logs a final timing summary when its loop finishes. The loop also ends the other way: the server takes the last step it was asked for, closes with plugrl-server-stop, and the agent raises ServerStopped out of an infer. `run()` already treats that as the happy path - the comment there records that it used to escape as an unhandled exception and end a successful run in a traceback. It was not happy enough. The exception left rollout() before the summary, so a run that ended exactly as intended reported no timings at all. What went with them matters more than the timings: env_steps is the only client-side record of how far the rollout got, and it is what a server's global_step has to be reconciled against. Without it there is no way to tell whether the two sides agree about how much work happened. Found while building E9, whose whole point is that reconciliation across several clients. The harness read the count from a log line that a normal finish never printed, so every run looked like it had taken zero steps. The summary now runs in a finally, so it belongs to the rollout however the rollout ends. Tests drive the real loop rather than a stand-in, because the defect was in an exit path and a mocked rollout has none. Both fail against the previous behaviour: one finds no summary at all, the other finds no step count in it. Co-Authored-By: Claude Opus 5 (1M context) --- src/plugrl_env_client/runner/rollout.py | 232 ++++++++++--------- tests/test_rollout_reports_on_server_stop.py | 74 ++++++ 2 files changed, 194 insertions(+), 112 deletions(-) create mode 100644 tests/test_rollout_reports_on_server_stop.py diff --git a/src/plugrl_env_client/runner/rollout.py b/src/plugrl_env_client/runner/rollout.py index be076b8..403e59a 100644 --- a/src/plugrl_env_client/runner/rollout.py +++ b/src/plugrl_env_client/runner/rollout.py @@ -111,120 +111,128 @@ def log_timing_summary(final: bool = False) -> None: ) finished_episodes = 0 - while finished_episodes < num_episodes: - need_infer = np.nonzero(plan_pos >= plan_len)[0] - if need_infer.size: - infer_obs_pack_started_at = time.perf_counter() - obs_msg = dataclasses.asdict(_select_obs(obs, need_infer)) - timing.infer_obs_pack += time.perf_counter() - infer_obs_pack_started_at - - infer_wait_started_at = time.perf_counter() - action_chunk = agent.infer( - obs_msg, - env_indices=need_infer, - step_ids=step_id[need_infer], - )["action"] - timing.infer_wait += time.perf_counter() - infer_wait_started_at - timing.infer_calls += 1 - - steps = replan_steps or len(action_chunk) - if len(action_chunk) < steps: - raise ValueError( - f"replan_steps={replan_steps} exceeds predicted steps={len(action_chunk)}" - ) + # The loop also ends when the server says it is done, which raises + # ServerStopped out of an agent call. That is the happy path - the + # algorithm took every step it was asked for - and it used to skip the + # summary below, so a run that finished normally reported no timings at + # all and left nothing to reconcile the client's step count against the + # server's. The summary belongs to the rollout either way. + try: + while finished_episodes < num_episodes: + need_infer = np.nonzero(plan_pos >= plan_len)[0] + if need_infer.size: + infer_obs_pack_started_at = time.perf_counter() + obs_msg = dataclasses.asdict(_select_obs(obs, need_infer)) + timing.infer_obs_pack += time.perf_counter() - infer_obs_pack_started_at + + infer_wait_started_at = time.perf_counter() + action_chunk = agent.infer( + obs_msg, + env_indices=need_infer, + step_ids=step_id[need_infer], + )["action"] + timing.infer_wait += time.perf_counter() - infer_wait_started_at + timing.infer_calls += 1 + + steps = replan_steps or len(action_chunk) + if len(action_chunk) < steps: + raise ValueError( + f"replan_steps={replan_steps} exceeds predicted steps={len(action_chunk)}" + ) + + a = np.asarray(action_chunk[:steps], dtype=expected_action_dtype) + if a.shape[:2] != (steps, need_infer.size): + raise ValueError( + f"Expected action shape ({steps}, {need_infer.size}, da), got {a.shape}" + ) + if a.shape[2:] != expected_action_shape: + raise ValueError( + f"Expected action shape tail {expected_action_shape}, got {a.shape[2:]}" + ) + + need_capacity = max(int(action_plan.shape[1]), int(steps)) + need_shape = (num_envs, need_capacity) + expected_action_shape + if action_plan.shape != need_shape: + new_plan = np.empty(need_shape, dtype=expected_action_dtype) + cap = min(int(action_plan.shape[1]), need_capacity) + if cap: + new_plan[:, :cap, ...] = action_plan[:, :cap, ...] + action_plan = new_plan + + action_plan[need_infer, :steps, ...] = a.swapaxes(0, 1) + plan_pos[need_infer] = 0 + plan_len[need_infer] = steps + + if np.any(plan_pos >= plan_len): + raise RuntimeError("Action plan is not ready for all envs") + + actions = action_plan[np.arange(num_envs), plan_pos] + env_step_started_at = time.perf_counter() + obs, reward, terminated, truncated, info = env.step(actions) + timing.env_step += time.perf_counter() - env_step_started_at + timing.env_steps += num_envs + if recorder is not None: + recorder.on_step(obs, reward, terminated, truncated, info) + + plan_pos += 1 + + reward = np.asarray(reward, dtype=np.float32) + terminated = np.asarray(terminated, dtype=np.bool_) + truncated = np.asarray(truncated, dtype=np.bool_) + done = np.logical_or(terminated, truncated) + + chunk_reward += reward + + done_indices = np.nonzero(done)[0] + if done_indices.size: + plan_pos[done_indices] = plan_len[done_indices] - a = np.asarray(action_chunk[:steps], dtype=expected_action_dtype) - if a.shape[:2] != (steps, need_infer.size): - raise ValueError( - f"Expected action shape ({steps}, {need_infer.size}, da), got {a.shape}" + feedback_indices = np.nonzero(plan_pos >= plan_len)[0] + if feedback_indices.size: + feedback_started_at = time.perf_counter() + feedback_obs_pack_started_at = time.perf_counter() + feedback_obs = dataclasses.asdict(_select_obs(obs, feedback_indices)) + timing.feedback_obs_pack += ( + time.perf_counter() - feedback_obs_pack_started_at ) - if a.shape[2:] != expected_action_shape: - raise ValueError( - f"Expected action shape tail {expected_action_shape}, got {a.shape[2:]}" + feedback_info_pack_started_at = time.perf_counter() + feedback_info = _select_info(info, feedback_indices, num_envs=num_envs) + timing.feedback_info_pack += ( + time.perf_counter() - feedback_info_pack_started_at ) + agent.feedback( + obs=feedback_obs, + rewards=chunk_reward[feedback_indices], + terminated=terminated[feedback_indices], + truncated=truncated[feedback_indices], + info=feedback_info, + env_indices=feedback_indices, + step_ids=step_id[feedback_indices], + ) + timing.feedback += time.perf_counter() - feedback_started_at + timing.feedback_calls += 1 + chunk_reward[feedback_indices] = 0.0 + step_id[feedback_indices] += 1 - need_capacity = max(int(action_plan.shape[1]), int(steps)) - need_shape = (num_envs, need_capacity) + expected_action_shape - if action_plan.shape != need_shape: - new_plan = np.empty(need_shape, dtype=expected_action_dtype) - cap = min(int(action_plan.shape[1]), need_capacity) - if cap: - new_plan[:, :cap, ...] = action_plan[:, :cap, ...] - action_plan = new_plan - - action_plan[need_infer, :steps, ...] = a.swapaxes(0, 1) - plan_pos[need_infer] = 0 - plan_len[need_infer] = steps - - if np.any(plan_pos >= plan_len): - raise RuntimeError("Action plan is not ready for all envs") - - actions = action_plan[np.arange(num_envs), plan_pos] - env_step_started_at = time.perf_counter() - obs, reward, terminated, truncated, info = env.step(actions) - timing.env_step += time.perf_counter() - env_step_started_at - timing.env_steps += num_envs - if recorder is not None: - recorder.on_step(obs, reward, terminated, truncated, info) - - plan_pos += 1 - - reward = np.asarray(reward, dtype=np.float32) - terminated = np.asarray(terminated, dtype=np.bool_) - truncated = np.asarray(truncated, dtype=np.bool_) - done = np.logical_or(terminated, truncated) - - chunk_reward += reward - - done_indices = np.nonzero(done)[0] - if done_indices.size: - plan_pos[done_indices] = plan_len[done_indices] - - feedback_indices = np.nonzero(plan_pos >= plan_len)[0] - if feedback_indices.size: - feedback_started_at = time.perf_counter() - feedback_obs_pack_started_at = time.perf_counter() - feedback_obs = dataclasses.asdict(_select_obs(obs, feedback_indices)) - timing.feedback_obs_pack += ( - time.perf_counter() - feedback_obs_pack_started_at - ) - feedback_info_pack_started_at = time.perf_counter() - feedback_info = _select_info(info, feedback_indices, num_envs=num_envs) - timing.feedback_info_pack += ( - time.perf_counter() - feedback_info_pack_started_at - ) - agent.feedback( - obs=feedback_obs, - rewards=chunk_reward[feedback_indices], - terminated=terminated[feedback_indices], - truncated=truncated[feedback_indices], - info=feedback_info, - env_indices=feedback_indices, - step_ids=step_id[feedback_indices], - ) - timing.feedback += time.perf_counter() - feedback_started_at - timing.feedback_calls += 1 - chunk_reward[feedback_indices] = 0.0 - step_id[feedback_indices] += 1 - - if done_indices.size: - if recorder is not None: - recorder.on_episode_done(done_indices, obs, info) - finished_episodes += done_indices.size - # Only reset if the loop is going to use what comes back. A reset - # after the last episode costs a simulator a wasted rollout, and - # costs a real robot a pointless move back to its home pose - for - # an observation nothing will ever read. - if finished_episodes < num_episodes: - obs, info = env.reset(options={"reset_indices": done_indices}) + if done_indices.size: if recorder is not None: - recorder.on_reset(obs, info, reset_indices=done_indices) - step_id[done_indices] = 0 - now = time.perf_counter() - if now - last_timing_log_at >= 30.0: - log_timing_summary(final=False) - last_timing_log_at = now - - log_timing_summary(final=True) - if recorder is not None: - recorder.record_timing(timing) + recorder.on_episode_done(done_indices, obs, info) + finished_episodes += done_indices.size + # Only reset if the loop is going to use what comes back. A reset + # after the last episode costs a simulator a wasted rollout, and + # costs a real robot a pointless move back to its home pose - for + # an observation nothing will ever read. + if finished_episodes < num_episodes: + obs, info = env.reset(options={"reset_indices": done_indices}) + if recorder is not None: + recorder.on_reset(obs, info, reset_indices=done_indices) + step_id[done_indices] = 0 + now = time.perf_counter() + if now - last_timing_log_at >= 30.0: + log_timing_summary(final=False) + last_timing_log_at = now + + finally: + log_timing_summary(final=True) + if recorder is not None: + recorder.record_timing(timing) diff --git a/tests/test_rollout_reports_on_server_stop.py b/tests/test_rollout_reports_on_server_stop.py new file mode 100644 index 0000000..5a7dec9 --- /dev/null +++ b/tests/test_rollout_reports_on_server_stop.py @@ -0,0 +1,74 @@ +"""A run the server ended still has to report what it did. + +`rollout()` logs a final timing summary when its loop finishes. The loop also +ends the other way: the server takes the last step it was asked for, closes +with `plugrl-server-stop`, and the agent raises `ServerStopped` out of an +infer. `run()` already treats that as the happy path rather than a failure. + +It was not happy enough. The exception left `rollout()` before the summary, +so a run that ended exactly as intended reported no timings at all - and with +them went the only client-side record of how many environment steps it took, +which is what a server's `global_step` has to be reconciled against. + +These drive the real rollout loop rather than a stand-in for it, because the +defect was in the loop's exit path and a mocked rollout cannot have one. +""" + +from __future__ import annotations + +import pytest +from loguru import logger + +from plugrl_env_client.agent.websocket_env_client_agent import ServerStopped +from plugrl_env_client.runner.rollout import rollout + +from test_protocol_alternation import _StaggeredEnv, _WireRecorder + + +class _StopsAfter(_WireRecorder): + """A server that ends the run after `n` infers, the way a finished one does.""" + + def __init__(self, horizon, stop_after): + super().__init__(horizon=horizon) + self.stop_after = stop_after + self.infers = 0 + + def infer(self, obs, *, env_indices, step_ids): + self.infers += 1 + if self.infers > self.stop_after: + raise ServerStopped("Server requested env client shutdown.") + return super().infer(obs, env_indices=env_indices, step_ids=step_ids) + + +def _messages_from(stop_after): + captured: list[str] = [] + sink = logger.add(lambda m: captured.append(m.record["message"]), level="INFO") + try: + with pytest.raises(ServerStopped): + rollout( + _StaggeredEnv(episode_lengths=(2, 6)), + _StopsAfter(horizon=1, stop_after=stop_after), + num_episodes=100, + replan_steps=1, + num_envs=2, + seed=0, + ) + finally: + logger.remove(sink) + return captured + + +def test_a_server_stop_still_reports_the_final_summary(): + messages = _messages_from(stop_after=5) + + finals = [m for m in messages if m.startswith("Final rollout timing summary")] + assert finals, f"no final summary was logged; got {messages}" + + +def test_the_summary_carries_the_steps_actually_taken(): + """The count is the point: it is what reconciles against the server.""" + messages = _messages_from(stop_after=5) + + final = next(m for m in messages if m.startswith("Final rollout timing summary")) + steps = int(final.split("env_steps=")[1].split()[0]) + assert steps > 0, f"summary reported {steps} steps: {final}"