transformer_lens.model_bridge.generalized_components.vision_classifier_head module

CLS-token classifier head bridge component.

HF’s ViTForImageClassification.forward() / DeiTForImageClassification.forward() already pool the CLS token before calling self.classifier:

sequence_output = outputs.last_hidden_state pooled_output = sequence_output[:, 0, :] logits = self.classifier(pooled_output)

So this component only ever receives an already-pooled (batch, hidden) tensor. No slicing needed here — this is a thin, hook-named pass-through.

Deliberately NOT covering DeiTForImageClassificationWithTeacher (dual cls_classifier + distillation_classifier head) — see vit.py’s docstring.

class transformer_lens.model_bridge.generalized_components.vision_classifier_head.VisionClassifierHeadBridge(name: str | None, config: Any | None = None, submodules: Dict[str, GeneralizedComponent] | None = None, conversion_rule: BaseTensorConversion | None = None, hook_alias_overrides: Dict[str, str] | None = None, optional: bool = False)

Bases: GeneralizedComponent

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

forward(pooled_output: Tensor, **kwargs: Any) Tensor

Generic forward pass for bridge components with input/output hooks.