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"])