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

1"""HuggingFace transformers Driver.""" 

2from __future__ import annotations 

3 

4from typing import Any, Iterator, Mapping 

5 

6import torch 

7from torch import nn 

8 

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 

15 

16 

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

20 

21 _supported_features = frozenset( 

22 {"gradients", "parameters", "state_dict", "weight_access", "intervention_callbacks"} 

23 ) 

24 

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 

29 

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) 

63 

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 

80 

81 return ForwardResult(logits=logits, raw_output=raw) 

82 

83 def parameters(self) -> Iterator[torch.Tensor]: 

84 return self._model.parameters() 

85 

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) 

93 

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 

98 

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