Skip to content
Closed
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
1 change: 1 addition & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
### New Features

### Bug fixes
- Fixed a bug where sampled nested hyperparameters (e.g. ``policy_kwargs``) overwrote the whole dictionary during hyperparameter optimization, dropping user-specified keys like ``features_extractor_class`` (@musaaero)
- Set `learning_starts` to 500 for `parking-v0` in TQC hyperparameters to ensure the first episode is finished
- Fixed "too many open files" issue when doing hyperparameter optimization: the trial environments are now closed in a `finally` block and `gc_after_trial=True` is passed to Optuna to free them between trials (also needs latest SB3 master version to fully work)

Expand Down
26 changes: 25 additions & 1 deletion rl_zoo3/exp_manager.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import argparse
import copy
import importlib
import os
import pickle as pkl
Expand Down Expand Up @@ -60,7 +61,30 @@
get_wrapper_class,
)

def deep_update(original: dict, update: dict) -> dict:
"""
Recursively merge ``update`` into ``original`` and return the merged dict.

Unlike ``dict.update()``, nested dictionaries are merged key-by-key instead of
being replaced wholesale, so sampled nested hyperparameters (e.g. ``policy_kwargs``)
no longer drop user-specified sibling keys (see issue #431).

The inputs are not mutated: ``original`` is deep-copied first, which also prevents
sampled values from leaking into ``self._hyperparams`` across optimization trials
(the previous code used a shallow ``dict.copy()`` that shared nested dicts).
"""
merged = copy.deepcopy(original)
_deep_update(merged, update)
return merged


def _deep_update(target: dict, update: dict) -> None:
for key, value in update.items():
if isinstance(value, dict) and isinstance(target.get(key), dict):
_deep_update(target[key], value)
else:
target[key] = value

class ExperimentManager:
"""
Experiment manager: read the hyperparameters,
Expand Down Expand Up @@ -812,7 +836,7 @@ def objective(self, trial: optuna.Trial) -> float:
# Pass n_actions to initialize DDPG/TD3 noise sampler
# Sample candidate hyperparameters
sampled_hyperparams = HYPERPARAMS_SAMPLER[self.algo](trial, self.n_actions, n_envs, additional_args)
kwargs.update(sampled_hyperparams)
kwargs = deep_update(kwargs, sampled_hyperparams)

env = self.create_envs(n_envs, no_log=True)

Expand Down
42 changes: 42 additions & 0 deletions tests/test_hyperparams_opt.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
import optuna
import pytest
from optuna.trial import TrialState
from rl_zoo3.exp_manager import deep_update


def _assert_eq(left, right):
Expand Down Expand Up @@ -131,3 +132,44 @@ def test_multiple_workers(tmp_path):

for worker in workers:
assert worker.returncode == 0, "STDOUT:\n{}\nSTDERR:\n{}\n".format(*worker.communicate())



def test_deep_update_merges_nested_dicts():
# Regression test for issue #431: sampled nested hyperparameters must merge
# into user-specified nested dicts instead of replacing them wholesale
base = {
"learning_rate": 3e-4,
"policy_kwargs": {
"features_extractor_class": "CustomFeatureExtractor",
"features_extractor_kwargs": {"features_dim": 10},
},
}
sampled = {
"policy_kwargs": {
"net_arch": {"pi": [64], "vf": [64]},
"activation_fn": "relu",
}
}

merged = deep_update(base, sampled)

# Sampled keys are applied
assert merged["policy_kwargs"]["net_arch"] == {"pi": [64], "vf": [64]}
assert merged["policy_kwargs"]["activation_fn"] == "relu"
# Sibling keys are NOT clobbered (the bug: dict.update() dropped them)
assert merged["policy_kwargs"]["features_extractor_class"] == "CustomFeatureExtractor"
assert merged["policy_kwargs"]["features_extractor_kwargs"]["features_dim"] == 10
# Flat keys behave like before
assert merged["learning_rate"] == 3e-4
# Inputs are not mutated (the old shallow copy leaked across trials)
assert "net_arch" not in base["policy_kwargs"]
assert "activation_fn" not in base["policy_kwargs"]


def test_deep_update_replaces_non_dict_values():
# Merging only happens dict-into-dict; anything else replaces, like dict.update()
assert deep_update({"a": {"b": 1}}, {"a": 5}) == {"a": 5}
assert deep_update({"a": 1}, {"a": {"b": 2}}) == {"a": {"b": 2}}
assert deep_update({}, {"a": 1}) == {"a": 1}
assert deep_update({"a": 1}, {}) == {"a": 1}