Coverage for transformer_lens/model_bridge/sources/transformers_driver.py: 96%
41 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-21 19:27 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-21 19:27 +0000
1"""HuggingFace transformers Driver."""
2from __future__ import annotations
4from typing import Any, Iterator, Mapping
6import torch
7from torch import nn
9from transformer_lens.model_bridge.driver_protocol import (
10 ForwardResult,
11 Intervention,
12 TensorLike,
13)
14from transformer_lens.model_bridge.sources._driver_base import DriverBase
17class TransformersDriver(DriverBase):
18 """Wraps an HF ``nn.Module``. PyTorch hooks fire via module replacement during the
19 real forward; this driver just runs the engine and threads the native output back."""
21 _supported_features = frozenset(
22 {"gradients", "parameters", "state_dict", "weight_access", "intervention_callbacks"}
23 )
25 def __init__(self, model: nn.Module, adapter: Any, tokenizer: Any) -> None:
26 super().__init__(adapter.cfg, tokenizer)
27 self._model = model
28 self._adapter = adapter
30 def forward(
31 self,
32 input_ids: TensorLike | None = None,
33 *,
34 capture: tuple[str, ...] = (),
35 intervene: Mapping[str, Intervention] | None = None,
36 max_new_tokens: int = 1,
37 return_logits: bool = True,
38 **kwargs: Any,
39 ) -> ForwardResult:
40 # Module-replacement dialect: these args are served by the bridge's
41 # HookPoints, not here — silently ignoring them would be a lie.
42 if capture:
43 raise NotImplementedError(
44 "TransformersDriver.forward does not serve capture=: on the HF "
45 "backend activations are captured through the bridge's HookPoint "
46 "system — use run_with_cache()/run_with_hooks() on the bridge."
47 )
48 if intervene is not None:
49 raise NotImplementedError(
50 "TransformersDriver.forward does not serve intervene=: on the HF "
51 "backend interventions are torch hooks — use run_with_hooks() on "
52 "the bridge."
53 )
54 if max_new_tokens != 1:
55 raise NotImplementedError(
56 "TransformersDriver.forward does not generate; use "
57 "bridge.generate() for multi-token decoding."
58 )
59 if input_ids is not None:
60 raw = self._model(input_ids, **kwargs)
61 else:
62 raw = self._model(**kwargs)
64 logits = None
65 if return_logits: 65 ↛ 81line 65 didn't jump to line 81 because the condition on line 65 was always true
66 if hasattr(raw, "logits"):
67 logits = raw.logits
68 elif isinstance(raw, tuple) and len(raw) > 0:
69 # HF tuple outputs prepend loss when labels are supplied.
70 logits = raw[1] if kwargs.get("labels") is not None and len(raw) > 1 else raw[0]
71 elif hasattr(raw, "last_hidden_state"):
72 # Bare encoder models (ViTModel, DeiTModel, BertModel, etc. without
73 # a task head) return e.g. BaseModelOutput/BaseModelOutputWithPooling,
74 # which has neither `.logits` nor tuple semantics. Fall back to
75 # `last_hidden_state` so return_type="logits" still yields a plain
76 # tensor rather than silently handing back the raw HF output object.
77 logits = raw.last_hidden_state
78 else:
79 logits = raw
81 return ForwardResult(logits=logits, raw_output=raw)
83 def parameters(self) -> Iterator[torch.Tensor]:
84 return self._model.parameters()
86 def named_parameters(
87 self,
88 prefix: str = "",
89 recurse: bool = True,
90 remove_duplicate: bool = True,
91 ) -> Iterator[tuple[str, torch.Tensor]]:
92 return self._model.named_parameters(prefix, recurse, remove_duplicate)
94 @property
95 def underlying_model(self) -> nn.Module:
96 """Escape hatch for code that needs the raw HF module. Driver-specific."""
97 return self._model
99 def set_underlying_model(self, value: nn.Module) -> None:
100 """Used by weight-processing paths that move the model to a different
101 device. Non-torch drivers don't implement this."""
102 self._model = value