Skip to content

Commit 75fff3a

Browse files
Copilotaraffin
andauthored
Fix n_updates logging in MaskablePPO and RecurrentPPO when using early stopping (#325)
* Initial plan * Fix MaskablePPO inaccurate n_updates counting when target_kl early exits Co-authored-by: araffin <1973948+araffin@users.noreply.github.com> * Fix for recurrent ppo too --------- Co-authored-by: copilot-swe-agent[bot] <198982749+Copilot@users.noreply.github.com> Co-authored-by: araffin <1973948+araffin@users.noreply.github.com> Co-authored-by: Antonin RAFFIN <antonin.raffin@ensta.org>
1 parent ee23ad3 commit 75fff3a

5 files changed

Lines changed: 7 additions & 7 deletions

File tree

docs/misc/changelog.md

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22

33
# Changelog
44

5-
## Release 2.8.0a1 (WIP)
5+
## Release 2.8.0a4 (WIP)
66

77
### Breaking Changes:
88

@@ -16,7 +16,8 @@
1616

1717
### Bug Fixes:
1818

19-
- Fix RecurrentPPO and MaskablePPO forward and predict do not reshape action before clip it (@immortal-boy)
19+
- Fix `MaskablePPO` and `RecurrentPPO` inaccurate `n_updates` counting when `target_kl` early exits the training loop
20+
- Fix `RecurrentPPO` and `MaskablePPO` forward and predict do not reshape action before clip it (@immortal-boy)
2021
- Do not call `forward()` method directly in `RecurrentPPO` (@immortal-boy)
2122

2223
### Deprecations:

sb3_contrib/ppo_mask/ppo_mask.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -407,10 +407,9 @@ def train(self) -> None:
407407
th.nn.utils.clip_grad_norm_(self.policy.parameters(), self.max_grad_norm)
408408
self.policy.optimizer.step()
409409

410+
self._n_updates += 1
410411
if not continue_training:
411412
break
412-
413-
self._n_updates += self.n_epochs
414413
explained_var = explained_variance(self.rollout_buffer.values.flatten(), self.rollout_buffer.returns.flatten())
415414

416415
# Logs

sb3_contrib/ppo_recurrent/ppo_recurrent.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -416,10 +416,10 @@ def train(self) -> None:
416416
th.nn.utils.clip_grad_norm_(self.policy.parameters(), self.max_grad_norm)
417417
self.policy.optimizer.step()
418418

419+
self._n_updates += 1
419420
if not continue_training:
420421
break
421422

422-
self._n_updates += self.n_epochs
423423
explained_var = explained_variance(self.rollout_buffer.values.flatten(), self.rollout_buffer.returns.flatten())
424424

425425
# Logs

sb3_contrib/version.txt

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1 @@
1-
2.8.0a1
1+
2.8.0a4

setup.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -67,7 +67,7 @@
6767
packages=[package for package in find_packages() if package.startswith("sb3_contrib")],
6868
package_data={"sb3_contrib": ["py.typed", "version.txt"]},
6969
install_requires=[
70-
"stable_baselines3>=2.8.0a0,<3.0",
70+
"stable_baselines3>=2.8.0a4,<3.0",
7171
],
7272
description="Contrib package of Stable Baselines3, experimental code.",
7373
author="Antonin Raffin",

0 commit comments

Comments
 (0)