From c85500542907c35d72cf3fae9419e856d7cdeb0b Mon Sep 17 00:00:00 2001 From: Victor Morand Date: Tue, 21 Jul 2026 19:02:38 +0200 Subject: [PATCH 1/2] fix: don't tag GridSearch Param if unique value in the Grid --- src/experimaestro/experiments/grid.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/src/experimaestro/experiments/grid.py b/src/experimaestro/experiments/grid.py index 8f0ccdd2..1a3d3b89 100644 --- a/src/experimaestro/experiments/grid.py +++ b/src/experimaestro/experiments/grid.py @@ -296,8 +296,9 @@ def converter(value: Any) -> Any: for combination in grid_combinations: cfg_tags = {} new_cfg = copy.deepcopy(base_cfg) - for path, value in zip(param_paths, combination): - cfg_tags[path] = value + for i, (path, value) in enumerate(zip(param_paths, combination)): + if len(value_options[i]) > 1: + cfg_tags[path] = value set_nested_attr(new_cfg, path, value) # Convert any remaining single-value GenericParams to scalars From a8a29d442a65f5135e6f95c702c8b70691b4a80c Mon Sep 17 00:00:00 2001 From: Victor Morand Date: Thu, 23 Jul 2026 14:19:51 +0200 Subject: [PATCH 2/2] fix: tests --- src/experimaestro/tests/test_grid.py | 14 ++++++++++++-- src/experimaestro/tests/test_grid_validation.py | 17 ++++++++++++++++- 2 files changed, 28 insertions(+), 3 deletions(-) diff --git a/src/experimaestro/tests/test_grid.py b/src/experimaestro/tests/test_grid.py index b2c5e55f..ee1e1a66 100644 --- a/src/experimaestro/tests/test_grid.py +++ b/src/experimaestro/tests/test_grid.py @@ -114,7 +114,17 @@ def test_unique_value_in_tags(): configs, tags = generate_grid(cfg) assert len(configs) == 1 - assert tags[0]["lr"] == 0.1 - assert tags[0]["batch_size"] == 32 + assert tags[0] == {} + + # Test that multi-value params are tagged while single-value params in the same grid are not + cfg_multi = MyConfig(id="test", lr=[0.1, 0.01], batch_size=32) + cfg_multi.lr = GenericParams.from_any(cfg_multi.lr) + cfg_multi.batch_size = GenericParams.from_any(cfg_multi.batch_size) + + configs_multi, tags_multi = generate_grid(cfg_multi) + assert len(configs_multi) == 2 + assert tags_multi[0] == {"lr": 0.1} + assert tags_multi[1] == {"lr": 0.01} + diff --git a/src/experimaestro/tests/test_grid_validation.py b/src/experimaestro/tests/test_grid_validation.py index 7db2336c..73edf0ba 100644 --- a/src/experimaestro/tests/test_grid_validation.py +++ b/src/experimaestro/tests/test_grid_validation.py @@ -79,7 +79,22 @@ def test_unique_value_in_tags_from_validation(): configs, tags = generate_grid(cfg) assert len(configs) == 1 - assert tags[0] == {"lr": 0.05, "sub.value": 10} + assert tags[0] == {} + + # Test with multi-value param + data_multi = { + "id": "test", + "lr": [0.05, 0.1], + "sub": {"value": 10} + } + + cfg_multi = validate_attrs(MainConfig, data_multi) + configs_multi, tags_multi = generate_grid(cfg_multi) + + assert len(configs_multi) == 2 + assert tags_multi[0] == {"lr": 0.05} + assert tags_multi[1] == {"lr": 0.1} + def test_unrecognized_key_in_validation():