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
32 changes: 13 additions & 19 deletions src/plugrl_server/buffer/rollout_buffer.py
Original file line number Diff line number Diff line change
Expand Up @@ -340,35 +340,29 @@ 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}"
)
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 = (
Expand Down
39 changes: 26 additions & 13 deletions tests/test_fpo_value_loss_coeff.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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():
Expand Down
151 changes: 151 additions & 0 deletions tests/test_gae_episode_boundary.py
Original file line number Diff line number Diff line change
@@ -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])
Loading