Coverage for transformer_lens/utilities/heterogeneous_config.py: 91%
39 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
1"""Safe attribute access for transformers>=5.15 heterogeneous configs.
3Configs whose attention geometry varies across layers (e.g. Gemma 4) register
4those fields as per-layer: reading one on the global config raises
5``AmbiguousGlobalPerLayerAttributeError`` from ``__getattribute__``, which
6neither ``hasattr()`` nor ``getattr(..., default)`` suppresses. These helpers
7let config-probing code stay safe on any transformers version — the per-layer
8machinery is simply absent pre-5.15, where every helper degrades to plain
9attribute access.
10"""
12from typing import Any
15def per_layer_attr_names(config: Any) -> frozenset:
16 """Fields a heterogeneous config refuses to serve globally.
18 Empty for homogeneous configs and pre-5.15 transformers.
19 """
20 if not getattr(config, "is_heterogeneous", False):
21 return frozenset()
22 return frozenset(getattr(config, "per_layer_attributes", None) or ())
25def per_layer_values(config: Any, name: str) -> list:
26 """Collect a per-layer attribute from a heterogeneous config's layer configs."""
27 per_layer_config = config.per_layer_config
28 return [getattr(per_layer_config[i], name, None) for i in range(len(per_layer_config))]
31def majority_value(values: list) -> Any:
32 """Most common value in a per-layer list; ties break toward the earliest layer.
34 Counts by equality, not hashing: per-layer values can be dicts
35 (rope_parameters in 5.x), and a TypeError here escapes hasattr
36 """
37 distinct: list = []
38 counts: list[int] = []
39 for value in values:
40 for index, seen in enumerate(distinct):
41 if seen == value:
42 counts[index] += 1
43 break
44 else:
45 distinct.append(value)
46 counts.append(1)
47 best = max(range(len(distinct)), key=lambda i: (counts[i], -values.index(distinct[i])))
48 return distinct[best]
51def safe_config_get(config: Any, name: str, default: Any = None) -> Any:
52 """getattr that resolves per-layer-registered fields to their majority value."""
53 if name in per_layer_attr_names(config): 53 ↛ 54line 53 didn't jump to line 54 because the condition on line 53 was never true
54 values = [v for v in per_layer_values(config, name) if v is not None]
55 return majority_value(values) if values else default
56 return getattr(config, name, default)
59class HetSafeConfigView:
60 """Read-only getattr proxy: per-layer-registered fields resolve to their
61 majority-layer value instead of raising; everything else passes through.
63 Wrap a config once and downstream ``hasattr``/``getattr`` probes need no
64 per-field awareness of heterogeneity.
65 """
67 def __init__(self, config: Any) -> None:
68 object.__setattr__(self, "_config", config)
69 object.__setattr__(self, "_het_attrs", per_layer_attr_names(config))
71 def __getattr__(self, name: str) -> Any:
72 config = object.__getattribute__(self, "_config")
73 if name in object.__getattribute__(self, "_het_attrs"):
74 values = [v for v in per_layer_values(config, name) if v is not None]
75 if not values: 75 ↛ 76line 75 didn't jump to line 76 because the condition on line 75 was never true
76 raise AttributeError(name)
77 return majority_value(values)
78 return getattr(config, name)
81def het_safe_view(config: Any) -> Any:
82 """Wrap heterogeneous configs in a :class:`HetSafeConfigView`; pass others through."""
83 return HetSafeConfigView(config) if per_layer_attr_names(config) else config