clip-vitb-mini-distilled / modeling_clip_mini.py
AbstractPhil's picture
AutoModel format: modeling_clip_mini.py
28a619a verified
Raw
History Blame Contribute Delete
3.96 kB
"""HuggingFace modeling file for the distilled CLIP-mini student.
Architecture identical to the training bed's Student class: ViT with
d=240, depth 12, heads 4, patch 16, 160px input, CLS readout, linear
head to the 512-d projection space. If the checkpoint carries a
`rotation` buffer (the consensus-distilled champion does), forward
applies it by default, returning embeddings in the CLIP-B/16 (LAION-2B)
deployment frame — compatible with that teacher's text tower.
Inputs are standard `pixel_values`: images resized so the shorter edge
is 182 (bicubic), center-cropped to 160, rescaled to [0,1], normalized
with the CLIP mean/std (see preprocessor_config.json). The published
evaluations used the torchvision transform pipeline with exactly those
constants.
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import PreTrainedModel
from transformers.modeling_outputs import BaseModelOutputWithPooling
from .configuration_clip_mini import ClipMiniConfig
class _Block(nn.Module):
def __init__(self, d, heads):
super().__init__()
self.n1 = nn.LayerNorm(d)
self.qkv = nn.Linear(d, 3 * d)
self.proj = nn.Linear(d, d)
self.n2 = nn.LayerNorm(d)
self.fc1 = nn.Linear(d, 4 * d)
self.fc2 = nn.Linear(4 * d, d)
self.heads = heads
def forward(self, x):
B, N, C = x.shape
q, k, v = (self.qkv(self.n1(x))
.reshape(B, N, 3, self.heads, C // self.heads)
.permute(2, 0, 3, 1, 4))
a = F.scaled_dot_product_attention(q, k, v)
x = x + self.proj(a.transpose(1, 2).reshape(B, N, C))
return x + self.fc2(F.gelu(self.fc1(self.n2(x))))
class ClipMiniModel(PreTrainedModel):
config_class = ClipMiniConfig
main_input_name = "pixel_values"
def __init__(self, config):
super().__init__(config)
d = config.hidden_size
n_tok = (config.image_size // config.patch_size) ** 2 + 1
self.patch = nn.Conv2d(3, d, config.patch_size, config.patch_size)
self.cls = nn.Parameter(torch.zeros(1, 1, d))
self.pos = nn.Parameter(torch.zeros(1, n_tok, d))
self.blocks = nn.ModuleList(
_Block(d, config.num_attention_heads)
for _ in range(config.num_hidden_layers))
self.norm = nn.LayerNorm(d)
self.head = nn.Linear(d, config.projection_dim)
if config.has_rotation:
self.register_buffer(
"rotation", torch.eye(config.projection_dim), persistent=True)
self.post_init()
def _init_weights(self, module):
if isinstance(module, (nn.Linear, nn.Conv2d)):
nn.init.trunc_normal_(module.weight, std=0.02)
if module.bias is not None:
nn.init.zeros_(module.bias)
def forward_features(self, pixel_values):
x = self.patch(pixel_values).flatten(2).transpose(1, 2)
x = torch.cat([self.cls.expand(x.shape[0], -1, -1), x], 1) + self.pos
for b in self.blocks:
x = b(x)
return self.norm(x)[:, 0]
def get_image_features(self, pixel_values, apply_rotation=None):
"""L2-normalized (B, 512) image embeddings. With the rotation
applied (default when present), outputs live in the CLIP-B/16
LAION-2B frame."""
z = F.normalize(self.head(self.forward_features(pixel_values)),
dim=-1)
rot = (self.config.apply_rotation
if apply_rotation is None else apply_rotation)
if rot and self.config.has_rotation:
z = F.normalize(z.double() @ self.rotation.double(),
dim=-1).to(z.dtype)
return z
def forward(self, pixel_values, apply_rotation=None, **kwargs):
feats = self.get_image_features(pixel_values, apply_rotation)
return BaseModelOutputWithPooling(
last_hidden_state=feats.unsqueeze(1), pooler_output=feats)