Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
67 changes: 52 additions & 15 deletions src/plugrl_server/algorithm/ppo/ppo.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand All @@ -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
Expand Down
7 changes: 6 additions & 1 deletion src/plugrl_server/algorithm/ppo/ppo_buffer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand Down
57 changes: 56 additions & 1 deletion src/plugrl_server/algorithm/ppo/ppo_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
3 changes: 2 additions & 1 deletion src/plugrl_server/policy/dppo/__init__.py
Original file line number Diff line number Diff line change
@@ -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__ = []
Loading
Loading