Coverage for transformer_lens/utilities/aliases.py: 87%
59 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"""Utilities for handling hook aliases in the bridge system."""
3import warnings
4from typing import Any, Dict, List, Optional, Set, Union
7def resolve_alias(
8 target_object: Any,
9 requested_name: str,
10 aliases: Dict[str, str] | Dict[str, Union[str, List[str]]],
11) -> Optional[Any]:
12 """Resolve a hook alias to the actual hook object.
14 Args:
15 target_object: The object to get the resolved attribute from
16 requested_name: The name being requested (potentially an alias)
17 aliases: Dictionary mapping alias names to target names
19 Returns:
20 The resolved hook object if alias found, None otherwise
21 """
22 if requested_name in aliases:
23 target_name = aliases[requested_name]
25 if hasattr(target_object, "disable_warnings") and target_object.disable_warnings == False:
26 warnings.warn(
27 f"Hook '{requested_name}' is deprecated and will be removed in a future version. "
28 f"Use '{target_name}' instead.",
29 FutureWarning,
30 stacklevel=3, # Adjusted for utility function call
31 )
33 def _resolve_single_target(target_name: str) -> Any:
34 """Helper function to resolve a single target name."""
35 target_name_split = target_name.split(".")
36 # Resolve dotted paths; list-based aliases try each option
37 if len(target_name_split) > 1:
38 current_attr = target_object
39 for i in range(len(target_name_split) - 1):
40 if not hasattr(current_attr, target_name_split[i]):
41 # Raise so list-based aliases can try next option
42 raise AttributeError(
43 f"'{type(current_attr).__name__}' object has no attribute '{target_name_split[i]}'"
44 )
45 current_attr = getattr(current_attr, target_name_split[i])
47 # Check if the final attribute exists
48 if not hasattr(current_attr, target_name_split[-1]): 48 ↛ 49line 48 didn't jump to line 49 because the condition on line 48 was never true
49 raise AttributeError(
50 f"'{type(current_attr).__name__}' object has no attribute '{target_name_split[-1]}'"
51 )
52 next_attr = getattr(current_attr, target_name_split[-1])
53 return next_attr
54 else:
55 # Check if the target attribute exists before getting it
56 if not hasattr(target_object, target_name):
57 raise AttributeError(
58 f"'{type(target_object).__name__}' object has no attribute '{target_name}'"
59 )
60 return getattr(target_object, target_name)
62 # if the target_name is a list, we check all elements
63 if isinstance(target_name, list): 63 ↛ 64line 63 didn't jump to line 64 because the condition on line 63 was never true
64 for target_name_item in target_name:
65 try:
66 result = _resolve_single_target(target_name_item)
67 return result
68 except AttributeError:
69 continue
70 # If we get here, none of the targets in the list were found
71 raise AttributeError(
72 f"None of the target names {target_name} could be resolved on '{type(target_object).__name__}' object"
73 )
74 else:
75 return _resolve_single_target(target_name)
76 return None
79def _collect_aliases_from_module(
80 module: Any, path: str, aliases: Dict[str, str], visited: Set[int] = set()
81) -> None:
82 """Helper function to collect all aliases from a single module.
83 Args:
84 module: The module to collect aliases from
85 path: Current path prefix for building full names
86 aliases: Dictionary to populate with aliases (modified in-place)
87 visited: Set of already visited module IDs to prevent infinite recursion
88 """
89 mod_id = id(module)
90 if mod_id in visited:
91 return
92 visited.add(mod_id)
94 if hasattr(module, "hook_aliases"):
95 for alias_name, target_name in module.hook_aliases.items():
96 if alias_name == "":
97 # Empty string creates cache alias: embed -> embed.hook_out
98 if path: # Only add if we have a meaningful path
99 aliases[path] = f"{path}.{target_name}"
100 else:
101 # Named hook alias: embed.hook_embed -> embed.hook_out
102 # Handle special case, hook_pos_embed and hook_embed should not be prefixed
103 if path and not (alias_name == "hook_pos_embed" or alias_name == "hook_embed"):
104 full_alias = f"{path}.{alias_name}"
105 full_target = f"{path}.{target_name}"
106 else:
107 full_alias = alias_name
108 full_target = f"{path}.{target_name}" if path else target_name
110 aliases[full_alias] = full_target
112 # Recursively collect from submodules, excluding original_model
113 if hasattr(module, "named_children"):
114 for child_name, child_module in module.named_children():
115 # Skip the original_model to avoid collecting hooks from HuggingFace model
116 if child_name == "original_model" or child_name == "_original_component":
117 continue
119 child_path = f"{path}.{child_name}" if path else child_name
120 _collect_aliases_from_module(child_module, child_path, aliases, visited)
123def collect_aliases_recursive(module: Any, prefix: str = "") -> Dict[str, str]:
124 """Recursively collect all aliases from a module and its children.
125 This unified function collects both:
126 - Named hook aliases: old_hook_name -> new_hook_name
127 - Cache aliases: component_name -> component_name.hook_out (from empty string keys)
128 Args:
129 module: The module to collect aliases from
130 prefix: Path prefix for building full names
131 Returns:
132 Dictionary mapping all alias names to target names
133 """
134 aliases: Dict[str, str] = {}
135 visited: Set[int] = set()
136 _collect_aliases_from_module(module, prefix, aliases, visited)
137 return aliases