File size: 3,958 Bytes
28a619a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
"""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)