Coverage for transformer_lens/model_bridge/generalized_components/vision_classifier_head.py: 85%

11 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-08-11 18:50 +0000

1"""CLS-token classifier head bridge component. 

2 

3HF's ViTForImageClassification.forward() / DeiTForImageClassification.forward() 

4already pool the CLS token before calling self.classifier: 

5 

6 sequence_output = outputs.last_hidden_state 

7 pooled_output = sequence_output[:, 0, :] 

8 logits = self.classifier(pooled_output) 

9 

10So this component only ever receives an already-pooled (batch, hidden) tensor. 

11No slicing needed here — this is a thin, hook-named pass-through. 

12 

13Deliberately NOT covering DeiTForImageClassificationWithTeacher (dual 

14cls_classifier + distillation_classifier head) — see vit.py's docstring. 

15""" 

16 

17from typing import Any 

18 

19from torch import Tensor 

20 

21from transformer_lens.model_bridge.generalized_components.base import ( 

22 GeneralizedComponent, 

23) 

24 

25 

26class VisionClassifierHeadBridge(GeneralizedComponent): 

27 """Wraps the classifier nn.Linear that HF calls with an already-pooled CLS token.""" 

28 

29 def forward(self, pooled_output: Tensor, **kwargs: Any) -> Tensor: 

30 original_component = self.original_component 

31 if original_component is None: 31 ↛ 32line 31 didn't jump to line 32 because the condition on line 31 was never true

32 raise RuntimeError( 

33 f"Original component not set for {self.name}. Call set_original_component() first." 

34 ) 

35 pooled_output = self.hook_in(pooled_output) 

36 logits = original_component(pooled_output) 

37 return self.hook_out(logits)