diff --git a/CHANGELOG.md b/CHANGELOG.md index bf743fa22..84c0017de 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -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) diff --git a/rl_zoo3/exp_manager.py b/rl_zoo3/exp_manager.py index 770a56ba8..b228b00f1 100644 --- a/rl_zoo3/exp_manager.py +++ b/rl_zoo3/exp_manager.py @@ -1,4 +1,5 @@ import argparse +import copy import importlib import os import pickle as pkl @@ -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, @@ -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) diff --git a/tests/test_hyperparams_opt.py b/tests/test_hyperparams_opt.py index f6d5d6e38..89a0f970d 100644 --- a/tests/test_hyperparams_opt.py +++ b/tests/test_hyperparams_opt.py @@ -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): @@ -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}