A clean, extensible factory for creating HuggingFace-based Vision Transformer models (ViT, DeiT, DINO, DINOv2, DINOv3, CLIP) with flexible heads and easy backbone freezing.
Installation
pip install vit-zoo
From source:
git clone https://github.com/jbindaAI/vit_zoo.git
cd vit_zoo
pip install -e .
For development: pip install -e ".[dev]"
Quick start
from vit_zoo import ViTModel
model = ViTModel("facebook/dinov2-base", head=10, freeze_backbone=True)
outputs = model(images)
logits = outputs["predictions"] # (batch_size, 10)
Basic usage
from vit_zoo import ViTModel
# Simple classification - pass any HuggingFace model ID
model = ViTModel("google/vit-base-patch16-224", head=10, freeze_backbone=True)
outputs = model(images)
predictions = outputs["predictions"] # Shape: (batch_size, 10)
Custom MLP Head
from vit_zoo import ViTModel, MLPHead
mlp_head = MLPHead(
input_dim=768,
hidden_dims=[512, 256],
output_dim=100,
dropout=0.1,
activation="gelu" # or 'relu', 'tanh', or nn.Module
)
model = ViTModel("facebook/dinov2-base", head=mlp_head)
Embedding Extraction
from vit_zoo import ViTModel
from transformers import CLIPVisionModel
model = ViTModel("openai/clip-vit-base-patch16", backbone_cls=CLIPVisionModel, head=None)
outputs = model(images, output_hidden_states=True)
hidden_states = outputs["last_hidden_state"] # (batch_size, seq_len, embedding_dim)
cls_embedding = hidden_states[:, 0, :] # (batch_size, embedding_dim)
predictions = outputs["predictions"] # same as cls_embedding when head=None (IdentityHead)
Attention Weights
from vit_zoo import ViTModel
model = ViTModel(
"google/vit-base-patch16-224",
head=10,
config_kwargs={"attn_implementation": "eager"}
)
outputs = model(images, output_attentions=True)
attentions = outputs["attentions"]
Custom Head
from vit_zoo import ViTModel, BaseHead
import torch.nn as nn
class CustomHead(BaseHead):
def __init__(self, input_dim: int, num_classes: int):
super().__init__()
self._input_dim = input_dim
self.fc = nn.Linear(input_dim, num_classes)
@property
def input_dim(self) -> int:
return self._input_dim
def forward(self, embeddings):
return self.fc(embeddings)
head = CustomHead(input_dim=768, num_classes=10)
model = ViTModel("google/vit-base-patch16-224", head=head)
Multi-modal Models (CLIP)
For CLIP and other multi-modal models, pass backbone_cls to load only the vision encoder (AutoModel would load the full model):
from vit_zoo import ViTModel
from transformers import CLIPVisionModel
model = ViTModel("openai/clip-vit-base-patch16", backbone_cls=CLIPVisionModel, head=10)
Any HuggingFace Model
ViTModel uses AutoModel to auto-detect the model type from the HuggingFace Hub. Any ViT-compatible model works:
from vit_zoo import ViTModel
model = ViTModel("google/vit-large-patch16-224", head=10)
model = ViTModel("facebook/deit-base-distilled-patch16-224", head=10)
model = ViTModel("facebook/dinov2-with-registers-base", head=10)
API Reference
ViTModel
Single entry point: construct from a HuggingFace model name or from a pre-built backbone.
ViTModel(
model_name: Optional[str] = None,
head: Optional[Union[int, BaseHead]] = None,
backbone: Optional[nn.Module] = None,
backbone_cls: Optional[Type] = None,
freeze_backbone: bool = False,
load_pretrained: bool = True,
backbone_dropout: float = 0.0,
config_kwargs: Optional[Dict[str, Any]] = None,
)
Parameters:
model_name: HuggingFace model identifier (e.g."google/vit-base-patch16-224"). Required unlessbackboneis provided.head:int(creates LinearHead),BaseHeadinstance, orNone(embedding extraction).backbone: Optional pre-built backbone; if set,model_nameand backbone-loading args are ignored.backbone_cls: Optional HuggingFace model class (e.g.CLIPVisionModel). Use for multi-modal models.freeze_backbone: Freeze all backbone parameters.load_pretrained: Load pretrained weights when usingmodel_name.backbone_dropout: Dropout probability in backbone.config_kwargs: Extra config options (e.g.{"attn_implementation": "eager"}).
Usage:
ViTModel("google/vit-base-patch16-224", head=10)ViTModel("facebook/dinov2-base", head=None)(embedding extraction)- Custom backbone:
backbone = vit_zoo.utils._load_backbone(...); ViTModel(backbone=backbone, head=10)
ViTModel.forward()
forward(
pixel_values: torch.Tensor,
output_attentions: bool = False,
output_hidden_states: bool = False,
) -> Dict[str, Any]
Always returns a dict. Keys:
"predictions": head output tensor (always present)"attentions": optional, whenoutput_attentions=True"last_hidden_state": optional, whenoutput_hidden_states=True; shape(batch_size, seq_len, embedding_dim)
Freezing the backbone
model.freeze_backbone(freeze: bool = True) # Freeze/unfreeze backbone
The backbone is the raw HuggingFace model (e.g., model.backbone.encoder.layer.11 for ViT), so you can register hooks and access layers directly without an extra wrapper.
Supported Models
Any ViT-compatible model on the HuggingFace Hub works. Examples:
google/vit-base-patch16-224,google/vit-large-patch16-224(ViT)facebook/deit-base-distilled-patch16-224(DeiT)facebook/dino-vitb16(DINO)facebook/dinov2-base,facebook/dinov2-with-registers-base(DINOv2)facebook/dinov3-vitb16-pretrain-lvd1689m(DINOv3)openai/clip-vit-base-patch16(CLIP Vision; passbackbone_cls=CLIPVisionModel)
Browse the HuggingFace Hub for more models.
Import Patterns
You can import the public API from the root package or from submodules:
# One-line style (recommended)
from vit_zoo import ViTModel, BaseHead, LinearHead, MLPHead, IdentityHead
# Submodule style (explicit namespaces)
from vit_zoo import ViTModel
from vit_zoo.components import BaseHead, LinearHead, MLPHead, IdentityHead
from vit_zoo.utils import _load_backbone # for custom backbone path (private)
Architecture
- Public API (
vit_zoo.__all__):ViTModel,BaseHead,LinearHead,MLPHead,IdentityHead. - Layout:
vit_zoo/model.py(ViT model),vit_zoo/utils/backbone.py(_load_backbone,_get_embedding_dim,_get_cls_token_embedding),vit_zoo/components/(heads). - Extending: Add new heads in
components; usevit_zoo.utils._load_backbonefor custom backbones, thenViTModel(backbone=..., head=...).
Available Heads
LinearHead: Simple linear layer (auto-created whenhead=int)MLPHead: Multi-layer perceptron with configurable depth, activation, dropoutIdentityHead: Returns embeddings unchanged
All heads must implement input_dim property. Custom heads by subclassing BaseHead.
License
GPL-3.0
Release files for vit-zoo 0.2.1
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Source distribution (sdist)
| File | Size | Uploaded | |
|---|---|---|---|
| vit_zoo-0.2.1.tar.gz | 15.2 kB | Details |
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| vit_zoo-0.2.1-py3-none-any.whl | Python 3 | none | any | Details |
Total release size: 25.7 kB
Release files / vit_zoo-0.2.1.tar.gz
| Download URL | vit_zoo-0.2.1.tar.gz |
|---|---|
| Size | 15.2 kB |
| Tags | Source |
|
SHA-256 checksum How to use checksums |
459e903c46c9917329ff5cfc83a83a086c096a5359d97af9c2e770b84f099aab
|
|
BLAKE2b-256 checksum How to use checksums |
0d9706d4395884b2c67d7aa13204ea9899fd407f7f12fb7766bd8e8a49432fb4
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
Yes |
| Uploaded via |
twine/6.1.0 CPython/3.13.7
|
Provenance
Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.
PyPI Publish Attestation
PyPI verified that this artifact, at this checksum, originated from the publisher listed below.
Signed by GitHub Actions, verified by PyPI on Feb 16, 2026.
Transparency logRelease files / vit_zoo-0.2.1-py3-none-any.whl
| Download URL | vit_zoo-0.2.1-py3-none-any.whl |
|---|---|
| Size | 10.5 kB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
551dbe804429ae0abf960583ca1302d18dfa20d84ee852a1da6aba6a4408b432
|
|
BLAKE2b-256 checksum How to use checksums |
a054dcd11c2bc62bbcf506ee7ad000ab09b1de358deb1cd01eeb589b45cd9ac8
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
Yes |
| Uploaded via |
twine/6.1.0 CPython/3.13.7
|
Provenance
Provenance describes where a file came from. On PyPI, provenance is shared via attestations, which provide a verifiable record of the build or publishing details. View details, limitations and caveats.
PyPI Publish Attestation
PyPI verified that this artifact, at this checksum, originated from the publisher listed below.
Signed by GitHub Actions, verified by PyPI on Feb 16, 2026.
Transparency log