yyy commited on
Commit
fe8b212
·
verified ·
1 Parent(s): 04ba9ff

Upload vision_encoder.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. vision_encoder.py +167 -0
vision_encoder.py ADDED
@@ -0,0 +1,167 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """SigLIP2-base-patch16-512 vision encoder, adapted from lusxvr/nanoVLM (MIT).
2
+ Loads real pretrained weights from google/siglip2-base-patch16-512.
3
+ """
4
+ import math
5
+ import torch
6
+ import torch.nn as nn
7
+ import torch.nn.functional as F
8
+
9
+ VIT_HIDDEN_DIM = 768
10
+ VIT_INTER_DIM = 3072
11
+ VIT_PATCH_SIZE = 16
12
+ VIT_IMG_SIZE = 512
13
+ VIT_N_HEADS = 12
14
+ VIT_N_BLOCKS = 12
15
+ VIT_LN_EPS = 1e-6
16
+ VIT_MODEL_TYPE = "google/siglip2-base-patch16-512"
17
+
18
+
19
+ class ViTPatchEmbeddings(nn.Module):
20
+ def __init__(self):
21
+ super().__init__()
22
+ self.num_patches = (VIT_IMG_SIZE // VIT_PATCH_SIZE) ** 2
23
+ self.conv = nn.Conv2d(3, VIT_HIDDEN_DIM, kernel_size=VIT_PATCH_SIZE, stride=VIT_PATCH_SIZE, padding="valid")
24
+ self.position_embedding = nn.Parameter(torch.rand(1, self.num_patches, VIT_HIDDEN_DIM))
25
+
26
+ def forward(self, x):
27
+ x = self.conv(x)
28
+ x = x.flatten(2).transpose(1, 2)
29
+ x = x + self.position_embedding
30
+ return x
31
+
32
+
33
+ class ViTMultiHeadAttention(nn.Module):
34
+ def __init__(self):
35
+ super().__init__()
36
+ self.n_heads = VIT_N_HEADS
37
+ self.head_dim = VIT_HIDDEN_DIM // VIT_N_HEADS
38
+ self.qkv_proj = nn.Linear(VIT_HIDDEN_DIM, 3 * VIT_HIDDEN_DIM, bias=True)
39
+ self.out_proj = nn.Linear(VIT_HIDDEN_DIM, VIT_HIDDEN_DIM, bias=True)
40
+
41
+ def forward(self, x):
42
+ B, T, C = x.size()
43
+ qkv = self.qkv_proj(x)
44
+ q, k, v = qkv.split(C, dim=2)
45
+ q = q.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
46
+ k = k.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
47
+ v = v.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
48
+ y = F.scaled_dot_product_attention(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=False)
49
+ y = y.transpose(1, 2).contiguous().view(B, T, C)
50
+ return self.out_proj(y)
51
+
52
+
53
+ class ViTMLP(nn.Module):
54
+ def __init__(self):
55
+ super().__init__()
56
+ self.fc1 = nn.Linear(VIT_HIDDEN_DIM, VIT_INTER_DIM)
57
+ self.fc2 = nn.Linear(VIT_INTER_DIM, VIT_HIDDEN_DIM)
58
+ self.act = nn.GELU(approximate="tanh")
59
+
60
+ def forward(self, x):
61
+ return self.fc2(self.act(self.fc1(x)))
62
+
63
+
64
+ class ViTBlock(nn.Module):
65
+ def __init__(self):
66
+ super().__init__()
67
+ self.ln1 = nn.LayerNorm(VIT_HIDDEN_DIM, eps=VIT_LN_EPS)
68
+ self.attn = ViTMultiHeadAttention()
69
+ self.ln2 = nn.LayerNorm(VIT_HIDDEN_DIM, eps=VIT_LN_EPS)
70
+ self.mlp = ViTMLP()
71
+
72
+ def forward(self, x):
73
+ x = x + self.attn(self.ln1(x))
74
+ x = x + self.mlp(self.ln2(x))
75
+ return x
76
+
77
+
78
+ class ViT(nn.Module):
79
+ def __init__(self):
80
+ super().__init__()
81
+ self.patch_embedding = ViTPatchEmbeddings()
82
+ self.blocks = nn.ModuleList([ViTBlock() for _ in range(VIT_N_BLOCKS)])
83
+ self.layer_norm = nn.LayerNorm(VIT_HIDDEN_DIM, eps=VIT_LN_EPS)
84
+
85
+ def forward(self, x):
86
+ x = self.patch_embedding(x)
87
+ for block in self.blocks:
88
+ x = block(x)
89
+ return self.layer_norm(x)
90
+
91
+ @classmethod
92
+ def from_pretrained(cls):
93
+ from huggingface_hub import hf_hub_download
94
+ import safetensors
95
+
96
+ model = cls()
97
+ safetensors_file = hf_hub_download(repo_id=VIT_MODEL_TYPE, filename="model.safetensors")
98
+ sd = model.state_dict()
99
+ mapping = {
100
+ "vision_model.embeddings.patch_embedding.weight": "patch_embedding.conv.weight",
101
+ "vision_model.embeddings.patch_embedding.bias": "patch_embedding.conv.bias",
102
+ "vision_model.embeddings.position_embedding.weight": "patch_embedding.position_embedding",
103
+ "vision_model.post_layernorm.weight": "layer_norm.weight",
104
+ "vision_model.post_layernorm.bias": "layer_norm.bias",
105
+ }
106
+ for i in range(VIT_N_BLOCKS):
107
+ mapping[f"vision_model.encoder.layers.{i}.layer_norm1.weight"] = f"blocks.{i}.ln1.weight"
108
+ mapping[f"vision_model.encoder.layers.{i}.layer_norm1.bias"] = f"blocks.{i}.ln1.bias"
109
+ mapping[f"vision_model.encoder.layers.{i}.layer_norm2.weight"] = f"blocks.{i}.ln2.weight"
110
+ mapping[f"vision_model.encoder.layers.{i}.layer_norm2.bias"] = f"blocks.{i}.ln2.bias"
111
+ mapping[f"vision_model.encoder.layers.{i}.mlp.fc1.weight"] = f"blocks.{i}.mlp.fc1.weight"
112
+ mapping[f"vision_model.encoder.layers.{i}.mlp.fc1.bias"] = f"blocks.{i}.mlp.fc1.bias"
113
+ mapping[f"vision_model.encoder.layers.{i}.mlp.fc2.weight"] = f"blocks.{i}.mlp.fc2.weight"
114
+ mapping[f"vision_model.encoder.layers.{i}.mlp.fc2.bias"] = f"blocks.{i}.mlp.fc2.bias"
115
+ mapping[f"vision_model.encoder.layers.{i}.self_attn.out_proj.weight"] = f"blocks.{i}.attn.out_proj.weight"
116
+ mapping[f"vision_model.encoder.layers.{i}.self_attn.out_proj.bias"] = f"blocks.{i}.attn.out_proj.bias"
117
+ with safetensors.safe_open(filename=safetensors_file, framework="pt", device="cpu") as f:
118
+ for hf_key, our_key in mapping.items():
119
+ tensor = f.get_tensor(hf_key)
120
+ if tensor.shape == sd[our_key].shape:
121
+ sd[our_key].copy_(tensor)
122
+ elif "position_embedding" in hf_key:
123
+ sd[our_key].copy_(tensor.unsqueeze(0))
124
+ else:
125
+ raise ValueError(f"shape mismatch {hf_key}: {tensor.shape} vs {sd[our_key].shape}")
126
+ for i in range(VIT_N_BLOCKS):
127
+ q = f.get_tensor(f"vision_model.encoder.layers.{i}.self_attn.q_proj.weight")
128
+ k = f.get_tensor(f"vision_model.encoder.layers.{i}.self_attn.k_proj.weight")
129
+ v = f.get_tensor(f"vision_model.encoder.layers.{i}.self_attn.v_proj.weight")
130
+ sd[f"blocks.{i}.attn.qkv_proj.weight"].copy_(torch.cat([q, k, v], dim=0))
131
+ qb = f.get_tensor(f"vision_model.encoder.layers.{i}.self_attn.q_proj.bias")
132
+ kb = f.get_tensor(f"vision_model.encoder.layers.{i}.self_attn.k_proj.bias")
133
+ vb = f.get_tensor(f"vision_model.encoder.layers.{i}.self_attn.v_proj.bias")
134
+ sd[f"blocks.{i}.attn.qkv_proj.bias"].copy_(torch.cat([qb, kb, vb], dim=0))
135
+ model.load_state_dict(sd)
136
+ n_params = sum(p.numel() for p in model.parameters())
137
+ print(f"Loaded {VIT_MODEL_TYPE}: {n_params:,} params")
138
+ return model
139
+
140
+
141
+ class ModalityProjector(nn.Module):
142
+ """Pixel-shuffle (x4) + linear projection into the target LLM's hidden size."""
143
+
144
+ def __init__(self, llm_hidden_size, scale_factor=4):
145
+ super().__init__()
146
+ self.scale_factor = scale_factor
147
+ self.input_dim = VIT_HIDDEN_DIM * (scale_factor ** 2)
148
+ self.output_dim = llm_hidden_size
149
+ self.proj = nn.Linear(self.input_dim, self.output_dim, bias=False)
150
+ nn.init.normal_(self.proj.weight, mean=0.0, std=0.02)
151
+
152
+ def pixel_shuffle(self, x):
153
+ bsz, seq, embed_dim = x.size()
154
+ seq_root = int(seq ** 0.5)
155
+ assert seq_root ** 2 == seq
156
+ assert seq_root % self.scale_factor == 0
157
+ h = w = seq_root
158
+ x = x.view(bsz, h, w, embed_dim)
159
+ h_out, w_out = h // self.scale_factor, w // self.scale_factor
160
+ x = x.reshape(bsz, h_out, self.scale_factor, w_out, self.scale_factor, embed_dim)
161
+ x = x.permute(0, 1, 3, 2, 4, 5).contiguous()
162
+ x = x.reshape(bsz, h_out * w_out, embed_dim * self.scale_factor ** 2)
163
+ return x
164
+
165
+ def forward(self, x):
166
+ x = x.to(self.proj.weight.dtype)
167
+ return self.proj(self.pixel_shuffle(x))