AQIT 0.1.0
Loading...
Searching...
No Matches
hf_llm_shim.py
Go to the documentation of this file.
1# Copyright (c) 2025-present Aquin Labs Private Limited. All Rights Reserved.
2# This file is part of the Aquin Engine. Unauthorized copying, modification,
3# distribution, or use of this file, via any medium, is strictly prohibited.
4# Proprietary and confidential. See LICENSE for terms.
5
6"""HuggingFace causal LM wrapper with a minimal HookedTransformer-like surface."""
7
8from __future__ import annotations
9
10import re
11from dataclasses import dataclass
12from typing import Any, Callable
13
14import torch
15import torch.nn as nn
16from transformers import AutoModelForCausalLM, AutoTokenizer
17
18_RESID_POST = re.compile(r"^blocks\.(\d+)\.hook_resid_post$")
19_RESID_PRE = re.compile(r"^blocks\.(\d+)\.hook_resid_pre$")
20
21
22@dataclass
24 n_layers: int
25 d_model: int
26 n_heads: int
27 model_name: str
28
29
30def _names_filter_match(names_filter: Any, name: str) -> bool:
31 if names_filter is None:
32 return True
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
39 return False
40
41
42def _hf_load_device_map(cfg: dict[str, Any], device: str) -> str | dict[str, Any]:
43 """Use device_map='auto' for large HF-native models (e.g. Sarvam 30B MoE)."""
44 explicit = cfg.get("device_map")
45 if explicit:
46 return explicit
47 if cfg.get("hf_only") and int(cfg.get("d_model", 0)) >= 4096:
48 return "auto"
49 return device
50
51
52def _model_has_device_map(hf_model: Any) -> bool:
53 return bool(getattr(hf_model, "hf_device_map", None))
54
55
56def _resolve_modules(hf_model: Any) -> tuple[Any, Any, Any, Any]:
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()
61
62 backbone = getattr(hf_model, "model", hf_model)
63 if hasattr(backbone, "get_decoder"):
64 backbone = backbone.get_decoder()
65
66 layers = (
67 getattr(backbone, "layers", None)
68 or getattr(backbone, "h", None)
69 or getattr(backbone, "decoder", None)
70 )
71 if layers is not None and hasattr(layers, "layers"):
72 layers = layers.layers
73
74 embed = (
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)
79 )
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"):
83 embed = embed.wte
84
85 if layers is None or embed is None:
86 raise RuntimeError(
87 "Could not locate decoder layers / embeddings on HuggingFace model "
88 f"(type={type(hf_model).__name__}, backbone={type(backbone).__name__})"
89 )
90 return backbone, layers, embed, lm_head
91
92
93class _IdentityNorm(nn.Module):
94 def forward(self, x: torch.Tensor) -> torch.Tensor:
95 return x
96
99 def __init__(self, weight: torch.Tensor, bias: torch.Tensor | None = None) -> None:
100 self.W_U = weight
101 if bias is not None:
102 self.b_U = bias
103 else:
104 self.b_U = torch.zeros(weight.shape[0], device=weight.device, dtype=weight.dtype)
105
107def project_residual_to_logits(model: Any, resid: torch.Tensor) -> torch.Tensor:
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)
112 else:
113 w = model.W_U
114 b = None
115
116 d_model = resid.shape[-1]
117 if w.shape[0] == d_model:
118 logits = resid @ w
119 elif w.shape[1] == d_model:
120 logits = resid @ w.transpose(-1, -2)
121 else:
122 raise RuntimeError(
123 f"Unexpected lm_head shape {tuple(w.shape)} for d_model={d_model}"
124 )
125
126 if b is not None and b.numel():
127 logits = logits + b.to(dtype=logits.dtype, device=logits.device)
128 return logits
129
130
131class HfLlmShim:
132 """Enough of HookedTransformer for SAE capture, steering, and eval logits."""
133
134 def __init__(
135 self,
136 hf_model: Any,
137 tokenizer: Any,
138 *,
139 hf_name: str,
140 n_layers: int,
141 d_model: int,
142 n_heads: int,
143 ) -> None:
144 self.hf_model = hf_model
145 self.tokenizer = tokenizer
146 self.cfg = _ShimCfg(
147 n_layers=n_layers,
148 d_model=d_model,
149 n_heads=n_heads,
150 model_name=hf_name,
151 )
152 self._backbone, self._layers, self._embed, self._lm_head = _resolve_modules(hf_model)
153 self.device = next(hf_model.parameters()).device
154 self._final_norm = (
155 getattr(self._backbone, "norm", None)
156 or getattr(self._backbone, "ln_f", None)
157 or getattr(self._backbone, "final_layer_norm", None)
159 head_bias = getattr(self._lm_head, "bias", None) if self._lm_head is not None else None
160 self._unembed = _UnembedShim(self.W_U, head_bias)
161
162 @property
163 def ln_final(self) -> nn.Module:
164 return self._final_norm if self._final_norm is not None else _IdentityNorm()
165
166 @property
167 def unembed(self) -> _UnembedShim:
168 return self._unembed
169
170 @classmethod
172 cls,
173 hf_name: str,
174 cfg: dict[str, Any],
175 *,
176 dtype: torch.dtype,
177 device: str,
178 hf_kwargs: dict[str, Any] | None = None,
179 ) -> HfLlmShim:
180 trust = bool(cfg.get("trust_remote_code", False))
181 kw = dict(hf_kwargs or {})
182 load_device_map = _hf_load_device_map(cfg, device)
183 common = dict(
184 torch_dtype=dtype,
185 trust_remote_code=trust,
186 low_cpu_mem_usage=True,
187 **kw,
188 )
189 if device in ("mps", "cpu"):
190 hf_model = AutoModelForCausalLM.from_pretrained(hf_name, **common)
191 hf_model = hf_model.to(device)
192 else:
193 hf_model = AutoModelForCausalLM.from_pretrained(
194 hf_name,
195 device_map=load_device_map,
196 **common,
197 )
198 hf_model.eval()
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)))
201 return cls(
202 hf_model,
203 tokenizer,
204 hf_name=hf_name,
205 n_layers=int(cfg["n_layers"]),
206 d_model=int(cfg["d_model"]),
207 n_heads=n_heads,
208 )
209
210 @property
211 def W_E(self) -> torch.Tensor:
212 return self._embed.weight
213
214 @property
215 def W_U(self) -> torch.Tensor:
216 if self._lm_head is None:
217 raise RuntimeError("Model has no lm_head")
218 return self._lm_head.weight
220 @property
221 def blocks(self) -> list[Any]:
222 return list(self._layers)
223
224 def eval(self) -> HfLlmShim:
226 return self
227
228 def parameters(self, recurse: bool = True):
229 return self.hf_model.parameters(recurse=recurse)
230
231 def named_parameters(self, prefix: str = "", recurse: bool = True):
232 return self.hf_model.named_parameters(prefix=prefix, recurse=recurse)
233
234 def to(self, device: str | torch.device) -> HfLlmShim:
236 return self
237 self.hf_model.to(device)
238 self.device = next(self.hf_model.parameters()).device
239 return self
240
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)
244 encoded = self.tokenizer(
245 text,
246 return_tensors="pt",
247 add_special_tokens=prepend_bos,
248 )
249 return encoded.input_ids.to(self.device)
250
251 def to_string(self, tokens: Any) -> str:
252 if isinstance(tokens, torch.Tensor):
253 tokens = tokens.tolist()
254 if isinstance(tokens, list) and tokens and isinstance(tokens[0], list):
255 tokens = tokens[0]
256 return self.tokenizer.decode(tokens, skip_special_tokens=True)
257
258 def __call__(self, tokens: torch.Tensor) -> torch.Tensor:
259 return self._forward_logits(tokens)
260
261 def _forward_logits(self, tokens: torch.Tensor) -> torch.Tensor:
262 with torch.no_grad():
263 out = self.hf_model(input_ids=tokens)
264 return out.logits
266 def _hooks_from_filter(self, names_filter: Any) -> tuple[set[int], set[int], bool]:
267 if names_filter is None:
268 return set(range(self.cfg.n_layers)), set(), False
269
270 resid_post: set[int] = set()
271 resid_pre: set[int] = set()
272 need_embed = _names_filter_match(names_filter, "hook_embed")
273 for layer in range(self.cfg.n_layers):
274 if _names_filter_match(names_filter, f"blocks.{layer}.hook_resid_post"):
275 resid_post.add(layer)
276 if _names_filter_match(names_filter, f"blocks.{layer}.hook_resid_pre"):
277 resid_pre.add(layer)
278 return resid_post, resid_pre, need_embed
279
280 def run_with_cache(
281 self,
282 tokens: torch.Tensor,
283 names_filter: Any = None,
284 return_type: str | None = "logits",
285 ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
286 resid_post, resid_pre, need_embed = self._hooks_from_filter(names_filter)
287 cache: dict[str, torch.Tensor] = {}
288
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]
296 logits = out.logits
297 return logits, cache
298
299 handles: list[Any] = []
300
301 if need_embed:
302 def _embed_hook(_module: Any, _inp: Any, output: torch.Tensor) -> torch.Tensor:
303 cache["hook_embed"] = output
304 return output
305
306 handles.append(self._embed.register_forward_hook(_embed_hook))
307
308 for layer in resid_pre:
309 def _pre_hook(_module: Any, args: tuple[Any, ...], li: int = layer) -> None:
310 if args:
311 cache[f"blocks.{li}.hook_resid_pre"] = args[0]
312
313 handles.append(self._layers[layer].register_forward_pre_hook(_pre_hook))
314
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
319 return output
320
321 handles.append(self._layers[layer].register_forward_hook(_post_hook))
322
323 try:
324 with torch.no_grad():
325 logits = self._forward_logits(tokens)
326 finally:
327 for handle in handles:
328 handle.remove()
329
330 return logits, cache
331
332 def run_with_hooks(
333 self,
334 tokens: torch.Tensor,
335 fwd_hooks: list[tuple[str, Callable[..., torch.Tensor]]] | None = None,
336 return_type: str | None = "logits",
337 ) -> torch.Tensor:
338 handles: list[Any] = []
339
340 for hook_name, hook_fn in fwd_hooks or []:
341 if hook_name == "hook_embed":
342
343 def _embed_hook(_module: Any, _inp: Any, output: torch.Tensor, fn=hook_fn) -> torch.Tensor:
344 return fn(output, None)
345
346 handles.append(self._embed.register_forward_hook(_embed_hook))
347 continue
348
349 m_post = _RESID_POST.match(hook_name)
350 if m_post:
351 layer = int(m_post.group(1))
352
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:]
358 return modified
359
360 handles.append(self._layers[layer].register_forward_hook(_post_hook))
361 continue
362
363 m_pre = _RESID_PRE.match(hook_name)
364 if m_pre:
365 layer = int(m_pre.group(1))
366
367 def _pre_hook(_module: Any, args: tuple[Any, ...], fn=hook_fn) -> tuple[Any, ...] | None:
368 if not args:
369 return None
370 modified = fn(args[0], None)
371 return (modified,) + args[1:]
372
373 handles.append(self._layers[layer].register_forward_pre_hook(_pre_hook))
374
375 try:
376 with torch.no_grad():
377 logits = self._forward_logits(tokens)
378 finally:
379 for handle in handles:
380 handle.remove()
381
382 return logits
named_parameters(self, str prefix="", bool recurse=True)
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)