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
« prev ^ index » next coverage.py v7.10.1, created at 2026-08-11 18:50 +0000
1"""CLS-token classifier head bridge component.
3HF's ViTForImageClassification.forward() / DeiTForImageClassification.forward()
4already pool the CLS token before calling self.classifier:
6 sequence_output = outputs.last_hidden_state
7 pooled_output = sequence_output[:, 0, :]
8 logits = self.classifier(pooled_output)
10So this component only ever receives an already-pooled (batch, hidden) tensor.
11No slicing needed here — this is a thin, hook-named pass-through.
13Deliberately NOT covering DeiTForImageClassificationWithTeacher (dual
14cls_classifier + distillation_classifier head) — see vit.py's docstring.
15"""
17from typing import Any
19from torch import Tensor
21from transformer_lens.model_bridge.generalized_components.base import (
22 GeneralizedComponent,
23)
26class VisionClassifierHeadBridge(GeneralizedComponent):
27 """Wraps the classifier nn.Linear that HF calls with an already-pooled CLS token."""
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)