From d9eb5707d86354e15ce303577f5c7ce7a201b779 Mon Sep 17 00:00:00 2001 From: tactino <18781106300@163.com> Date: Sun, 27 Sep 2026 22:06:28 -0400 Subject: [PATCH] dppo-gaussian-policy and ppo dppo-square: DPPO's Gaussian PPO on square DPPO fine-tunes a Gaussian MLP with PPO as its robomimic baseline, from pretrained checkpoints it releases. This adds that policy, built on the dppo package's ResidualMLP and CriticObs so the released square checkpoint loads by its own keys (2,152,485 parameters, as the paper states), and the options ppo needs to run DPPO's square config: critic_learning_rate a second parameter group (DPPO: actor 1e-4, critic 1e-3) n_critic_warmup_itrs iterations in which the actor gets no gradient max_grad_norm = None no clipping adam_eps CleanRL's 1e-5 by default; DPPO keeps 1e-8 reward_scaling_gamma the running return's own discount (DPPO: 0.99) `ppo dppo-square` carries ft_ppo_gaussian_mlp.yaml's values at irom-lab/dppo cc7234ad. The released checkpoint's logvar_max (a deviation of 1) replaces the config's 0.2 on load, as in DPPO; the policy keeps that and says so. --- src/plugrl_server/algorithm/ppo/ppo.py | 67 +++-- src/plugrl_server/algorithm/ppo/ppo_buffer.py | 7 +- src/plugrl_server/algorithm/ppo/ppo_config.py | 57 ++++- src/plugrl_server/policy/dppo/__init__.py | 3 +- .../policy/dppo/dppo_gaussian_policy.py | 242 ++++++++++++++++++ tests/test_dppo_gaussian_policy.py | 211 +++++++++++++++ tests/test_ppo_dppo_square.py | 187 ++++++++++++++ 7 files changed, 756 insertions(+), 18 deletions(-) create mode 100644 src/plugrl_server/policy/dppo/dppo_gaussian_policy.py create mode 100644 tests/test_dppo_gaussian_policy.py create mode 100644 tests/test_ppo_dppo_square.py diff --git a/src/plugrl_server/algorithm/ppo/ppo.py b/src/plugrl_server/algorithm/ppo/ppo.py index e67f688..9676958 100644 --- a/src/plugrl_server/algorithm/ppo/ppo.py +++ b/src/plugrl_server/algorithm/ppo/ppo.py @@ -58,16 +58,40 @@ def __init__(self, config: PPOAlgoConfig, policy: GaussianPolicy): gae_lambda=config.gae_lambda, normalize_rewards=config.normalize_rewards, reward_clip=config.reward_clip, + reward_scaling_gamma=config.reward_scaling_gamma, ) self.global_step = 0 self.curr_train_itrs = 0 self.last_saved_itr = 0 + def _critic_param_ids(self) -> set[int]: + return {id(p) for p in self.policy.critic.parameters()} + + def _actor_params(self) -> list[torch.nn.Parameter]: + critic = self._critic_param_ids() + return [p for p in self.policy.parameters() if id(p) not in critic] + def init_optimizers(self) -> None: + config = self.config + if config.critic_learning_rate is None: + self.optimizer = torch.optim.Adam( + self.policy.parameters(), lr=config.learning_rate, eps=config.adam_eps + ) + return self.optimizer = torch.optim.Adam( - self.policy.parameters(), lr=self.config.learning_rate, eps=1e-5 + [ + dict(params=self._actor_params(), lr=config.learning_rate), + dict( + params=list(self.policy.critic.parameters()), + lr=config.critic_learning_rate, + ), + ], + eps=config.adam_eps, ) + def in_critic_warmup(self) -> bool: + return self.curr_train_itrs < self.config.n_critic_warmup_itrs + def infer(self, obs: dict) -> tuple[np.ndarray, PolicyRuntimeState]: with torch.inference_mode(): return self.policy.get_action_and_runtime_state(obs) @@ -130,12 +154,14 @@ def get_learn_progress_total(self) -> int | None: self.config.buffer_size / self.config.batch_size ) - def _learning_rate(self) -> float: - """CleanRL's: (1 - (iteration - 1) / num_iterations) * learning_rate.""" + def _anneal_fraction(self) -> float: + """CleanRL's: lr = (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 + return 1.0 + return 1.0 - self.curr_train_itrs / self.config.train_itrs + + def _learning_rate(self) -> float: + return self._anneal_fraction() * self.config.learning_rate def _loss(self, obs, action, oldlogprob, value, advantage, ret): config = self.config @@ -165,8 +191,12 @@ def _loss(self, obs, action, oldlogprob, value, advantage, ret): def learn_impl(self) -> tuple[int, dict]: config = self.config lr = self._learning_rate() - for group in self.optimizer.param_groups: - group["lr"] = lr + frac = self._anneal_fraction() + base = [config.learning_rate, config.critic_learning_rate] + for group, rate in zip(self.optimizer.param_groups, base): + group["lr"] = frac * rate + in_warmup = self.in_critic_warmup() + actor_params = self._actor_params() if in_warmup else [] dataloader = torch.utils.data.DataLoader( self.rollout_buffer, @@ -204,13 +234,11 @@ def learn_impl(self) -> tuple[int, dict]: ) self.optimizer.zero_grad() loss.backward() - grad_norms.append( - float( - nn.utils.clip_grad_norm_( - self.policy.parameters(), config.max_grad_norm - ) - ) - ) + # During a critic warmup the actor's parameters get no + # gradient, so the optimizer skips them entirely. + for param in actor_params: + param.grad = None + grad_norms.append(self._clip_or_measure()) self.optimizer.step() progress += 1 self.report_learn_progress(progress, progress_total) @@ -233,10 +261,19 @@ def learn_impl(self) -> tuple[int, dict]: train=dict( max_grad_norm=max(grad_norms) if grad_norms else 0.0, train_itrs=float(self.curr_train_itrs), + critic_warmup=float(in_warmup), ), rollout=rollout_summary, ) + def _clip_or_measure(self) -> float: + """The gradient's norm, clipped to `max_grad_norm` when that is set.""" + params = [p for p in self.policy.parameters() if p.grad is not None] + if self.config.max_grad_norm is not None: + return float(nn.utils.clip_grad_norm_(params, self.config.max_grad_norm)) + norms = torch.stack([p.grad.detach().norm() for p in params]) + return float(torch.linalg.vector_norm(norms)) + 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 diff --git a/src/plugrl_server/algorithm/ppo/ppo_buffer.py b/src/plugrl_server/algorithm/ppo/ppo_buffer.py index 9cfe268..caba7b8 100644 --- a/src/plugrl_server/algorithm/ppo/ppo_buffer.py +++ b/src/plugrl_server/algorithm/ppo/ppo_buffer.py @@ -32,8 +32,13 @@ def __init__( normalize_rewards: bool = True, reward_clip: float = 10.0, epsilon: float = 1e-8, + reward_scaling_gamma: float | None = None, ): super().__init__(buffer_size, example_train_state, gamma, gae_lambda) + # The running return's discount; GAE's own when unset. + self.reward_scaling_gamma = ( + gamma if reward_scaling_gamma is None else reward_scaling_gamma + ) self.normalize_rewards = normalize_rewards self.reward_clip = reward_clip self.epsilon = epsilon @@ -74,7 +79,7 @@ def add_frame( and not (terminated or truncated) ) self.rets[current_idx] = float(reward) + ( - self.gamma * self.rets[prev_idx] if continues else 0.0 + self.reward_scaling_gamma * self.rets[prev_idx] if continues else 0.0 ) return node diff --git a/src/plugrl_server/algorithm/ppo/ppo_config.py b/src/plugrl_server/algorithm/ppo/ppo_config.py index 4bf0cd3..c461a09 100644 --- a/src/plugrl_server/algorithm/ppo/ppo_config.py +++ b/src/plugrl_server/algorithm/ppo/ppo_config.py @@ -31,7 +31,17 @@ class PPOAlgoConfig(BaseAlgoConfig): clip_vloss: bool = True ent_coef: float = 0.0 vf_coef: float = 0.5 - max_grad_norm: float = 0.5 + # None: no clipping (DPPO's Gaussian PPO). + max_grad_norm: float | None = 0.5 + # Adam's epsilon: CleanRL's 1e-5. DPPO keeps torch's 1e-8. + adam_eps: float = 1e-5 + # A learning rate of the critic's own, in a parameter group of its own; + # None keeps CleanRL's one rate for actor and critic. Annealing scales + # both. + critic_learning_rate: float | None = None + # Iterations at the start in which only the critic learns (DPPO: 1 on + # square, after its evaluation-only iteration 0). + n_critic_warmup_itrs: int = 0 # Stop an iteration's epochs once the approximate KL passes this. None, # CleanRL's default, never stops. target_kl: float | None = None @@ -40,8 +50,53 @@ class PPOAlgoConfig(BaseAlgoConfig): # discounted return, then clipped. normalize_rewards: bool = True reward_clip: float = 10.0 + # The discount of the running return that rewards are scaled by, when it + # is not `gamma`. DPPO's RunningRewardScaler keeps its own 0.99 while its + # GAE discounts by 0.999. + reward_scaling_gamma: float | None = None train_itrs: int = 488 save_interval: int = 50 def __post_init__(self): self.global_steps = self.train_itrs * self.buffer_size + + +@register_algo_config(UID, "dppo-square") +@dataclasses.dataclass +class PPOAlgoConfigDPPOSquare(PPOAlgoConfig): + """DPPO's Gaussian-policy PPO on robomimic square, state input. + + `cfg/robomimic/finetune/square/ft_ppo_gaussian_mlp.yaml` at irom-lab/dppo + cc7234ad, for `dppo-gaussian-policy` started from DPPO's released + checkpoint. One iteration is 50 environments x 400 action chunks = 20,000 + chunks (80,000 steps), two minibatches of 10,000 over ten epochs. The + rates are the current config's: 1e-4 and 1e-3, constant (the paper's + table lists an actor rate of 1e-5, which the config used before v0.7). + + Not carried over: DPPO's deterministic evaluation every tenth iteration, + which trains nothing; and its GAE, which masks only termination and so, + on square, where episodes only time out, bootstraps across resets. Here a + time-out ends the episode for GAE, as for every other algorithm. + """ + + learning_rate: float = 1e-4 + critic_learning_rate: float | None = 1e-3 + anneal_lr: bool = False + buffer_size: int = 20000 + batch_size: int = 10000 + update_epochs: int = 10 + gamma: float = 0.999 + gae_lambda: float = 0.95 + clip_coef: float = 0.01 + clip_vloss: bool = False + ent_coef: float = 0.0 + vf_coef: float = 0.5 + max_grad_norm: float | None = None + target_kl: float | None = 1.0 + adam_eps: float = 1e-8 + normalize_rewards: bool = True + reward_clip: float = 10.0 + reward_scaling_gamma: float | None = 0.99 + n_critic_warmup_itrs: int = 1 + train_itrs: int = 40 + save_interval: int = 10 diff --git a/src/plugrl_server/policy/dppo/__init__.py b/src/plugrl_server/policy/dppo/__init__.py index a7b110a..bd1b422 100644 --- a/src/plugrl_server/policy/dppo/__init__.py +++ b/src/plugrl_server/policy/dppo/__init__.py @@ -1,6 +1,7 @@ try: + from .dppo_gaussian_policy import DPPOGaussianPolicy as DPPOGaussianPolicy from .dppo_policy import DPPOPolicy as DPPOPolicy - __all__ = ["DPPOPolicy"] + __all__ = ["DPPOGaussianPolicy", "DPPOPolicy"] except Exception: __all__ = [] diff --git a/src/plugrl_server/policy/dppo/dppo_gaussian_policy.py b/src/plugrl_server/policy/dppo/dppo_gaussian_policy.py new file mode 100644 index 0000000..ef069a5 --- /dev/null +++ b/src/plugrl_server/policy/dppo/dppo_gaussian_policy.py @@ -0,0 +1,242 @@ +"""DPPO's Gaussian MLP policy, loadable from the checkpoints DPPO releases. + +DPPO (Ren et al. 2024; irom-lab/dppo, MIT) fine-tunes a Gaussian MLP with PPO +as its baseline on robomimic, from pretrained checkpoints it releases. This +is that policy as `model/common/mlp_gaussian.py`'s `Gaussian_MLP` builds it +with `fixed_std` and `learn_fixed_std`, and as `model/common/gaussian.py`'s +`GaussianModel` samples it, built on the `dppo` package's own `ResidualMLP` +and `CriticObs` so that a released checkpoint loads by its own keys: + + mean tanh(ResidualMLP([obs_dim, 1024, 1024, 1024, horizon x action])) + with Mish; one standard deviation per action dimension, + exp(0.5 x clamp(logvar, logvar_min, logvar_max)), repeated over + the chunk's steps and independent of the observation. + sample a normal draw clamped to the mean +- `randn_clip_value` + deviations; the mean when `deterministic`. + logprob the mean over the chunk's elements of each element's log + density, clamped to [-5, 2], as `PPO_Gaussian` scores it. + critic `CriticObs` with a residual [256, 256, 256] Mish MLP. + +Observations are scaled to [-1, 1] and actions back to the environment's +units by DPPO's stored `normalization.npz`, as `dppo-policy` does. + +One behaviour worth knowing, kept on purpose. The released square checkpoint +carries `network.logvar_max = 0` (a deviation of 1), pretraining's default, +and DPPO loads it non-strictly over the fine-tuning config's 0.2. So DPPO's +fine-tuning actually bounds the deviation at 1.0, not the 0.2 its config and +paper state. This does the same. +""" + +from __future__ import annotations + +import dataclasses +import math +import pathlib +from typing import Any + +import numpy as np +import torch +import torch.nn as nn +from dppo.model.common.critic import CriticObs +from dppo.model.common.mlp import ResidualMLP + +from plugrl_server.common.logging_utils import get_logger +from plugrl_server.paths import PACKAGE_DIR + +from ..base_torch_policy import BaseTorchPolicy, BaseTorchPolicyConfig +from ..registration import register_policy, register_policy_config + +logger = get_logger(__name__) + +UID = "dppo-gaussian-policy" + + +@dataclasses.dataclass +class DPPOGaussianRuntimeState: + obs: torch.Tensor # (B, obs_dim), scaled to [-1, 1] + action: torch.Tensor # (B, horizon x action_dim), the normalised sample + logprob: torch.Tensor # (B,) + value: torch.Tensor # (B,) + + +@register_policy_config(UID) +@dataclasses.dataclass +class DPPOGaussianPolicyConfig(BaseTorchPolicyConfig): + """DPPO's square Gaussian (`cfg/robomimic/finetune/square/ft_ppo_gaussian_mlp.yaml`).""" + + env_type: str = "robomimic" + env_name: str = "square" + state_keys: tuple[str, ...] = ( + "robot0_eef_pos", + "robot0_eef_quat", + "robot0_gripper_qpos", + "object", + ) + obs_dim: int = 23 + action_dim: int = 7 + horizon_steps: int = 4 + mlp_dims: tuple[int, ...] = (1024, 1024, 1024) + fixed_std: float = 0.1 + std_min: float = 0.01 + # Replaced by a checkpoint's own `logvar_max`; see the module docstring. + std_max: float = 0.2 + randn_clip_value: float = 3.0 + logprob_min: float = -5.0 + logprob_max: float = 2.0 + critic_mlp_dims: tuple[int, ...] = (256, 256, 256) + # DPPO's released pretrained policy; its `model` weights are loaded. + checkpoint_path: pathlib.Path | None = None + # Act with the mean: for evaluation. `ppo` refuses it. + deterministic: bool = False + + +class _GaussianMLP(nn.Module): + """The attributes, and so the state-dict keys, of DPPO's `Gaussian_MLP`.""" + + def __init__(self, config: DPPOGaussianPolicyConfig) -> None: + super().__init__() + self.action_dim = config.action_dim + self.horizon_steps = config.horizon_steps + self.mlp_mean = ResidualMLP( + [ + config.obs_dim, + *config.mlp_dims, + config.action_dim * config.horizon_steps, + ], + activation_type="Mish", + out_activation_type="Identity", + ) + self.logvar = nn.Parameter( + torch.full((config.action_dim,), math.log(config.fixed_std**2)) + ) + self.logvar_min = nn.Parameter( + torch.tensor(math.log(config.std_min**2)), requires_grad=False + ) + self.logvar_max = nn.Parameter( + torch.tensor(math.log(config.std_max**2)), requires_grad=False + ) + + def forward(self, state: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + mean = torch.tanh(self.mlp_mean(state)) + logvar = torch.clamp(self.logvar, self.logvar_min, self.logvar_max) + scale = torch.exp(0.5 * logvar).repeat(self.horizon_steps) + return mean, scale.expand_as(mean) + + +@register_policy(UID) +class DPPOGaussianPolicy(BaseTorchPolicy): + config: DPPOGaussianPolicyConfig + + def __init__(self, config: DPPOGaussianPolicyConfig): + super().__init__(config) + self.action_dim = config.action_dim + self.action_horizon = config.horizon_steps + self.network = _GaussianMLP(config) + self.critic = CriticObs( + config.obs_dim, + mlp_dims=list(config.critic_mlp_dims), + activation_type="Mish", + residual_style=True, + ) + self.normalization = dict( + np.load( + PACKAGE_DIR + / "meta" + / "dppo" + / "asset" + / config.env_type + / config.env_name + / "normalization.npz" + ) + ) + if config.checkpoint_path is not None: + self._load_released(config.checkpoint_path) + self.to(self.device) + + def _load_released(self, path: pathlib.Path) -> None: + """Load a DPPO pretraining checkpoint's `model`, as `GaussianModel` does. + + DPPO's load is non-strict: fine-tuning creates `logvar`, which + pretraining did not save. Anything else missing is refused rather than + left at its initialisation. + """ + checkpoint = torch.load(path, map_location="cpu", weights_only=True) + missing, unexpected = self.load_state_dict(checkpoint["model"], strict=False) + missing = [k for k in missing if not k.startswith("critic.")] + if unexpected or set(missing) - {"network.logvar"}: + raise ValueError( + f"{path} does not fit DPPO's Gaussian MLP: missing " + f"{sorted(set(missing) - {'network.logvar'})}, unexpected {unexpected}" + ) + std_max = math.exp(0.5 * float(self.network.logvar_max)) + logger.info( + f"Loaded the model weights of {path}; the deviation is bounded at " + f"{std_max:g} by the checkpoint's logvar_max" + ) + + def extract_model_obs_tensor(self, _obs: dict[str, Any]) -> torch.Tensor: + states = _obs["states"] + raw = np.concatenate([states[k] for k in self.config.state_keys], axis=-1) + n = self.normalization + z = 2 * (raw - n["obs_min"]) / (n["obs_max"] - n["obs_min"]) - 1 + return torch.as_tensor(z, dtype=torch.float32, device=self.device) + + def _logprob(self, dist: torch.distributions.Normal, action: torch.Tensor): + return ( + dist.log_prob(action) + .mean(-1) + .clamp(self.config.logprob_min, self.config.logprob_max) + ) + + def _value(self, z: torch.Tensor) -> torch.Tensor: + return self.critic(z).squeeze(-1) + + def get_action_and_runtime_state( + self, _obs: dict[str, Any] + ) -> tuple[np.ndarray, DPPOGaussianRuntimeState]: + z = self.extract_model_obs_tensor(_obs) + mean, scale = self.network(z) + dist = torch.distributions.Normal(mean, scale) + if self.config.deterministic: + sample = mean + else: + clip = self.config.randn_clip_value + sample = torch.clamp( + dist.sample(), mean - clip * scale, mean + clip * scale + ) + state = DPPOGaussianRuntimeState( + obs=z, + action=sample, + logprob=self._logprob(dist, sample), + value=self._value(z), + ) + batch = sample.shape[0] + chunk = ( + sample.reshape(batch, self.action_horizon, self.action_dim).cpu().numpy() + ) + n = self.normalization + action = ( + 0.5 * (chunk + 1) * (n["action_max"] - n["action_min"]) + n["action_min"] + ) + action = np.clip(action, n["action_min"], n["action_max"]) + return action.astype(np.float32), 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.""" + mean, scale = self.network(obs) + dist = torch.distributions.Normal(mean, scale) + return self._logprob(dist, action), dist.entropy().mean(-1), self._value(obs) + + def get_value(self, _obs: dict[str, Any]) -> torch.Tensor: + return self._value(self.extract_model_obs_tensor(_obs)).cpu() + + def fake_runtime_state(self, batch_size: int) -> DPPOGaussianRuntimeState: + flat = self.action_horizon * self.action_dim + return DPPOGaussianRuntimeState( + obs=torch.zeros(batch_size, self.config.obs_dim), + action=torch.zeros(batch_size, flat), + logprob=torch.zeros(batch_size), + value=torch.zeros(batch_size), + ) diff --git a/tests/test_dppo_gaussian_policy.py b/tests/test_dppo_gaussian_policy.py new file mode 100644 index 0000000..426c101 --- /dev/null +++ b/tests/test_dppo_gaussian_policy.py @@ -0,0 +1,211 @@ +"""`dppo-gaussian-policy`: DPPO's Gaussian MLP, loadable from its released checkpoints. + +DPPO (Ren et al. 2024, irom-lab/dppo) fine-tunes a Gaussian MLP with PPO as +the baseline on robomimic, from pretrained checkpoints it releases. For +square, `square_pre_gaussian_mlp_ta4/.../state_5000.pt` holds `model` and +`ema`, each with: + + network.mlp_mean.layers.{0,1.l1,1.l2,2}.{weight,bias} a ResidualMLP + network.logvar_min, network.logvar_max the std's bounds + +and no `network.logvar`, which fine-tuning creates at log(0.1^2). The +checkpoint's `logvar_max` is log(1^2) = 0, from pretraining's default; the +fine-tuning config asks for 0.2, and DPPO's non-strict load overwrites it +with the checkpoint's. The policy keeps DPPO's behaviour and says so. + +These check what DPPO's `Gaussian_MLP` and `GaussianModel` do, on a +checkpoint built to the same keys, since the released one is not in the +repository. +""" + +from __future__ import annotations + +import math + +import numpy as np +import pytest +import torch + +pytest.importorskip("dppo") + +from plugrl_server.policy.dppo.dppo_gaussian_policy import ( # noqa: E402 + DPPOGaussianPolicy, + DPPOGaussianPolicyConfig, +) + +KEYS = ("robot0_eef_pos", "robot0_eef_quat", "robot0_gripper_qpos", "object") +SIZES = (3, 4, 2, 14) + + +def _policy(**overrides) -> DPPOGaussianPolicy: + torch.manual_seed(0) + return DPPOGaussianPolicy(DPPOGaussianPolicyConfig(device="cpu", **overrides)) + + +def _obs(batch: int, rng: np.random.Generator) -> dict: + return { + "states": { + k: rng.uniform(-0.3, 0.3, size=(batch, n)).astype(np.float32) + for k, n in zip(KEYS, SIZES) + } + } + + +def _released_style_checkpoint(path, policy: DPPOGaussianPolicy) -> dict: + """A checkpoint with the released one's keys: no logvar, logvar_max 0.""" + torch.manual_seed(1) + donor = DPPOGaussianPolicy(DPPOGaussianPolicyConfig(device="cpu")) + model = { + k: v.clone() + for k, v in donor.state_dict().items() + if k.startswith("network.") and k != "network.logvar" + } + model["network.logvar_max"] = torch.tensor(0.0) + torch.save({"epoch": 5000, "model": model, "ema": model}, path) + return model + + +class TestTheNetwork: + def test_its_keys_are_the_released_checkpoints(self): + keys = {k for k in _policy().state_dict() if k.startswith("network.")} + + assert keys == { + "network.logvar", + "network.logvar_min", + "network.logvar_max", + *( + f"network.mlp_mean.layers.{layer}.{p}" + for layer in ("0", "1.l1", "1.l2", "2") + for p in ("weight", "bias") + ), + } + + def test_the_mean_head_has_the_released_shapes(self): + state = _policy().state_dict() + + assert state["network.mlp_mean.layers.0.weight"].shape == (1024, 23) + assert state["network.mlp_mean.layers.1.l1.weight"].shape == (1024, 1024) + assert state["network.mlp_mean.layers.2.weight"].shape == (28, 1024) + assert state["network.logvar"].shape == (7,) + + def test_the_std_starts_at_the_fixed_std(self): + policy = _policy() + + torch.testing.assert_close( + policy.network.logvar.detach(), torch.full((7,), math.log(0.1**2)) + ) + + def test_it_declares_its_action_shape(self): + policy = _policy() + + assert (policy.action_dim, policy.action_horizon) == (7, 4) + + +class TestLoading: + def test_the_released_keys_load_and_logvar_max_comes_with_them(self, tmp_path): + """DPPO's load is non-strict, so the checkpoint's logvar_max (std 1.0) + replaces the configured 0.2. Kept, so the policy behaves as DPPO's.""" + path = tmp_path / "state_5000.pt" + model = _released_style_checkpoint(path, _policy()) + + policy = _policy(checkpoint_path=path) + + for k, v in model.items(): + torch.testing.assert_close(policy.state_dict()[k], v, msg=k) + assert policy.network.logvar_max.item() == 0.0 + torch.testing.assert_close( + policy.network.logvar.detach(), torch.full((7,), math.log(0.1**2)) + ) + + def test_it_loads_model_not_ema(self, tmp_path): + """DPPO's GaussianModel reads `model`; its diffusion model reads `ema`.""" + path = tmp_path / "state_5000.pt" + model = _released_style_checkpoint(path, _policy()) + ema = {k: v + 1.0 for k, v in model.items()} + torch.save({"epoch": 5000, "model": model, "ema": ema}, path) + + policy = _policy(checkpoint_path=path) + + weight = "network.mlp_mean.layers.0.weight" + torch.testing.assert_close(policy.state_dict()[weight], model[weight]) + + def test_a_checkpoint_missing_a_mean_layer_is_refused(self, tmp_path): + path = tmp_path / "state_5000.pt" + model = _released_style_checkpoint(path, _policy()) + del model["network.mlp_mean.layers.2.weight"] + torch.save({"model": model}, path) + + with pytest.raises(ValueError, match="layers.2.weight"): + _policy(checkpoint_path=path) + + +class TestActing: + def test_actions_are_chunks_in_the_environments_units(self): + policy = _policy() + with torch.inference_mode(): + action, state = policy.get_action_and_runtime_state( + _obs(16, np.random.default_rng(0)) + ) + + n = policy.normalization + assert action.shape == (16, 4, 7) and action.dtype == np.float32 + assert (action >= n["action_min"] - 1e-6).all() + assert (action <= n["action_max"] + 1e-6).all() + assert state.action.shape == (16, 28) + assert state.logprob.shape == state.value.shape == (16,) + + def test_observations_are_scaled_to_minus_one_one_by_the_stored_range(self): + policy = _policy() + obs = _obs(4, np.random.default_rng(0)) + + z = policy.extract_model_obs_tensor(obs) + + n = policy.normalization + raw = np.concatenate([obs["states"][k] for k in KEYS], axis=-1) + expected = 2 * (raw - n["obs_min"]) / (n["obs_max"] - n["obs_min"]) - 1 + torch.testing.assert_close(z, torch.as_tensor(expected, dtype=torch.float32)) + + def test_a_sample_stays_within_three_deviations_of_the_mean(self): + policy = _policy() + policy.network.logvar.data.fill_(0.0) # std 1, so the clip binds often + obs = _obs(256, np.random.default_rng(0)) + with torch.inference_mode(): + _, state = policy.get_action_and_runtime_state(obs) + mean, scale = policy.network(state.obs) + + assert ((state.action - mean).abs() <= 3 * scale + 1e-6).all() + + def test_the_logprob_is_the_mean_over_the_chunk_clamped_to_dppos_range(self): + policy = _policy() + obs = _obs(8, np.random.default_rng(0)) + with torch.inference_mode(): + _, state = policy.get_action_and_runtime_state(obs) + mean, scale = policy.network(state.obs) + + expected = torch.distributions.Normal(mean, scale).log_prob(state.action) + torch.testing.assert_close(state.logprob, expected.mean(-1).clamp(-5.0, 2.0)) + + def test_learning_scores_exactly_what_was_sampled(self): + policy = _policy() + with torch.inference_mode(): + _, state = policy.get_action_and_runtime_state( + _obs(8, np.random.default_rng(0)) + ) + + # Fresh tensors, as they come out of the rollout buffer. + logprob, entropy, value = policy.evaluate_actions( + state.obs.clone(), state.action.clone() + ) + + torch.testing.assert_close(logprob, state.logprob) + torch.testing.assert_close(value, state.value) + assert entropy.shape == (8,) + + def test_deterministic_acts_with_the_mean(self): + policy = _policy(deterministic=True) + obs = _obs(4, np.random.default_rng(0)) + with torch.inference_mode(): + _, state = policy.get_action_and_runtime_state(obs) + mean, _ = policy.network(state.obs) + + torch.testing.assert_close(state.action, mean) diff --git a/tests/test_ppo_dppo_square.py b/tests/test_ppo_dppo_square.py new file mode 100644 index 0000000..fe7dee8 --- /dev/null +++ b/tests/test_ppo_dppo_square.py @@ -0,0 +1,187 @@ +"""`ppo dppo-square`: DPPO's Gaussian-policy PPO on robomimic square. + +DPPO fine-tunes its released Gaussian MLP on square with +`cfg/robomimic/finetune/square/ft_ppo_gaussian_mlp.yaml` (irom-lab/dppo +cc7234ad). Where that differs from CleanRL's PPO, `ppo` gains an option, each +off by default: + + critic_learning_rate a separate rate for the critic (1e-3; actor 1e-4) + n_critic_warmup_itrs iterations in which only the critic learns + max_grad_norm = None no gradient clipping + reward_scaling_gamma the discount of the running return that rewards are + scaled by (DPPO's RunningRewardScaler: 0.99), apart + from GAE's discount (0.999) +""" + +from __future__ import annotations + +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.algorithm.registration import REGISTERED_ALGO_CONFIGS +from test_gaussian_ppo import _algo, _collect, _iteration + + +def _state(module) -> dict[str, torch.Tensor]: + return {k: v.detach().clone() for k, v in module.state_dict().items()} + + +def _moved(before, module) -> list[str]: + return [k for k, v in module.state_dict().items() if not torch.equal(v, before[k])] + + +def test_the_variant_carries_dppos_square_settings(): + config = REGISTERED_ALGO_CONFIGS["ppo"]["dppo-square"] + + assert (config.learning_rate, config.critic_learning_rate) == (1e-4, 1e-3) + assert config.anneal_lr is False + assert (config.buffer_size, config.batch_size, config.update_epochs) == ( + 20000, + 10000, + 10, + ) + assert (config.gamma, config.gae_lambda) == (0.999, 0.95) + assert (config.clip_coef, config.clip_vloss, config.vf_coef) == (0.01, False, 0.5) + assert config.ent_coef == 0.0 and config.max_grad_norm is None + assert config.target_kl == 1.0 and config.adam_eps == 1e-8 + assert config.normalize_rewards and config.reward_scaling_gamma == 0.99 + assert config.n_critic_warmup_itrs == 1 + + +class TestTheOptimizer: + def test_unset_it_is_cleanrls_one_adam(self): + (group,) = _algo().optimizer.param_groups + + assert group["eps"] == 1e-5 + + def test_a_critic_rate_splits_actor_and_critic(self): + algo = _algo(critic_learning_rate=1e-3, learning_rate=1e-4, adam_eps=1e-8) + + actor, critic = algo.optimizer.param_groups + critic_ids = {id(p) for p in algo.policy.critic.parameters()} + assert {id(p) for p in critic["params"]} == critic_ids + assert {id(p) for p in actor["params"]} == { + id(p) for p in algo.policy.parameters() if id(p) not in critic_ids + } + assert (actor["lr"], critic["lr"]) == (1e-4, 1e-3) + assert actor["eps"] == critic["eps"] == 1e-8 + + def test_annealing_scales_both_rates(self): + algo = _algo(critic_learning_rate=1e-3, learning_rate=1e-4, train_itrs=4) + rng = np.random.default_rng(0) + _iteration(algo, rng) + _iteration(algo, rng) + + actor, critic = algo.optimizer.param_groups + assert actor["lr"] == pytest.approx(1e-4 * 0.75) + assert critic["lr"] == pytest.approx(1e-3 * 0.75) + + +class TestTheCriticWarmup: + def test_the_actor_does_not_move_during_it_and_the_critic_does(self): + algo = _algo(n_critic_warmup_itrs=1, critic_learning_rate=1e-3) + critic_ids = {id(p) for p in algo.policy.critic.parameters()} + before = {n: p.detach().clone() for n, p in algo.policy.named_parameters()} + + _iteration(algo, np.random.default_rng(0)) + + for name, p in algo.policy.named_parameters(): + moved = not torch.equal(p.detach(), before[name]) + assert moved == (id(p) in critic_ids), name + + def test_after_it_the_actor_moves(self): + algo = _algo(n_critic_warmup_itrs=1) + rng = np.random.default_rng(0) + _iteration(algo, rng) + before = _state(algo.policy.actor_mean) + + _iteration(algo, rng) + + assert _moved(before, algo.policy.actor_mean) != [] + + +def test_no_clipping_when_the_norm_is_none(monkeypatch): + calls = [] + monkeypatch.setattr( + ppo_module.nn.utils, "clip_grad_norm_", lambda *a, **k: calls.append(1) + ) + + _, metrics = _iteration(_algo(max_grad_norm=None), np.random.default_rng(0)) + + assert calls == [] + assert np.isfinite(metrics["train"]["max_grad_norm"]) + + +def test_rewards_are_scaled_by_a_return_with_its_own_discount(): + algo = _algo() + buffer = PPOBuffer( + buffer_size=8, + example_train_state=algo.example_train_state(batch_size=1), + gamma=0.999, + reward_scaling_gamma=0.5, + ) + one = algo.example_train_state(batch_size=1) + node = (-1, "") + for _ in range(3): + node = buffer.add_frame( + prev_node=node, train_state=one, reward=1.0, terminated=False, + truncated=False, last_value=None, next_terminated=False, + next_truncated=False, + ) # fmt: skip + + np.testing.assert_allclose(buffer.rets[:3], [1.0, 1.5, 1.75]) + + +def test_the_config_default_keeps_the_rets_discount_at_gamma(): + assert PPOAlgoConfig().reward_scaling_gamma is None + + +def test_it_drives_dppos_gaussian_policy_end_to_end(): + pytest.importorskip("dppo") + from plugrl_server.policy.dppo.dppo_gaussian_policy import ( + DPPOGaussianPolicy, + DPPOGaussianPolicyConfig, + ) + import test_gaussian_ppo as harness + + torch.manual_seed(0) + policy = DPPOGaussianPolicy(DPPOGaussianPolicyConfig(device="cpu")) + config = REGISTERED_ALGO_CONFIGS["ppo"]["dppo-square"] + algo = PPOAlgorithm( + PPOAlgoConfig(**{**config.__dict__, "buffer_size": 64, "batch_size": 32, + "n_critic_warmup_itrs": 0}), + policy, + ) # fmt: skip + algo.init_optimizers() + keys = policy.config.state_keys + sizes = (3, 4, 2, 14) + + def obs(envs, rng): + return { + "states": { + k: rng.uniform(-0.3, 0.3, size=(envs, n)).astype(np.float32) + for k, n in zip(keys, sizes) + } + } + + original = harness._obs + harness._obs = obs + try: + before = {n: p.detach().clone() for n, p in policy.named_parameters()} + _collect(algo, np.random.default_rng(0), reward_fn=lambda a: float(a.mean())) + algo.pre_learn() + _, metrics = algo.learn() + algo.post_learn() + finally: + harness._obs = original + + moved = [n for n, p in policy.named_parameters() if not torch.equal(p, before[n])] + assert "network.logvar" in moved + assert any(n.startswith("network.mlp_mean") for n in moved) + assert any(n.startswith("critic.") for n in moved) + assert np.isfinite(metrics["losses"]["policy_loss"])