Skip to content
Merged
Show file tree
Hide file tree
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
232 changes: 120 additions & 112 deletions src/plugrl_env_client/runner/rollout.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
74 changes: 74 additions & 0 deletions tests/test_rollout_reports_on_server_stop.py
Original file line number Diff line number Diff line change
@@ -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}"
Loading