From 4ccf245801be2d230de0c77e6229e8c1cc8ae670 Mon Sep 17 00:00:00 2001 From: tactino <18781106300@163.com> Date: Sun, 27 Sep 2026 17:49:22 -0400 Subject: [PATCH 1/2] fix: GAE ends an episode at the step that ended it, not one later Both servers call feedback for step t with terminated/truncated set to step t-1's outcome and next_terminated/next_truncated to step t's, and chain an episode's first frame to the last frame of the episode before. A frame carrying `terminated` therefore begins an episode. Until 48042e5 the GAE recursion cut a linked frame with the successor's dones[next_idx] - step t's own outcome. 48042e5 made it read the frame's own dones[step] (terminated[step]/truncated[step]) - step t-1's - so since then the step that ended an episode bootstrapped from the next episode's first state and kept accumulating its advantages, and each episode's first step was cut off from the rest. It also made a chain end carrying `dones` count as terminal and drop its bootstrap value. Both branches now read whether a frame's own transition ended from next_*. Three 3-step episodes of reward 1 at zero value, gamma = lambda = 1: main gave 4 3 2 1 3 2 1 2 1; the answer is 3 2 1 three times. --- src/plugrl_server/buffer/rollout_buffer.py | 32 ++--- tests/test_gae_episode_boundary.py | 151 +++++++++++++++++++++ 2 files changed, 164 insertions(+), 19 deletions(-) create mode 100644 tests/test_gae_episode_boundary.py diff --git a/src/plugrl_server/buffer/rollout_buffer.py b/src/plugrl_server/buffer/rollout_buffer.py index 0e73c9e..ffae014 100644 --- a/src/plugrl_server/buffer/rollout_buffer.py +++ b/src/plugrl_server/buffer/rollout_buffer.py @@ -340,22 +340,22 @@ def compute_advantages_and_returns( for step in reversed(range(self.idx)): next_idx = self.next_indices[step] + # Whether this frame's own transition ended its episode. The + # servers pass a step's `terminated`/`truncated` as the step + # before it's outcome - a frame carrying them begins an episode - + # and `next_*` as its own, and they chain the first frame of an + # episode to the last of the one before. Reading the frame's own + # `dones` here, as this did from 48042e5 on, ended every episode + # one step late. + if self.treat_truncated_as_done: + terminated = float(self.next_done[step]) + truncated = 0.0 + else: + terminated = float(self.next_terminated[step]) + truncated = float(self.next_truncated[step]) if next_idx == 0: next_values = self.last_values[step] next_gae_lam = 0 - if self.treat_truncated_as_done: - if self.dones[step]: - terminated = float(self.dones[step]) - else: - terminated = float(self.next_done[step]) - truncated = 0.0 - else: - if self.dones[step]: - terminated = float(self.terminated[step]) - truncated = float(self.truncated[step]) - else: - terminated = float(self.next_terminated[step]) - truncated = float(self.next_truncated[step]) # check last values not overflow or abs extreme large assert abs(next_values).max() < 1e6, ( f"last_values overflow: {next_values}" @@ -363,12 +363,6 @@ def compute_advantages_and_returns( else: next_values = self.values[next_idx] next_gae_lam = self.advantages[next_idx] - if self.treat_truncated_as_done: - terminated = float(self.dones[step]) - truncated = 0.0 - else: - terminated = float(self.terminated[step]) - truncated = float(self.truncated[step]) next_non_terminal = 1.0 - terminated trunc_mask = 1.0 - truncated delta = ( diff --git a/tests/test_gae_episode_boundary.py b/tests/test_gae_episode_boundary.py new file mode 100644 index 0000000..36476d1 --- /dev/null +++ b/tests/test_gae_episode_boundary.py @@ -0,0 +1,151 @@ +"""GAE has to end an episode at the step that ended it, not one step later. + +Both servers (`websocket_agent_server`, `ray_agent_server`) call `feedback` +for step t with + + terminated, truncated = step t-1's outcome + next_terminated, next_truncated = step t's outcome + prev_node = the node returned for step t-1 + +and never reset `prev_node` at an episode boundary, so the first frame of an +episode is linked to the last frame of the one before it. A frame carrying +`terminated` therefore *begins* an episode; whether its own transition ended +one is in `next_*`. + +`GAEBuffer` read the wrong pair for linked frames. Until 48042e5 it cut the +recursion with the successor's `dones[next_idx]`, which under this +convention is step t's own outcome. 48042e5 changed that to the frame's own +`dones[step]` - step t-1's - so since then every episode has been cut one +step late: the step that fell over in Hopper bootstrapped from the value of +the next episode's first state, and each episode's first step was cut off +from everything after it. FPO, DPPO and anything else on `GAEBuffer` all +computed advantages this way. + +These feed the buffer exactly as the servers do and check the advantages +against ones worked out by hand. +""" + +from __future__ import annotations + +import numpy as np +import pytest +import torch + +from plugrl_server.buffer.rollout_buffer import GAEBuffer + + +def _train_state(value: float = 0.0) -> dict: + return dict( + obs=np.zeros((1, 2), np.float32), + action=np.zeros((1, 1), np.float32), + logprob=np.zeros((1,), np.float32), + value=np.full((1,), value, np.float32), + ) + + +class _ConstantValue: + def __init__(self, value: float) -> None: + self.value = value + + def get_value(self, obs: dict) -> torch.Tensor: + return torch.full((np.asarray(obs["v"]).shape[0],), self.value) + + +def _as_the_server_feeds_it( + outcomes: list[str], + *, + treat_truncated_as_done: bool, + bootstrap: float = 0.0, +) -> GAEBuffer: + """One env, reward 1 at every step, every stored value 0, gamma = lambda = 1. + + `outcomes[t]` is how step t ended: "" (it did not), "terminated" or + "truncated". The client resets after either, and the server keeps the + chain going across the reset. + """ + buffer = GAEBuffer( + len(outcomes), + _train_state(), + gamma=1.0, + gae_lambda=1.0, + treat_truncated_as_done=treat_truncated_as_done, + ) + prev_node: tuple = (-1, "") + previous = "" + for outcome in outcomes: + prev_node = buffer.add_frame( + prev_node=prev_node, + train_state=_train_state(), + reward=1.0, + terminated=previous == "terminated", + truncated=previous == "truncated", + last_value=None, + next_terminated=outcome == "terminated", + next_truncated=outcome == "truncated", + ) + buffer.add_next_obs_value_request( + obs={"v": np.zeros((1, 2), np.float32)}, end_node=prev_node + ) + previous = outcome + buffer.compute_advantages_and_returns( + policy=_ConstantValue(bootstrap), batch_size=4 + ) + return buffer + + +THREE_EPISODES = ["", "", "terminated"] * 3 + + +@pytest.mark.parametrize("treat_truncated_as_done", [True, False]) +def test_each_terminated_episode_is_its_own(treat_truncated_as_done): + """Three steps of reward 1 and a fall: 3, 2, 1 in every episode. + + Before the fix: 4, 3, 2, 1, 3, 2, 1, 2, 1 - the first fall carried on + into the second episode, and the last episode lost its first step. + """ + buffer = _as_the_server_feeds_it( + THREE_EPISODES, treat_truncated_as_done=treat_truncated_as_done + ) + + np.testing.assert_array_equal(buffer.advantages[:9], [3, 2, 1] * 3) + + +def test_a_truncated_episode_ends_like_a_terminated_one_when_told_to(): + buffer = _as_the_server_feeds_it( + ["", "", "truncated"] * 3, treat_truncated_as_done=True + ) + + np.testing.assert_array_equal(buffer.advantages[:9], [3, 2, 1] * 3) + + +def test_otherwise_the_truncated_step_is_left_out_and_nothing_crosses_it(): + """With treat_truncated_as_done off, the step that was truncated has its + delta masked to zero - GAEBuffer's treatment - and the episode before it + still does not reach into the next.""" + buffer = _as_the_server_feeds_it( + ["", "", "truncated"] * 3, treat_truncated_as_done=False + ) + + np.testing.assert_array_equal(buffer.advantages[:9], [2, 1, 0] * 3) + + +def test_a_chain_that_runs_off_the_buffer_bootstraps_even_if_it_just_began(): + """The buffer fills on the first step of a new episode. That step did not + end anything, so it bootstraps from the value of what came after it. + + Before the fix a frame carrying `terminated` at the end of a chain was + taken as terminal itself, and its bootstrap value was dropped. + """ + buffer = _as_the_server_feeds_it( + ["", "", "terminated", ""], treat_truncated_as_done=True, bootstrap=5.0 + ) + + np.testing.assert_array_equal(buffer.advantages[:4], [3, 2, 1, 6]) + + +def test_a_chain_that_runs_off_the_buffer_mid_episode_bootstraps(): + buffer = _as_the_server_feeds_it( + ["", "", "terminated", "", ""], treat_truncated_as_done=False, bootstrap=5.0 + ) + + np.testing.assert_array_equal(buffer.advantages[:5], [3, 2, 1, 7, 6]) From 3906e45c95799120231fc2cac3059e55b96731cd Mon Sep 17 00:00:00 2001 From: tactino <18781106300@163.com> Date: Sun, 27 Sep 2026 18:48:43 -0400 Subject: [PATCH 2/2] test: value_loss_coeff changes nothing to float32 rounding, not to the bit The factors are powers of two, but Adam's eps is added unscaled and leaves a trace in the last bits of elements whose gradient is within a few orders of it. After this branch changed the toy's returns, one parameter came out a few bits apart on CI and identical on Windows and on a Linux workstation, all torch 2.7.1. --- tests/test_fpo_value_loss_coeff.py | 39 ++++++++++++++++++++---------- 1 file changed, 26 insertions(+), 13 deletions(-) diff --git a/tests/test_fpo_value_loss_coeff.py b/tests/test_fpo_value_loss_coeff.py index b717e11..4ae2565 100644 --- a/tests/test_fpo_value_loss_coeff.py +++ b/tests/test_fpo_value_loss_coeff.py @@ -17,9 +17,10 @@ trunk and the disjoint critic and nothing else Measured on the bandit: coefficients 0.25, 1.0 and 4.0 give byte-identical -results across five seeds. The apparent effect below about 0.05 is Adam's -epsilon becoming comparable to the scaled second moment, not the critic -learning differently. +results across five seeds on the machine that measured it - identical to +float32 rounding is what holds everywhere (see the test below). The apparent +effect below about 0.05 is Adam's epsilon becoming comparable to the scaled +second moment, not the critic learning differently. This is not a behaviour change. Redefining the knob would silently alter what it means for anyone reading it, and nothing can depend on it today because it @@ -70,20 +71,32 @@ def test_the_critic_receives_no_gradient_from_the_policy_loss(): def test_scaling_the_value_loss_changes_nothing(): - """0.25, 1.0 and 4.0 give the same weights, to the bit.""" + """0.25, 1.0 and 4.0 give the same weights, to float32 rounding. + + Not to the bit. The factors are powers of two, so the scaled gradients + and moments are exact, but `eps` is added to sqrt(v) unscaled and leaves + a trace in the last bits of any element whose gradient is within a few + orders of it. Which elements, and whether the trace survives rounding, + depends on the platform's kernels: after #86 changed this toy's returns, + one parameter came out a few bits apart on CI and identical on Windows + and on a Linux workstation, all on torch 2.7.1. + """ baseline = _run(0.25) for coeff in (1.0, 4.0): other = _run(coeff) assert len(baseline) == len(other) - differing = [ - i for i, (a, b) in enumerate(zip(baseline, other)) if not torch.equal(a, b) - ] - assert not differing, ( - f"value_loss_coeff={coeff} changed {len(differing)} parameters " - "against 0.25. If this starts failing, either the critic is no " - "longer disjoint or the optimizer is no longer scale-invariant, " - "and the docstring above needs rewriting rather than the test." - ) + for i, (a, b) in enumerate(zip(baseline, other)): + torch.testing.assert_close( + b, + a, + msg=lambda m, i=i: ( + f"value_loss_coeff={coeff} moved parameter {i} against 0.25 " + f"beyond float32 rounding: {m} If this fails, either the " + "critic is no longer disjoint or the optimizer is no " + "longer scale-invariant, and the docstring above needs " + "rewriting rather than the test." + ), + ) def test_a_learning_rate_is_not_cancelled():