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

1"""Safe attribute access for transformers>=5.15 heterogeneous configs. 

2 

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""" 

11 

12from typing import Any 

13 

14 

15def per_layer_attr_names(config: Any) -> frozenset: 

16 """Fields a heterogeneous config refuses to serve globally. 

17 

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 ()) 

23 

24 

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))] 

29 

30 

31def majority_value(values: list) -> Any: 

32 """Most common value in a per-layer list; ties break toward the earliest layer. 

33 

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] 

49 

50 

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) 

57 

58 

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. 

62 

63 Wrap a config once and downstream ``hasattr``/``getattr`` probes need no 

64 per-field awareness of heterogeneity. 

65 """ 

66 

67 def __init__(self, config: Any) -> None: 

68 object.__setattr__(self, "_config", config) 

69 object.__setattr__(self, "_het_attrs", per_layer_attr_names(config)) 

70 

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) 

79 

80 

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