diff --git a/src/transformers/integrations/heterogeneity/configuration_utils.py b/src/transformers/integrations/heterogeneity/configuration_utils.py index 836060b9d343..6a8216387db9 100644 --- a/src/transformers/integrations/heterogeneity/configuration_utils.py +++ b/src/transformers/integrations/heterogeneity/configuration_utils.py @@ -227,13 +227,21 @@ def __getitem__(self, layer_idx: int | slice | str) -> PreTrainedConfig | list[P f"Available layer types: {set(layer_types)}" ) - configs = [self[i] for i, layer_type in enumerate(layer_types) if layer_type == layer_idx] - if any(config != configs[0] for config in configs): - raise ValueError( - f"Layer type '{layer_idx}' is not homogeneous across layers. " - f"Use an integer index to access a specific layer's config." - ) - return configs[0] + # Config is actually homogeneous so just return the global config + if not self._config.is_heterogeneous: + return self._config + + # Ensure that all layers of the requested type have the same overrides + layer_overrides = self._config._heterogeneity_spec.per_layer_overrides + reference_overrides = layer_overrides.get(layer_types.index(layer_idx), {}) + for idx, layer_type in enumerate(layer_types): + if layer_type == layer_idx and layer_overrides.get(idx, {}) != reference_overrides: + raise ValueError( + f"Layer type '{layer_idx}' is not homogeneous across layers (layer {idx} differs). " + f"Use an integer index to access a specific layer's config." + ) + + return _get_layer_config(self._config, reference_overrides) # Return a list of configs for a slice of layers if isinstance(layer_idx, slice): diff --git a/tests/heterogeneity/test_configuration_utils.py b/tests/heterogeneity/test_configuration_utils.py index 91624b67583b..356d5d553e7b 100644 --- a/tests/heterogeneity/test_configuration_utils.py +++ b/tests/heterogeneity/test_configuration_utils.py @@ -131,6 +131,50 @@ def test_uniform_per_layer_values_do_not_overwrite_global(self): for layer_idx in range(4): self.assertEqual(config.per_layer_config[layer_idx].num_key_value_heads, 2) + def test_indexing_by_layer_type(self): + config = _tiny_llama_config( + per_layer_config={1: {"num_key_value_heads": 2}, 3: {"num_key_value_heads": 2}}, + layer_types=["full_attention", "sliding_attention"] * 2, + ) + + self.assertEqual(config.per_layer_config["full_attention"].num_key_value_heads, 4) + self.assertEqual(config.per_layer_config["sliding_attention"].num_key_value_heads, 2) + + def test_indexing_homogeneous_config_by_layer_type_returns_global_config(self): + config = _tiny_llama_config(layer_types=["full_attention", "sliding_attention"] * 2) + + self.assertIs(config.per_layer_config["sliding_attention"], config) + + def test_indexing_by_layer_type_ignores_other_layer_types(self): + """Layers of a different type may differ, only the requested type has to be homogeneous.""" + config = _tiny_llama_config( + per_layer_config={0: {"intermediate_size": 32}, 2: {"intermediate_size": 256}}, + layer_types=["full_attention", "sliding_attention"] * 2, + ) + + self.assertEqual(config.per_layer_config["sliding_attention"].intermediate_size, 128) + + def test_indexing_by_heterogeneous_layer_type_raises(self): + config = _tiny_llama_config( + per_layer_config={0: {"intermediate_size": 32}}, + layer_types=["full_attention", "sliding_attention"] * 2, + ) + + with self.assertRaisesRegex(ValueError, "'full_attention' is not homogeneous across layers"): + config.per_layer_config["full_attention"] + + def test_indexing_by_unknown_layer_type_raises(self): + config = _tiny_llama_config(layer_types=["full_attention"] * 4) + + with self.assertRaisesRegex(ValueError, "'sliding_attention' not found in config.layer_types"): + config.per_layer_config["sliding_attention"] + + def test_indexing_by_layer_type_without_layer_types_raises(self): + config = _tiny_llama_config(per_layer_config={0: {"intermediate_size": 32}}) + + with self.assertRaisesRegex(ValueError, "config.layer_types is not defined"): + config.per_layer_config["full_attention"] + def test_explicit_serialization_restores_pruned_global_values(self): per_layer = {layer_idx: {"num_key_value_heads": 4} for layer_idx in range(4)} sparse_config = _tiny_llama_config(per_layer_config=per_layer)