6"""HuggingFace causal LM wrapper with a minimal HookedTransformer-like surface."""
8from __future__
import annotations
11from dataclasses
import dataclass
12from typing
import Any, Callable
16from transformers
import AutoModelForCausalLM, AutoTokenizer
18_RESID_POST = re.compile(
r"^blocks\.(\d+)\.hook_resid_post$")
19_RESID_PRE = re.compile(
r"^blocks\.(\d+)\.hook_resid_pre$")
31 if names_filter
is None:
33 if callable(names_filter):
34 return bool(names_filter(name))
35 if isinstance(names_filter, str):
36 return name == names_filter
37 if isinstance(names_filter, (set, list, tuple, frozenset)):
38 return name
in names_filter
43 """Use device_map='auto' for large HF-native models (e.g. Sarvam 30B MoE)."""
44 explicit = cfg.get(
"device_map")
47 if cfg.get(
"hf_only")
and int(cfg.get(
"d_model", 0)) >= 4096:
53 return bool(getattr(hf_model,
"hf_device_map",
None))
57 """Return (backbone, layer_list, embed_module, lm_head)."""
58 lm_head = getattr(hf_model,
"lm_head",
None)
59 if lm_head
is None and hasattr(hf_model,
"get_output_embeddings"):
60 lm_head = hf_model.get_output_embeddings()
62 backbone = getattr(hf_model,
"model", hf_model)
63 if hasattr(backbone,
"get_decoder"):
64 backbone = backbone.get_decoder()
67 getattr(backbone,
"layers",
None)
68 or getattr(backbone,
"h",
None)
69 or getattr(backbone,
"decoder",
None)
71 if layers
is not None and hasattr(layers,
"layers"):
72 layers = layers.layers
75 getattr(backbone,
"embed_tokens",
None)
76 or getattr(backbone,
"word_embeddings",
None)
77 or getattr(backbone,
"wte",
None)
78 or getattr(backbone,
"transformer",
None)
80 if embed
is None and hasattr(hf_model,
"get_input_embeddings"):
81 embed = hf_model.get_input_embeddings()
82 if embed
is not None and hasattr(embed,
"wte"):
85 if layers
is None or embed
is None:
87 "Could not locate decoder layers / embeddings on HuggingFace model "
88 f
"(type={type(hf_model).__name__}, backbone={type(backbone).__name__})"
90 return backbone, layers, embed, lm_head
94 def forward(self, x: torch.Tensor) -> torch.Tensor:
99 def __init__(self, weight: torch.Tensor, bias: torch.Tensor |
None =
None) ->
None:
104 self.
b_U = torch.zeros(weight.shape[0], device=weight.device, dtype=weight.dtype)
108 """Map a residual-space direction to vocab logits (TL or HF lm_head layout)."""
109 if hasattr(model,
"unembed"):
110 w = model.unembed.W_U
111 b = getattr(model.unembed,
"b_U",
None)
116 d_model = resid.shape[-1]
117 if w.shape[0] == d_model:
119 elif w.shape[1] == d_model:
120 logits = resid @ w.transpose(-1, -2)
123 f
"Unexpected lm_head shape {tuple(w.shape)} for d_model={d_model}"
126 if b
is not None and b.numel():
127 logits = logits + b.to(dtype=logits.dtype, device=logits.device)
132 """Enough of HookedTransformer for SAE capture, steering, and eval logits."""
153 self.
device = next(hf_model.parameters()).device
159 head_bias = getattr(self.
_lm_head,
"bias",
None)
if self.
_lm_head is not None else None
178 hf_kwargs: dict[str, Any] |
None =
None,
180 trust = bool(cfg.get(
"trust_remote_code",
False))
181 kw = dict(hf_kwargs
or {})
185 trust_remote_code=trust,
186 low_cpu_mem_usage=
True,
189 if device
in (
"mps",
"cpu"):
190 hf_model = AutoModelForCausalLM.from_pretrained(hf_name, **common)
191 hf_model = hf_model.to(device)
193 hf_model = AutoModelForCausalLM.from_pretrained(
195 device_map=load_device_map,
199 tokenizer = AutoTokenizer.from_pretrained(hf_name, trust_remote_code=trust, **kw)
200 n_heads = int(cfg.get(
"n_heads", cfg.get(
"num_attention_heads", 32)))
205 n_layers=int(cfg[
"n_layers"]),
206 d_model=int(cfg[
"d_model"]),
211 def W_E(self) -> torch.Tensor:
215 def W_U(self) -> torch.Tensor:
217 raise RuntimeError(
"Model has no lm_head")
221 def blocks(self) -> list[Any]:
224 def eval(self) -> HfLlmShim:
234 def to(self, device: str | torch.device) -> HfLlmShim:
241 def to_tokens(self, text: str | torch.Tensor, prepend_bos: bool =
False) -> torch.Tensor:
242 if not isinstance(text, str):
243 return text.to(self.
device)
247 add_special_tokens=prepend_bos,
249 return encoded.input_ids.to(self.
device)
252 if isinstance(tokens, torch.Tensor):
253 tokens = tokens.tolist()
254 if isinstance(tokens, list)
and tokens
and isinstance(tokens[0], list):
256 return self.
tokenizer.decode(tokens, skip_special_tokens=
True)
258 def __call__(self, tokens: torch.Tensor) -> torch.Tensor:
262 with torch.no_grad():
263 out = self.
hf_model(input_ids=tokens)
267 if names_filter
is None:
268 return set(range(self.
cfg.n_layers)), set(),
False
270 resid_post: set[int] = set()
271 resid_pre: set[int] = set()
273 for layer
in range(self.
cfg.n_layers):
275 resid_post.add(layer)
278 return resid_post, resid_pre, need_embed
282 tokens: torch.Tensor,
283 names_filter: Any =
None,
284 return_type: str |
None =
"logits",
285 ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
287 cache: dict[str, torch.Tensor] = {}
289 if not resid_pre
and not need_embed:
290 with torch.no_grad():
291 out = self.
hf_model(input_ids=tokens, output_hidden_states=
True)
292 hs = out.hidden_states
or ()
293 for layer
in resid_post:
294 if layer + 1 < len(hs):
295 cache[f
"blocks.{layer}.hook_resid_post"] = hs[layer + 1]
299 handles: list[Any] = []
302 def _embed_hook(_module: Any, _inp: Any, output: torch.Tensor) -> torch.Tensor:
303 cache[
"hook_embed"] = output
306 handles.append(self.
_embed.register_forward_hook(_embed_hook))
308 for layer
in resid_pre:
309 def _pre_hook(_module: Any, args: tuple[Any, ...], li: int = layer) ->
None:
311 cache[f
"blocks.{li}.hook_resid_pre"] = args[0]
313 handles.append(self.
_layers[layer].register_forward_pre_hook(_pre_hook))
315 for layer
in resid_post:
316 def _post_hook(_module: Any, _inp: Any, output: Any, li: int = layer) -> Any:
317 hidden = output[0]
if isinstance(output, tuple)
else output
318 cache[f
"blocks.{li}.hook_resid_post"] = hidden
321 handles.append(self.
_layers[layer].register_forward_hook(_post_hook))
324 with torch.no_grad():
327 for handle
in handles:
334 tokens: torch.Tensor,
335 fwd_hooks: list[tuple[str, Callable[..., torch.Tensor]]] |
None =
None,
336 return_type: str |
None =
"logits",
338 handles: list[Any] = []
340 for hook_name, hook_fn
in fwd_hooks
or []:
341 if hook_name ==
"hook_embed":
343 def _embed_hook(_module: Any, _inp: Any, output: torch.Tensor, fn=hook_fn) -> torch.Tensor:
344 return fn(output,
None)
346 handles.append(self.
_embed.register_forward_hook(_embed_hook))
349 m_post = _RESID_POST.match(hook_name)
351 layer = int(m_post.group(1))
353 def _post_hook(_module: Any, _inp: Any, output: Any, li: int = layer, fn=hook_fn) -> Any:
354 hidden = output[0]
if isinstance(output, tuple)
else output
355 modified = fn(hidden,
None)
356 if isinstance(output, tuple):
357 return (modified,) + output[1:]
360 handles.append(self.
_layers[layer].register_forward_hook(_post_hook))
363 m_pre = _RESID_PRE.match(hook_name)
365 layer = int(m_pre.group(1))
367 def _pre_hook(_module: Any, args: tuple[Any, ...], fn=hook_fn) -> tuple[Any, ...] |
None:
370 modified = fn(args[0],
None)
371 return (modified,) + args[1:]
373 handles.append(self.
_layers[layer].register_forward_pre_hook(_pre_hook))
376 with torch.no_grad():
379 for handle
in handles:
named_parameters(self, str prefix="", bool recurse=True)
_UnembedShim unembed(self)
torch.Tensor __call__(self, torch.Tensor tokens)
tuple[set[int], set[int], bool] _hooks_from_filter(self, Any names_filter)
torch.Tensor to_tokens(self, str|torch.Tensor text, bool prepend_bos=False)
parameters(self, bool recurse=True)
HfLlmShim from_pretrained(cls, str hf_name, dict[str, Any] cfg, *, torch.dtype dtype, str device, dict[str, Any]|None hf_kwargs=None)
str to_string(self, Any tokens)
None __init__(self, Any hf_model, Any tokenizer, *, str hf_name, int n_layers, int d_model, int n_heads)
HfLlmShim to(self, str|torch.device device)
torch.Tensor run_with_hooks(self, torch.Tensor tokens, list[tuple[str, Callable[..., torch.Tensor]]]|None fwd_hooks=None, str|None return_type="logits")
torch.Tensor _forward_logits(self, torch.Tensor tokens)
tuple[torch.Tensor, dict[str, torch.Tensor]] run_with_cache(self, torch.Tensor tokens, Any names_filter=None, str|None return_type="logits")
torch.Tensor forward(self, torch.Tensor x)
None __init__(self, torch.Tensor weight, torch.Tensor|None bias=None)
torch.Tensor project_residual_to_logits(Any model, torch.Tensor resid)
bool _names_filter_match(Any names_filter, str name)
tuple[Any, Any, Any, Any] _resolve_modules(Any hf_model)
bool _model_has_device_map(Any hf_model)
str|dict[str, Any] _hf_load_device_map(dict[str, Any] cfg, str device)