From 1ef3bc01049697ca9553e44204add7942dd08bc4 Mon Sep 17 00:00:00 2001 From: tactino <18781106300@163.com> Date: Sun, 27 Sep 2026 17:35:18 -0400 Subject: [PATCH 1/5] gaussian-policy and ppo: CleanRL's continuous-action PPO The coverage figure pairs expressive policies with algorithms built for them and has no row for the baseline they are measured against. This adds it, written to CleanRL's ppo_continuous_action.py with every default: gaussian-policy tanh MLPs of 64 for the mean and the value, a log std independent of the observation, orthogonal init (gain sqrt 2; 0.01 on the mean's head, 1.0 on the value's). ppo 2048-step rollouts, 10 epochs of 64-sample minibatches, clip 0.2 on the ratio and the value, per-minibatch advantage normalisation, one Adam at eps 1e-5, gradient norm 0.5, lr 3e-4 annealed linearly over 488 iterations. CleanRL clips actions and normalises observations and rewards in gymnasium wrappers, and PlugRL's client wraps nothing, so they happen on the server: the env receives the clipped sample and the buffer keeps the drawn one; observation statistics are buffers updated after each learn, so collection and learning share one normalisation; PPOBuffer divides rewards by the running deviation of the discounted return and clips at 10, restarting the return with each episode. tests/test_gaussian_ppo.py checks each piece and that the pair learns the bandit FPO's test uses and keeps it on five seeds (last/best at most 1.12). --- README.md | 26 + src/plugrl_server/algorithm/__init__.py | 2 + src/plugrl_server/algorithm/ppo/__init__.py | 0 src/plugrl_server/algorithm/ppo/ppo.py | 296 +++++++++++ src/plugrl_server/algorithm/ppo/ppo_buffer.py | 92 ++++ src/plugrl_server/algorithm/ppo/ppo_config.py | 47 ++ src/plugrl_server/policy/__init__.py | 2 + src/plugrl_server/policy/gaussian/__init__.py | 7 + .../policy/gaussian/gaussian_policy.py | 188 +++++++ tests/test_gaussian_ppo.py | 463 ++++++++++++++++++ 10 files changed, 1123 insertions(+) create mode 100644 src/plugrl_server/algorithm/ppo/__init__.py create mode 100644 src/plugrl_server/algorithm/ppo/ppo.py create mode 100644 src/plugrl_server/algorithm/ppo/ppo_buffer.py create mode 100644 src/plugrl_server/algorithm/ppo/ppo_config.py create mode 100644 src/plugrl_server/policy/gaussian/__init__.py create mode 100644 src/plugrl_server/policy/gaussian/gaussian_policy.py create mode 100644 tests/test_gaussian_ppo.py diff --git a/README.md b/README.md index c262871..4214016 100644 --- a/README.md +++ b/README.md @@ -126,10 +126,15 @@ We use `uv` to manage dependencies and development environments. - `dppo` - DPPO (Diffusion Policy Policy Optimization). No extras needed, and it runs against `fpo-policy`. - `dppo-dist` - Distributed DPPO (experimental) +- `ppo` - PPO as CleanRL's `ppo_continuous_action.py` runs it, every + default included. Drives `gaussian-policy`; no extras needed. - `eval` - Evaluation only, no learning **Policies:** - `fpo-policy` - Flow policy. Defaults to `obs_dim=17`, `action_dim=6` +- `gaussian-policy` - CleanRL's Gaussian MLP: tanh layers of 64, a log std + that does not depend on the observation. Defaults to `obs_dim=17`, + `action_dim=6` - `dummy-policy` - Outputs random actions (for testing) - `dppo-policy` - DPPO policy (requires `plugrl-server[dppo]` and a checkpoint) - `pi0-policy` - PI0 policy (OpenPI). Needs more than a checkpoint. The @@ -213,6 +218,27 @@ million steps therefore learns exactly once, at the very end - producing a single point rather than a curve. 4096 gives one update per 4096 environment steps. +#### The baseline: a Gaussian policy with PPO + +The pair every other one is measured against, written to CleanRL's +`ppo_continuous_action.py`: a rollout of 2048 steps, ten epochs of 32 +minibatches, clip 0.2, a learning rate of 3e-4 annealed to zero over 488 +iterations (a million steps), observations and rewards normalised. CleanRL +does the normalising and the action clipping in gymnasium wrappers; PlugRL's +client does not wrap, so the policy and the algorithm do it on the server. + +```bash +# Terminal 1: Hopper-v5 has an 11-dimensional observation and 3 actions +python -m plugrl_server.cli gaussian-policy default ppo default \ + --port 8000 --policy.device cpu \ + --policy.obs-dim 11 --policy.action-dim 3 + +# Terminal 2 +python -m plugrl_env_client.cli mujoco-v1 \ + --server-port 8000 --num-envs 1 --num-episodes 100000 \ + --env.name Hopper-v5 --runner.replan-steps 1 --runner.seed 0 +``` + #### Quick Start: Testing with Dummy Components To test the agent-server connection with dummy algorithm and policy: diff --git a/src/plugrl_server/algorithm/__init__.py b/src/plugrl_server/algorithm/__init__.py index 076c413..b520883 100644 --- a/src/plugrl_server/algorithm/__init__.py +++ b/src/plugrl_server/algorithm/__init__.py @@ -16,3 +16,5 @@ import plugrl_server.algorithm.dppo.dppo_config # noqa: F401,E402 import plugrl_server.algorithm.dppo.dppo_dist # noqa: F401,E402 import plugrl_server.algorithm.dppo.dppo_dist_config # noqa: F401,E402 +import plugrl_server.algorithm.ppo.ppo # noqa: F401,E402 +import plugrl_server.algorithm.ppo.ppo_config # noqa: F401,E402 diff --git a/src/plugrl_server/algorithm/ppo/__init__.py b/src/plugrl_server/algorithm/ppo/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/plugrl_server/algorithm/ppo/ppo.py b/src/plugrl_server/algorithm/ppo/ppo.py new file mode 100644 index 0000000..0bc1432 --- /dev/null +++ b/src/plugrl_server/algorithm/ppo/ppo.py @@ -0,0 +1,296 @@ +"""PPO, as CleanRL's `ppo_continuous_action.py` runs it. + +Clipped surrogate at 0.2, a clipped value loss, advantages normalised per +minibatch, one Adam (eps 1e-5) over the actor and the critic together, the +whole gradient clipped at norm 0.5, the learning rate annealed linearly to +zero. It drives `gaussian-policy`, and any policy that offers the same four +things: `get_action_and_runtime_state`, `evaluate_actions`, `get_value` and, +optionally, `update_obs_stats`. +""" + +import math + +import numpy as np +import torch +import torch.nn as nn + +from plugrl_server.algorithm.base_algorithm import BaseAlgorithm +from plugrl_server.algorithm.registration import register_algo +from plugrl_server.algorithm.train_utils import move_batch_to_device +from plugrl_server.buffer.rollout_buffer import ROLLOUT_BUFFER_SCHEMA_VERSION +from plugrl_server.common.checkpoint_manager import Checkpoint +from plugrl_server.common.logging_utils import get_logger +from plugrl_server.policy.gaussian.gaussian_policy import GaussianPolicy +from plugrl_server.policy.state import ( + PolicyRuntimeState, + PolicyTrainState, + to_numpy_state, +) + +from .ppo_buffer import PPOBuffer +from .ppo_config import UID, PPOAlgoConfig + +logger = get_logger(__name__) + + +@register_algo(UID) +class PPOAlgorithm(BaseAlgorithm): + config: PPOAlgoConfig + policy: GaussianPolicy + optimizer: torch.optim.Optimizer + + def __init__(self, config: PPOAlgoConfig, policy: GaussianPolicy): + super().__init__(config, policy) + if not hasattr(policy, "evaluate_actions"): + raise TypeError( + f"ppo needs a policy with evaluate_actions, such as " + f"gaussian-policy; {type(policy).__name__} has none" + ) + self.rollout_buffer = PPOBuffer( + buffer_size=config.buffer_size, + example_train_state=self.example_train_state(batch_size=1), + gamma=config.gamma, + gae_lambda=config.gae_lambda, + normalize_rewards=config.normalize_rewards, + reward_clip=config.reward_clip, + ) + self.global_step = 0 + self.curr_train_itrs = 0 + self.last_saved_itr = 0 + + def init_optimizers(self) -> None: + self.optimizer = torch.optim.Adam( + self.policy.parameters(), lr=self.config.learning_rate, eps=1e-5 + ) + + def infer(self, obs: dict) -> tuple[np.ndarray, PolicyRuntimeState]: + with torch.inference_mode(): + return self.policy.get_action_and_runtime_state(obs) + + def derive_train_state(self, runtime_state: PolicyRuntimeState) -> PolicyTrainState: + numpy_state = to_numpy_state(runtime_state) + if not isinstance(numpy_state, dict): + raise TypeError("ppo requires a mapping-like train_state export.") + return numpy_state + + def feedback( + self, + *, + obs: dict, + runtime_state: PolicyRuntimeState, + train_state: PolicyTrainState = None, + terminated: bool, + truncated: bool, + next_obs: dict, + reward: float, + next_terminated: bool, + next_truncated: bool, + info: dict, + prev_node: tuple, + ) -> tuple[tuple, int, dict]: + assert train_state is not None, "ppo requires train_state for rollout storage." + current_node = self.rollout_buffer.add_frame( + prev_node=prev_node, + train_state=train_state, + reward=reward, + terminated=terminated, + truncated=truncated, + last_value=None, + next_terminated=next_terminated, + next_truncated=next_truncated, + ) + self.rollout_buffer.add_next_obs_value_request( + obs=next_obs, end_node=current_node + ) + if next_terminated or next_truncated: + if "episode" in info and bool(info["episode"].get("mask", True)): + self.record_episode_metrics(info["episode"]) + self.rollout_buffer.finish_rollout(info=info) + self.global_step += 1 + return current_node, self.global_step, {} + + def pre_learn(self) -> None: + self.rollout_buffer.compute_advantages_and_returns( + policy=self.policy, batch_size=self.config.batch_size + ) + + def get_collect_progress_total(self) -> int | None: + return self.config.buffer_size + + def get_collect_progress_completed(self) -> int | None: + return len(self.rollout_buffer) + + def get_learn_progress_total(self) -> int | None: + return self.config.update_epochs * math.ceil( + self.config.buffer_size / self.config.batch_size + ) + + def _learning_rate(self) -> float: + """CleanRL's: (1 - (iteration - 1) / num_iterations) * learning_rate.""" + if not self.config.anneal_lr: + return self.config.learning_rate + frac = 1.0 - self.curr_train_itrs / self.config.train_itrs + return frac * self.config.learning_rate + + def _loss(self, obs, action, oldlogprob, value, advantage, ret): + config = self.config + newlogprob, entropy, newvalue = self.policy.evaluate_actions(obs, action) + logratio = newlogprob - oldlogprob + ratio = logratio.exp() + + if config.norm_adv and advantage.numel() > 1: + advantage = (advantage - advantage.mean()) / (advantage.std() + 1e-8) + pg_loss = torch.max( + -advantage * ratio, + -advantage * torch.clamp(ratio, 1 - config.clip_coef, 1 + config.clip_coef), + ).mean() + + if config.clip_vloss: + v_clipped = value + torch.clamp( + newvalue - value, -config.clip_coef, config.clip_coef + ) + v_loss = ( + 0.5 * torch.max((newvalue - ret) ** 2, (v_clipped - ret) ** 2).mean() + ) + else: + v_loss = 0.5 * ((newvalue - ret) ** 2).mean() + + return pg_loss, v_loss, entropy.mean(), logratio, ratio + + def learn_impl(self) -> tuple[int, dict]: + config = self.config + lr = self._learning_rate() + for group in self.optimizer.param_groups: + group["lr"] = lr + + dataloader = torch.utils.data.DataLoader( + self.rollout_buffer, + batch_size=config.batch_size, + shuffle=True, + drop_last=False, + num_workers=0, + collate_fn=self.rollout_buffer.collate_fn, + ) + rollout_summary = self.rollout_buffer.description() + pg_loss = v_loss = entropy_loss = torch.tensor(0.0) + old_approx_kl = approx_kl = torch.tensor(0.0) + clipfracs: list[float] = [] + grad_norms: list[float] = [] + progress_total = self.get_learn_progress_total() + progress = 0 + + for epoch in range(config.update_epochs): + for batch in dataloader: + obs, action, oldlogprob, _reward, value, advantage, ret = ( + move_batch_to_device(batch, device=self.policy.device) + ) + pg_loss, v_loss, entropy_loss, logratio, ratio = self._loss( + obs, action, oldlogprob, value, advantage, ret + ) + with torch.no_grad(): + old_approx_kl = (-logratio).mean() + approx_kl = ((ratio - 1) - logratio).mean() + clipfracs.append( + ((ratio - 1.0).abs() > config.clip_coef).float().mean().item() + ) + + loss = ( + pg_loss - config.ent_coef * entropy_loss + config.vf_coef * v_loss + ) + self.optimizer.zero_grad() + loss.backward() + grad_norms.append( + float( + nn.utils.clip_grad_norm_( + self.policy.parameters(), config.max_grad_norm + ) + ) + ) + self.optimizer.step() + progress += 1 + self.report_learn_progress(progress, progress_total) + + if config.target_kl is not None and approx_kl > config.target_kl: + logger.info("Stopping after epoch %d: approx_kl %.4f", epoch, approx_kl) + break + + self.curr_train_itrs += 1 + return self.global_step, dict( + models=dict(learning_rate=lr), + losses=dict( + value_loss=v_loss.item(), + policy_loss=pg_loss.item(), + entropy=entropy_loss.item(), + old_approx_kl=old_approx_kl.item(), + approx_kl=approx_kl.item(), + clipfrac=float(np.mean(clipfracs)) if clipfracs else 0.0, + ), + train=dict( + max_grad_norm=max(grad_norms) if grad_norms else 0.0, + train_itrs=float(self.curr_train_itrs), + ), + rollout=rollout_summary, + ) + + def post_learn(self) -> None: + # After learning and before the reset, as DPPO does for fpo-policy: + # this iteration collected and learned under one normalisation, and + # the next collects under the new one. + update = getattr(self.policy, "update_obs_stats", None) + filled = len(self.rollout_buffer) + if update is not None and filled > 0: + stored = self.rollout_buffer.train_state_storage.get_item(slice(0, filled)) + update(torch.as_tensor(stored, device=self.policy.device)) + self.rollout_buffer.reset() + super().post_learn() + + def should_learn(self) -> bool: + return self.rollout_buffer.full() + + def should_stop(self) -> bool: + return self.curr_train_itrs >= self.config.train_itrs + + def should_save(self) -> bool: + return (self.curr_train_itrs % self.config.save_interval == 0) and ( + self.curr_train_itrs > self.last_saved_itr + ) + + def create_checkpoint(self) -> Checkpoint: + self.last_saved_itr = self.curr_train_itrs + rms = self.rollout_buffer.ret_rms + return Checkpoint( + step=self.global_step, + model=self.policy.state_dict(), + optimizer={"adam": self.optimizer.state_dict()}, + meta={ + "train_itrs": self.curr_train_itrs, + "last_saved_itr": self.last_saved_itr, + # The reward scale, which a resume would otherwise relearn. + "return_rms": { + "mean": float(rms.mean), + "var": float(rms.var), + "count": float(rms.count), + }, + "rollout_buffer_schema_version": ROLLOUT_BUFFER_SCHEMA_VERSION, + }, + ) + + def load_checkpoint(self, checkpoint: Checkpoint) -> None: + self.global_step = checkpoint.step + if checkpoint.model is not None: + self.policy.load_state_dict(checkpoint.model) + if checkpoint.optimizer is not None: + self.optimizer.load_state_dict(checkpoint.optimizer["adam"]) + meta = checkpoint.meta + self.curr_train_itrs = meta.get("train_itrs", self.curr_train_itrs) + self.last_saved_itr = meta.get("last_saved_itr", self.last_saved_itr) + if "return_rms" in meta: + rms = self.rollout_buffer.ret_rms + rms.mean = np.float64(meta["return_rms"]["mean"]) + rms.var = np.float64(meta["return_rms"]["var"]) + rms.count = meta["return_rms"]["count"] + logger.info( + "Loaded checkpoint at step %d, train_itrs %d", + self.global_step, + self.curr_train_itrs, + ) diff --git a/src/plugrl_server/algorithm/ppo/ppo_buffer.py b/src/plugrl_server/algorithm/ppo/ppo_buffer.py new file mode 100644 index 0000000..9cfe268 --- /dev/null +++ b/src/plugrl_server/algorithm/ppo/ppo_buffer.py @@ -0,0 +1,92 @@ +import uuid + +import numpy as np + +from plugrl_server.algorithm.dppo.third_party.reward_scaling import RunningMeanStd +from plugrl_server.buffer.rollout_buffer import GAEBuffer +from plugrl_server.policy.base_policy import BasePolicy +from plugrl_server.policy.state import PolicyTrainState + + +class PPOBuffer(GAEBuffer): + """GAE, with rewards scaled the way gymnasium's NormalizeReward scales them. + + Each frame's discounted return is kept as it arrives; at learning time the + running variance takes in the buffer's returns and every reward is divided + by the running deviation, then clipped. NormalizeReward divides each reward + by the deviation as it stood at that step instead; over the same returns + the two differ only in how recent the estimate is. + + The discounted return restarts with each episode, as it does in + NormalizeReward and in DPPO's RunningRewardScaler. `DPPOBuffer` follows + the server's chain of frames, which runs on across episodes, so its return + carries over from one episode into the next. + """ + + def __init__( + self, + buffer_size, + example_train_state: PolicyTrainState, + gamma: float = 0.99, + gae_lambda: float = 0.95, + normalize_rewards: bool = True, + reward_clip: float = 10.0, + epsilon: float = 1e-8, + ): + super().__init__(buffer_size, example_train_state, gamma, gae_lambda) + self.normalize_rewards = normalize_rewards + self.reward_clip = reward_clip + self.epsilon = epsilon + self.ret_rms = RunningMeanStd(shape=()) + self.rets = np.zeros(buffer_size, dtype=np.float64) + + def add_frame( + self, + *, + prev_node: tuple[int, uuid.UUID], + train_state: PolicyTrainState, + reward: float, + terminated: bool, + truncated: bool, + last_value: np.ndarray | None, + next_terminated: bool, + next_truncated: bool, + ) -> tuple[int, uuid.UUID]: + node = super().add_frame( + prev_node=prev_node, + train_state=train_state, + reward=reward, + terminated=terminated, + truncated=truncated, + last_value=last_value, + next_terminated=next_terminated, + next_truncated=next_truncated, + ) + current_idx = node[0] + if current_idx == -1: + return node + prev_idx, prev_signature = prev_node + # `terminated` or `truncated` on a frame means the step before it + # ended an episode, so this frame begins one. + continues = ( + prev_idx != -1 + and prev_signature == self.buffer_signature + and not (terminated or truncated) + ) + self.rets[current_idx] = float(reward) + ( + self.gamma * self.rets[prev_idx] if continues else 0.0 + ) + return node + + def compute_advantages_and_returns( + self, policy: BasePolicy | None = None, batch_size: int = 1 + ): + if self.normalize_rewards and self.idx > 0: + self.ret_rms.update(self.rets[: self.idx]) + scale = np.float32(np.sqrt(self.ret_rms.var + self.epsilon)) + self.rewards[: self.idx] = np.clip( + self.rewards[: self.idx] / scale, -self.reward_clip, self.reward_clip + ) + return super().compute_advantages_and_returns( + policy=policy, batch_size=batch_size + ) diff --git a/src/plugrl_server/algorithm/ppo/ppo_config.py b/src/plugrl_server/algorithm/ppo/ppo_config.py new file mode 100644 index 0000000..4bf0cd3 --- /dev/null +++ b/src/plugrl_server/algorithm/ppo/ppo_config.py @@ -0,0 +1,47 @@ +import dataclasses + +from plugrl_server.algorithm.base_algorithm import BaseAlgoConfig +from plugrl_server.algorithm.registration import register_algo_config + +UID = "ppo" + + +@register_algo_config(UID) +@dataclasses.dataclass +class PPOAlgoConfig(BaseAlgoConfig): + """CleanRL `ppo_continuous_action.py`'s defaults, which it runs on every + MuJoCo task unchanged. + + CleanRL's rollout is `num_envs * num_steps` = 1 x 2048 transitions, split + into `num_minibatches` = 32; here that is `buffer_size` 2048 and + `batch_size` 64, however many environments the client runs. Its + 1,000,000 steps are `train_itrs` 488 iterations of 2048. + """ + + learning_rate: float = 3e-4 + # Linearly to zero across `train_itrs`, as CleanRL's `anneal_lr`. + anneal_lr: bool = True + buffer_size: int = 2048 + batch_size: int = 64 + update_epochs: int = 10 + gamma: float = 0.99 + gae_lambda: float = 0.95 + norm_adv: bool = True + clip_coef: float = 0.2 + clip_vloss: bool = True + ent_coef: float = 0.0 + vf_coef: float = 0.5 + max_grad_norm: float = 0.5 + # Stop an iteration's epochs once the approximate KL passes this. None, + # CleanRL's default, never stops. + target_kl: float | None = None + # gymnasium's NormalizeReward and TransformReward, which CleanRL wraps its + # environments in: rewards divided by the running deviation of the + # discounted return, then clipped. + normalize_rewards: bool = True + reward_clip: float = 10.0 + train_itrs: int = 488 + save_interval: int = 50 + + def __post_init__(self): + self.global_steps = self.train_itrs * self.buffer_size diff --git a/src/plugrl_server/policy/__init__.py b/src/plugrl_server/policy/__init__.py index 687c3b0..3000c9d 100644 --- a/src/plugrl_server/policy/__init__.py +++ b/src/plugrl_server/policy/__init__.py @@ -1,11 +1,13 @@ from . import dummy_policy as dummy_policy from . import dppo as dppo from . import fpo as fpo +from . import gaussian as gaussian from . import openpi as openpi __all__ = [ "dummy_policy", "dppo", "fpo", + "gaussian", "openpi", ] diff --git a/src/plugrl_server/policy/gaussian/__init__.py b/src/plugrl_server/policy/gaussian/__init__.py new file mode 100644 index 0000000..8451a6a --- /dev/null +++ b/src/plugrl_server/policy/gaussian/__init__.py @@ -0,0 +1,7 @@ +from .gaussian_policy import GaussianPolicy as GaussianPolicy +from .gaussian_policy import GaussianPolicyConfig as GaussianPolicyConfig + +__all__ = [ + "GaussianPolicy", + "GaussianPolicyConfig", +] diff --git a/src/plugrl_server/policy/gaussian/gaussian_policy.py b/src/plugrl_server/policy/gaussian/gaussian_policy.py new file mode 100644 index 0000000..32d95d3 --- /dev/null +++ b/src/plugrl_server/policy/gaussian/gaussian_policy.py @@ -0,0 +1,188 @@ +"""A Gaussian MLP policy, as CleanRL's `ppo_continuous_action.py` builds it. + +The baseline the expressive policies are measured against: a tanh MLP for the +mean, a log standard deviation that does not depend on the observation, and a +separate tanh MLP for the value. Layers are initialised orthogonally at gain +sqrt(2), except the mean's last layer at 0.01 - so the first policy is centred +on zero with unit deviation in every dimension - and the value's at 1.0. + +CleanRL does three things in gymnasium wrappers that PlugRL's environment +client does not do, so they are done here instead: + + ClipAction the environment receives the sample clipped to + [-action_clip, action_clip]; the runtime state keeps + the unclipped sample, whose density PPO's ratio is. + NormalizeObservation running mean and variance, stored as buffers and + updated by the algorithm between iterations rather + than at every step, so that an iteration's collection + and learning see the same normalisation. + TransformObservation the normalised observation clipped to +/-obs_clip. + +Rewards are normalised by `ppo`'s buffer, which has them. +""" + +from __future__ import annotations + +import dataclasses +from typing import Any + +import numpy as np +import torch +import torch.nn as nn + +from ..base_torch_policy import BaseTorchPolicy, BaseTorchPolicyConfig +from ..fpo.fpo_policy import _update_running_stats +from ..registration import register_policy, register_policy_config + +UID = "gaussian-policy" + + +def _layer(in_dim: int, out_dim: int, gain: float = float(np.sqrt(2))) -> nn.Linear: + layer = nn.Linear(in_dim, out_dim) + nn.init.orthogonal_(layer.weight, gain) + nn.init.constant_(layer.bias, 0.0) + return layer + + +def _tanh_mlp( + in_dim: int, hidden_dims: tuple[int, ...], out_dim: int, *, out_gain: float +) -> nn.Sequential: + layers: list[nn.Module] = [] + for hidden in hidden_dims: + layers += [_layer(in_dim, hidden), nn.Tanh()] + in_dim = hidden + layers.append(_layer(in_dim, out_dim, out_gain)) + return nn.Sequential(*layers) + + +@dataclasses.dataclass +class GaussianRuntimeState: + obs: torch.Tensor # (B, obs_dim), as the environment sent it + action: torch.Tensor # (B, action_dim), the sample before clipping + logprob: torch.Tensor # (B,), summed over the action's dimensions + value: torch.Tensor # (B,) + + +@register_policy_config(UID) +@dataclasses.dataclass +class GaussianPolicyConfig(BaseTorchPolicyConfig): + obs_dim: int = 17 + action_dim: int = 6 + hidden_dims: tuple[int, ...] = (64, 64) + # The observation's state keys, concatenated in this order, as for + # fpo-policy. MuJoCo sends one, "obs". + state_keys: tuple[str, ...] = ("obs",) + normalize_observations: bool = True + obs_clip: float = 10.0 + # Keep the observation statistics as loaded, as fpo-policy can. + freeze_obs_stats: bool = False + # HalfCheetah, Hopper and Walker2d all act in [-1, 1]. + action_clip: float = 1.0 + + +@register_policy(UID) +class GaussianPolicy(BaseTorchPolicy): + config: GaussianPolicyConfig + + def __init__(self, config: GaussianPolicyConfig): + super().__init__(config) + self.action_dim = config.action_dim + self.action_horizon = 1 + self.actor_mean = _tanh_mlp( + config.obs_dim, config.hidden_dims, config.action_dim, out_gain=0.01 + ) + self.actor_logstd = nn.Parameter(torch.zeros(1, config.action_dim)) + self.critic = _tanh_mlp(config.obs_dim, config.hidden_dims, 1, out_gain=1.0) + self.register_buffer("obs_stats_count", torch.zeros((), dtype=torch.float32)) + self.register_buffer("obs_stats_mean", torch.zeros(config.obs_dim)) + self.register_buffer("obs_stats_var_sum", torch.zeros(config.obs_dim)) + self.register_buffer("obs_stats_std", torch.ones(config.obs_dim)) + self.to(self.device) + + def extract_model_obs_tensor(self, _obs: dict[str, Any]) -> torch.Tensor: + states = _obs["states"] + keys = self.config.state_keys + missing = [k for k in keys if k not in states] + if missing: + raise KeyError( + f"state keys {missing} are not in the observation, which has " + f"{sorted(states)}; set GaussianPolicyConfig.state_keys" + ) + parts = [ + torch.as_tensor(states[k], dtype=torch.float32, device=self.device) + for k in keys + ] + x = parts[0] if len(parts) == 1 else torch.cat(parts, dim=-1) + if x.shape[-1] != self.config.obs_dim: + raise ValueError( + f"state keys {list(keys)} give {x.shape[-1]} values per " + f"observation, but obs_dim is {self.config.obs_dim}" + ) + return x + + def normalize_obs(self, x: torch.Tensor) -> torch.Tensor: + if not self.config.normalize_observations: + return x + z = (x - self.obs_stats_mean) / self.obs_stats_std + return z.clamp(-self.config.obs_clip, self.config.obs_clip) + + @torch.no_grad() + def update_obs_stats(self, x: torch.Tensor) -> None: + if not self.config.normalize_observations or self.config.freeze_obs_stats: + return + count, mean, var_sum, _ = _update_running_stats( + x=x.to(self.device), + count=self.obs_stats_count, + mean=self.obs_stats_mean, + var_sum=self.obs_stats_var_sum, + ) + self.obs_stats_count = count + self.obs_stats_mean = mean + self.obs_stats_var_sum = var_sum + # NormalizeObservation's divisor, sqrt(var + 1e-8). + self.obs_stats_std = torch.sqrt(var_sum / count + 1e-8) + + def _distribution(self, z: torch.Tensor) -> torch.distributions.Normal: + mean = self.actor_mean(z) + return torch.distributions.Normal(mean, self.actor_logstd.expand_as(mean).exp()) + + def get_action_and_runtime_state( + self, _obs: dict[str, Any] + ) -> tuple[np.ndarray, GaussianRuntimeState]: + x = self.extract_model_obs_tensor(_obs) + z = self.normalize_obs(x) + dist = self._distribution(z) + sample = dist.sample() + state = GaussianRuntimeState( + obs=x, + action=sample, + logprob=dist.log_prob(sample).sum(-1), + value=self.critic(z).squeeze(-1), + ) + clip = self.config.action_clip + action = sample.clamp(-clip, clip).cpu().numpy().astype(np.float32) + return action[:, None, :], state + + def evaluate_actions( + self, obs: torch.Tensor, action: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Log-probability, entropy and value, each (B,), for PPO's loss.""" + z = self.normalize_obs(obs) + dist = self._distribution(z) + return ( + dist.log_prob(action).sum(-1), + dist.entropy().sum(-1), + self.critic(z).squeeze(-1), + ) + + def get_value(self, _obs: dict[str, Any]) -> torch.Tensor: + z = self.normalize_obs(self.extract_model_obs_tensor(_obs)) + return self.critic(z).squeeze(-1).cpu() + + def fake_runtime_state(self, batch_size: int) -> GaussianRuntimeState: + return GaussianRuntimeState( + obs=torch.zeros(batch_size, self.config.obs_dim), + action=torch.zeros(batch_size, self.action_dim), + logprob=torch.zeros(batch_size), + value=torch.zeros(batch_size), + ) diff --git a/tests/test_gaussian_ppo.py b/tests/test_gaussian_ppo.py new file mode 100644 index 0000000..416c470 --- /dev/null +++ b/tests/test_gaussian_ppo.py @@ -0,0 +1,463 @@ +"""`gaussian-policy` and `ppo`: CleanRL's continuous-action PPO, through PlugRL. + +Every other policy-algorithm pair on the coverage figure trains an expressive +policy - a flow or a diffusion model - with an algorithm built for it. None is +the baseline those are measured against: a Gaussian MLP trained by PPO. These +two are that, written to CleanRL's `ppo_continuous_action.py` (Huang et al., +"The 37 Implementation Details of Proximal Policy Optimization", 2022), so +that if the pair fails to learn MuJoCo it is the platform that failed and not +an unfamiliar algorithm. + +CleanRL does part of its work in gymnasium wrappers on the environment - +clipping actions, normalising observations and rewards. PlugRL's environment +client does none of that, so here it is done on the server, and these tests +check each piece where PlugRL's plumbing could get it wrong: the sample the +log-probability was taken of is the one learned from, the ratio is one before +any update, the statistics do not move between collecting and learning, the +schedule anneals as CleanRL's does, and the whole thing improves a policy on a +task with a known answer. +""" + +from __future__ import annotations + +import math +import subprocess +import sys + +import numpy as np +import pytest +import torch + +from plugrl_server.algorithm.ppo import ppo as ppo_module +from plugrl_server.algorithm.ppo.ppo import PPOAlgorithm +from plugrl_server.algorithm.ppo.ppo_buffer import PPOBuffer +from plugrl_server.algorithm.ppo.ppo_config import PPOAlgoConfig +from plugrl_server.common.data_utils import unbatch_aggregate +from plugrl_server.policy.gaussian.gaussian_policy import ( + GaussianPolicy, + GaussianPolicyConfig, +) +from plugrl_server.policy.state import slice_policy_step_state + +OBS_DIM = 5 +ACTION_DIM = 3 +ENVS = 2 +TARGET = 0.5 + + +def _policy(seed: int = 0, **overrides) -> GaussianPolicy: + torch.manual_seed(seed) + return GaussianPolicy( + GaussianPolicyConfig( + obs_dim=OBS_DIM, action_dim=ACTION_DIM, device="cpu", **overrides + ) + ) + + +def _algo(seed: int = 0, **overrides) -> PPOAlgorithm: + fields = dict(buffer_size=128, batch_size=32, update_epochs=2, train_itrs=10) + fields.update(overrides) + algo = PPOAlgorithm(PPOAlgoConfig(**fields), _policy(seed)) + algo.init_optimizers() + return algo + + +def _obs(envs: int, rng: np.random.Generator) -> dict: + return {"states": {"obs": rng.normal(size=(envs, OBS_DIM)).astype(np.float32)}} + + +def _collect(algo, rng, *, reward_fn=lambda action: 1.0, episode_length=5): + """Fill the buffer the way websocket_agent_server does. Returns the rewards.""" + prev_node = {env: (-1, "") for env in range(ENVS)} + terminated = {env: False for env in range(ENVS)} + obs = _obs(ENVS, rng) + rewards: list[float] = [] + for round_index in range(10_000): + action, runtime_state = algo.infer(obs) + step_state = algo.build_step_state_from_runtime_state( + runtime_state, include_train_state=True + ) + next_obs = _obs(ENVS, rng) + obs_list = unbatch_aggregate(obs, aggregate_method="concat") + next_obs_list = unbatch_aggregate(next_obs, aggregate_method="concat") + done = round_index % episode_length == episode_length - 1 + for env in range(ENVS): + reward = float(reward_fn(action[env])) + rewards.append(reward) + one = slice_policy_step_state(step_state, slice(env, env + 1)) + prev_node[env], _, _ = algo.feedback( + obs=obs_list[env], + runtime_state=one.runtime_state, + train_state=one.train_state, + terminated=terminated[env], + truncated=False, + next_obs=next_obs_list[env], + reward=reward, + info={}, + next_terminated=done, + next_truncated=False, + prev_node=prev_node[env], + ) + terminated[env] = done + obs = next_obs + if algo.should_learn(): + return rewards + pytest.fail("the buffer never filled") + + +def _iteration(algo, rng, **collect_kwargs) -> tuple[float, dict]: + rewards = _collect(algo, rng, **collect_kwargs) + algo.pre_learn() + _, metrics = algo.learn() + algo.post_learn() + return float(np.mean(rewards)), metrics + + +def _bandit_reward(action: np.ndarray) -> float: + return -float(((action - TARGET) ** 2).mean()) + + +class TestThePolicy: + def test_it_declares_its_action_shape(self): + """The server's metadata message publishes both.""" + policy = _policy() + + assert (policy.action_dim, policy.action_horizon) == (ACTION_DIM, 1) + + def test_the_env_gets_a_clipped_sample_and_the_state_keeps_the_sample(self): + """ClipAction clips what the environment receives, not what was drawn. + + The log-probability is the density of the unclipped sample, and PPO's + ratio is only a ratio of densities if learning scores that same sample. + """ + policy = _policy() + policy.actor_logstd.data.fill_(1.0) # std e: many samples land outside + with torch.inference_mode(): + action, state = policy.get_action_and_runtime_state( + _obs(64, np.random.default_rng(0)) + ) + + assert action.shape == (64, 1, ACTION_DIM) + assert action.dtype == np.float32 + assert np.abs(action).max() <= 1.0 + assert state.action.shape == (64, ACTION_DIM) + assert (state.action.abs() > 1).any() + np.testing.assert_array_equal(action[:, 0], state.action.clamp(-1, 1).numpy()) + assert state.logprob.shape == state.value.shape == (64,) + + def test_learning_scores_exactly_what_was_sampled(self): + policy = _policy() + with torch.inference_mode(): + _, state = policy.get_action_and_runtime_state( + _obs(16, np.random.default_rng(0)) + ) + + logprob, entropy, value = policy.evaluate_actions(state.obs, state.action) + + torch.testing.assert_close(logprob, state.logprob) + torch.testing.assert_close(value, state.value) + # A diagonal Gaussian's entropy, summed over the action's dimensions. + per_dim = 0.5 + 0.5 * math.log(2 * math.pi) + policy.actor_logstd.detach() + torch.testing.assert_close(entropy, per_dim.sum().expand(16)) + + def test_it_is_initialised_the_way_cleanrl_initialises_it(self): + """Orthogonal weights at gain sqrt(2), 0.01 on the mean's last layer and + 1.0 on the value's, zero biases, a log std of zero, tanh between.""" + policy = _policy() + + assert torch.equal(policy.actor_logstd.detach(), torch.zeros(1, ACTION_DIM)) + for net, out, last_gain in ( + (policy.actor_mean, ACTION_DIM, 0.01), + (policy.critic, 1, 1.0), + ): + layers = [m for m in net if isinstance(m, torch.nn.Linear)] + assert [layer.out_features for layer in layers] == [64, 64, out] + assert sum(isinstance(m, torch.nn.Tanh) for m in net) == 2 + for layer in layers: + assert torch.equal(layer.bias.detach(), torch.zeros_like(layer.bias)) + first, hidden, last = (layer.weight.detach() for layer in layers) + # Orthogonal: a tall matrix's columns and a wide one's rows are + # orthonormal, times the gain. + torch.testing.assert_close(first.T @ first, 2.0 * torch.eye(OBS_DIM)) + torch.testing.assert_close(hidden @ hidden.T, 2.0 * torch.eye(64)) + torch.testing.assert_close( + last @ last.T, last_gain**2 * torch.eye(out), atol=1e-6, rtol=1e-5 + ) + + def test_observations_are_normalised_then_clipped_at_ten(self): + """NormalizeObservation, then TransformObservation's clip to [-10, 10].""" + policy = _policy() + rng = np.random.default_rng(0) + scales = np.array([0.1, 1.0, 10.0, 100.0, 1000.0]) + data = torch.as_tensor( + rng.normal(3.0, scales, size=(4096, OBS_DIM)), dtype=torch.float32 + ) + + policy.update_obs_stats(data) + z = policy.normalize_obs(data) + + torch.testing.assert_close(z.mean(0), torch.zeros(OBS_DIM), atol=1e-3, rtol=0) + torch.testing.assert_close(z.std(0), torch.ones(OBS_DIM), atol=1e-3, rtol=0) + far = policy.normalize_obs(torch.full((1, OBS_DIM), 1e6)) + assert far.max().item() == 10.0 + + def test_frozen_statistics_do_not_move(self): + policy = _policy(freeze_obs_stats=True) + + policy.update_obs_stats(torch.randn(64, OBS_DIM) * 5) + + assert policy.obs_stats_count.item() == 0.0 + assert torch.equal(policy.obs_stats_std, torch.ones(OBS_DIM)) + + def test_state_keys_are_concatenated_in_the_order_given(self): + policy = _policy(state_keys=("a", "b")) + obs = { + "states": { + "b": np.ones((2, 3), dtype=np.float32), + "a": np.zeros((2, 2), dtype=np.float32), + } + } + + x = policy.extract_model_obs_tensor(obs) + + assert torch.equal(x, torch.tensor([[0.0, 0, 1, 1, 1]] * 2)) + + +class TestTheBuffer: + def _buffer(self, **kwargs) -> tuple[PPOBuffer, dict]: + algo = _algo() + example = algo.example_train_state(batch_size=1) + return PPOBuffer(buffer_size=8, example_train_state=example, **kwargs), example + + def test_the_discounted_return_restarts_with_each_episode(self): + """As gymnasium's NormalizeReward and DPPO's RunningRewardScaler keep it. + + A frame's `terminated` says the step before it ended an episode, so a + frame carrying it is the first of a new one. + """ + buffer, one = self._buffer(gamma=0.5) + node = (-1, "") + for starts_episode in (False, False, True, False): + node = buffer.add_frame( + prev_node=node, + train_state=one, + reward=1.0, + terminated=starts_episode, + truncated=False, + last_value=None, + next_terminated=False, + next_truncated=False, + ) + + np.testing.assert_allclose(buffer.rets[:4], [1.0, 1.5, 1.0, 1.5]) + + @pytest.mark.parametrize("normalize", [True, False]) + def test_rewards_are_scaled_by_the_returns_deviation_and_clipped(self, normalize): + buffer, one = self._buffer(gamma=0.9, normalize_rewards=normalize) + raw = np.array([0.1, 50.0, -3.0, 2.0, 0.0, 400.0, 1.0, -1.0], np.float32) + node = (-1, "") + for reward in raw: + node = buffer.add_frame( + prev_node=node, + train_state=one, + reward=float(reward), + terminated=False, + truncated=False, + last_value=None, + next_terminated=False, + next_truncated=False, + ) + buffer.add_next_obs_value_request( + obs=_obs(1, np.random.default_rng(0)), end_node=node + ) + + buffer.compute_advantages_and_returns(policy=_policy(), batch_size=8) + + if not normalize: + np.testing.assert_array_equal(buffer.rewards[:8], raw) + return + scale = np.sqrt(buffer.ret_rms.var + 1e-8) + expected = np.clip(raw / scale, -10.0, 10.0) + np.testing.assert_allclose(buffer.rewards[:8], expected, rtol=1e-6) + assert np.abs(buffer.rewards[:8]).max() < np.abs(raw).max() + + +class TestTheAlgorithm: + def test_one_iteration_moves_every_parameter_and_reports_finite_losses(self): + algo = _algo() + before = {k: v.detach().clone() for k, v in algo.policy.named_parameters()} + + _, metrics = _iteration( + algo, np.random.default_rng(0), reward_fn=_bandit_reward + ) + + for name, parameter in algo.policy.named_parameters(): + assert not torch.equal(parameter.detach(), before[name]), name + for key in ( + "policy_loss", + "value_loss", + "entropy", + "old_approx_kl", + "approx_kl", + "clipfrac", + ): + assert math.isfinite(metrics["losses"][key]), key + assert algo.curr_train_itrs == 1 + assert len(algo.rollout_buffer) == 0 + + def test_before_any_update_every_ratio_is_one(self): + algo = _algo() + _collect(algo, np.random.default_rng(0)) + algo.pre_learn() + n = len(algo.rollout_buffer) + obs = torch.as_tensor( + algo.rollout_buffer.train_state_storage.get_item(slice(0, n)) + ) + + logprob, _, _ = algo.policy.evaluate_actions( + obs, torch.as_tensor(algo.rollout_buffer.actions[:n]) + ) + + torch.testing.assert_close( + logprob, torch.as_tensor(algo.rollout_buffer.logprobs[:n]) + ) + + def test_statistics_are_updated_after_learning_from_the_buffer(self): + """After, as DPPO does for fpo-policy: one iteration collects and learns + under one normalisation, and the next collects under the new one.""" + algo = _algo() + rng = np.random.default_rng(0) + _collect(algo, rng) + algo.pre_learn() + algo.learn() + assert algo.policy.obs_stats_count.item() == 0.0 + + algo.post_learn() + + assert algo.policy.obs_stats_count.item() == 128.0 + + def test_the_learning_rate_anneals_linearly_as_cleanrl_anneals_it(self): + """lr = (1 - (iteration - 1) / num_iterations) * learning_rate.""" + algo = _algo(train_itrs=4, learning_rate=1e-3) + rng = np.random.default_rng(0) + + seen = [_iteration(algo, rng)[1]["models"]["learning_rate"] for _ in range(4)] + + np.testing.assert_allclose(seen, [1e-3, 0.75e-3, 0.5e-3, 0.25e-3]) + assert algo.should_stop() + + def test_without_annealing_the_rate_stays_put(self): + algo = _algo(train_itrs=4, learning_rate=1e-3, anneal_lr=False) + rng = np.random.default_rng(0) + + seen = [_iteration(algo, rng)[1]["models"]["learning_rate"] for _ in range(2)] + + np.testing.assert_allclose(seen, [1e-3, 1e-3]) + + def test_one_adam_over_every_parameter_with_cleanrl_epsilon(self): + algo = _algo() + + (group,) = algo.optimizer.param_groups + assert isinstance(algo.optimizer, torch.optim.Adam) + assert group["eps"] == 1e-5 + assert {id(p) for p in group["params"]} == { + id(p) for p in algo.policy.parameters() + } + + def test_every_step_clips_the_whole_gradient_at_half(self, monkeypatch): + calls = [] + real = torch.nn.utils.clip_grad_norm_ + + def recording(parameters, max_norm, *args, **kwargs): + parameters = list(parameters) + calls.append((len(parameters), max_norm)) + return real(parameters, max_norm, *args, **kwargs) + + monkeypatch.setattr(ppo_module.nn.utils, "clip_grad_norm_", recording) + algo = _algo() + _iteration(algo, np.random.default_rng(0)) + + steps = algo.config.update_epochs * algo.config.buffer_size // 32 + assert calls == [(len(list(algo.policy.parameters())), 0.5)] * steps + + def test_a_checkpoint_resumes_where_it_left_off(self): + algo = _algo() + rng = np.random.default_rng(0) + _iteration(algo, rng, reward_fn=_bandit_reward) + checkpoint = algo.create_checkpoint() + + resumed = _algo(seed=1) + resumed.load_checkpoint(checkpoint) + + for (name, a), b in zip( + algo.policy.state_dict().items(), resumed.policy.state_dict().values() + ): + assert torch.equal(a, b), name + assert resumed.curr_train_itrs == 1 + assert resumed.global_step == algo.global_step + assert resumed.rollout_buffer.ret_rms.var == algo.rollout_buffer.ret_rms.var + assert resumed.rollout_buffer.ret_rms.count == algo.rollout_buffer.ret_rms.count + state, other = algo.optimizer.state_dict(), resumed.optimizer.state_dict() + for key in state["state"]: + assert torch.equal( + state["state"][key]["exp_avg"], other["state"][key]["exp_avg"] + ) + # The schedule is a function of the iteration, so it resumes with it. + _, metrics = _iteration(resumed, rng) + assert metrics["models"]["learning_rate"] == pytest.approx(3e-4 * 0.9) + + +@pytest.mark.parametrize("seed", (0, 1, 2, 3, 4)) +def test_ppo_learns_a_bandit_with_a_known_answer_and_keeps_it(seed: int): + """The bandit `test_fpo_learns_anything` gives FPO: reward is minus the + squared distance from the action to 0.5 in every dimension. + + A Gaussian at mean zero with unit deviation scores about -0.78 here. It + has to end at a quarter of that, and no worse than twice its best - the + bar four of FPO's ten cases fail. Measured, last over best: 1.00, 1.12, + 1.00, 1.00, 1.02, ending between -0.115 and -0.145. + + The rate is 3e-3, not CleanRL's 3e-4, because the floor on this reward is + set by how far the deviation narrows, and at 3e-4, annealed over thirty + iterations of this size, it cannot narrow far: seed 0 goes from -0.78 to + -0.23, still rising at the last iteration. + """ + algo = _algo( + seed=seed, + buffer_size=256, + batch_size=64, + update_epochs=10, + train_itrs=30, + learning_rate=3e-3, + ) + rng = np.random.default_rng(seed) + + history = [_iteration(algo, rng, reward_fn=_bandit_reward)[0] for _ in range(30)] + + best, last = max(history), history[-1] + assert np.isfinite(history).all(), history + assert last > history[0] / 4, ( + f"first {history[0]:+.4f}, best {best:+.4f}, last {last:+.4f}" + ) + assert last > 2 * best, f"best {best:+.4f}, last {last:+.4f}" + + +def test_both_are_on_the_cli_of_a_plain_install(): + """In a fresh interpreter, so it is the packages' own imports that register + them and not this file's.""" + result = subprocess.run( + [ + sys.executable, + "-m", + "plugrl_server.cli", + "gaussian-policy", + "default", + "--help", + ], + capture_output=True, + text=True, + timeout=180, + ) + + assert result.returncode == 0, result.stderr[-2000:] + assert "ppo" in result.stdout 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 2/5] 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 cfceaae6270ae15e04621f3256190c71b9ebe762 Mon Sep 17 00:00:00 2001 From: tactino <18781106300@163.com> Date: Sun, 27 Sep 2026 17:56:07 -0400 Subject: [PATCH 3/5] test: the PPO bandit's measured numbers, with the GAE fix merged in --- tests/test_gaussian_ppo.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/tests/test_gaussian_ppo.py b/tests/test_gaussian_ppo.py index 416c470..999739f 100644 --- a/tests/test_gaussian_ppo.py +++ b/tests/test_gaussian_ppo.py @@ -414,8 +414,9 @@ def test_ppo_learns_a_bandit_with_a_known_answer_and_keeps_it(seed: int): A Gaussian at mean zero with unit deviation scores about -0.78 here. It has to end at a quarter of that, and no worse than twice its best - the - bar four of FPO's ten cases fail. Measured, last over best: 1.00, 1.12, - 1.00, 1.00, 1.02, ending between -0.115 and -0.145. + bar `test_fpo_holds_what_it_learns` sets FPO, and which some of FPO's + cases fail. Measured, last over best: 1.00, 1.14, 1.00, 1.00, 1.03, + ending between -0.120 and -0.160. The rate is 3e-3, not CleanRL's 3e-4, because the floor on this reward is set by how far the deviation narrows, and at 3e-4, annealed over thirty From 124c04506e18141413fbea8a640e40437b0d9f34 Mon Sep 17 00:00:00 2001 From: tactino <18781106300@163.com> Date: Sun, 27 Sep 2026 18:10:41 -0400 Subject: [PATCH 4/5] gaussian-policy: --policy.deterministic acts with the mean, for evaluation; ppo refuses it --- src/plugrl_server/algorithm/ppo/ppo.py | 5 +++++ .../policy/gaussian/gaussian_policy.py | 5 ++++- tests/test_gaussian_ppo.py | 19 +++++++++++++++++++ 3 files changed, 28 insertions(+), 1 deletion(-) diff --git a/src/plugrl_server/algorithm/ppo/ppo.py b/src/plugrl_server/algorithm/ppo/ppo.py index 0bc1432..e67f688 100644 --- a/src/plugrl_server/algorithm/ppo/ppo.py +++ b/src/plugrl_server/algorithm/ppo/ppo.py @@ -46,6 +46,11 @@ def __init__(self, config: PPOAlgoConfig, policy: GaussianPolicy): f"ppo needs a policy with evaluate_actions, such as " f"gaussian-policy; {type(policy).__name__} has none" ) + if getattr(policy.config, "deterministic", False): + raise ValueError( + "ppo needs a policy that samples: --policy.deterministic is for " + "evaluation, and a mean action has no density to form a ratio from" + ) self.rollout_buffer = PPOBuffer( buffer_size=config.buffer_size, example_train_state=self.example_train_state(batch_size=1), diff --git a/src/plugrl_server/policy/gaussian/gaussian_policy.py b/src/plugrl_server/policy/gaussian/gaussian_policy.py index 32d95d3..a964438 100644 --- a/src/plugrl_server/policy/gaussian/gaussian_policy.py +++ b/src/plugrl_server/policy/gaussian/gaussian_policy.py @@ -78,6 +78,9 @@ class GaussianPolicyConfig(BaseTorchPolicyConfig): freeze_obs_stats: bool = False # HalfCheetah, Hopper and Walker2d all act in [-1, 1]. action_clip: float = 1.0 + # Act with the mean instead of a sample: for evaluation, which acts + # without training's sampling noise. `ppo` refuses it. + deterministic: bool = False @register_policy(UID) @@ -152,7 +155,7 @@ def get_action_and_runtime_state( x = self.extract_model_obs_tensor(_obs) z = self.normalize_obs(x) dist = self._distribution(z) - sample = dist.sample() + sample = dist.mean if self.config.deterministic else dist.sample() state = GaussianRuntimeState( obs=x, action=sample, diff --git a/tests/test_gaussian_ppo.py b/tests/test_gaussian_ppo.py index 999739f..9bb4fa4 100644 --- a/tests/test_gaussian_ppo.py +++ b/tests/test_gaussian_ppo.py @@ -201,6 +201,20 @@ def test_observations_are_normalised_then_clipped_at_ten(self): far = policy.normalize_obs(torch.full((1, OBS_DIM), 1e6)) assert far.max().item() == 10.0 + def test_deterministic_acts_with_the_mean(self): + """For evaluation, which acts without training's sampling noise.""" + policy = _policy(deterministic=True) + obs = _obs(8, np.random.default_rng(0)) + with torch.inference_mode(): + first, _ = policy.get_action_and_runtime_state(obs) + second, _ = policy.get_action_and_runtime_state(obs) + mean = policy.actor_mean( + policy.normalize_obs(policy.extract_model_obs_tensor(obs)) + ) + + np.testing.assert_array_equal(first, second) + np.testing.assert_allclose(first[:, 0], mean.clamp(-1, 1).numpy()) + def test_frozen_statistics_do_not_move(self): policy = _policy(freeze_obs_stats=True) @@ -354,6 +368,11 @@ def test_without_annealing_the_rate_stays_put(self): np.testing.assert_allclose(seen, [1e-3, 1e-3]) + def test_it_refuses_a_policy_that_does_not_sample(self): + """The ratio is of the density of what was done; a mean has none.""" + with pytest.raises(ValueError, match="deterministic"): + PPOAlgorithm(PPOAlgoConfig(), _policy(deterministic=True)) + def test_one_adam_over_every_parameter_with_cleanrl_epsilon(self): algo = _algo() 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 5/5] 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():