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}"