Views
No views yet
| Property | Value |
|---|---|
| Method | InfoNCE |
| Vision encoder | vit_base_patch16_224 (timm) |
| Text encoder | GPT-2 (12L / 12H / 768D) |
| Embedding dim | 512 |
| Projection head | Linear→BN→GELU→Linear (width 2048) |
| Training objective | Symmetric InfoNCE (CLIP-style) contrastive loss |
| Training data | DataComp-large |
| Training steps | 200,000 |
1import torch
2import timm
3from transformers import GPT2Config, GPT2Model, AutoTokenizer
4from safetensors.torch import load_file
5from torchvision.ops import MLP
6import torch.nn as nn
7
8HIDDEN = 768
9EMBED = 512
10
11# ── Vision encoder ───────────────────────────────────────────────────────────
12vision_encoder = timm.create_model(
13 "vit_base_patch16_224", pretrained=False, num_classes=0, dynamic_img_size=True
14)
15vision_pre_proj = nn.Sequential(
16 nn.Linear(HIDDEN, 2048), nn.BatchNorm1d(2048), nn.GELU(), nn.Linear(2048, EMBED)
17)
18
19# ── Text encoder ─────────────────────────────────────────────────────────────
20tokenizer = AutoTokenizer.from_pretrained("gpt2")
21tokenizer.pad_token = tokenizer.eos_token
22
23def tokenize_with_eos_readout(tokenizer, text, max_length=77):
24 ids = tokenizer(
25 text,
26 add_special_tokens=False,
27 truncation=True,
28 max_length=max_length - 1,
29 )["input_ids"] + [tokenizer.eos_token_id]
30 pad_len = max_length - len(ids)
31 input_ids = torch.tensor([ids + [tokenizer.pad_token_id] * pad_len])
32 attention_mask = torch.tensor([[1] * len(ids) + [0] * pad_len])
33 return dict(input_ids=input_ids, attention_mask=attention_mask)
34
35def last_unmasked_token(hidden, attention_mask):
36 lengths = attention_mask.sum(dim=1).clamp(min=1).long()
37 gather_idx = (lengths - 1).view(-1, 1, 1).expand(-1, 1, hidden.size(-1))
38 return hidden.gather(1, gather_idx).squeeze(1)
39
40text_encoder = GPT2Model(GPT2Config(
41 n_embd=HIDDEN, n_layer=12, n_head=12,
42 n_inner=HIDDEN * 4, vocab_size=tokenizer.vocab_size,
43 attn_pdrop=0.0, resid_pdrop=0.0, embd_pdrop=0.0,
44))
45text_pre_proj = nn.Sequential(
46 nn.Linear(HIDDEN, 2048), nn.BatchNorm1d(2048), nn.GELU(), nn.Linear(2048, EMBED)
47)
48
49# ── Load weights ─────────────────────────────────────────────────────────────
50from huggingface_hub import hf_hub_download
51
52vision_weights = load_file(hf_hub_download("lukaskuhndkfz/InfoNCE-ViT-B-DataComp-200k", "vision_encoder.safetensors"))
53text_weights = load_file(hf_hub_download("lukaskuhndkfz/InfoNCE-ViT-B-DataComp-200k", "text_encoder.safetensors"))
54
55vision_encoder.load_state_dict({k[len("encoder."):]: v for k, v in vision_weights.items() if k.startswith("encoder.")})
56vision_pre_proj.load_state_dict({k[len("pre_proj."):]: v for k, v in vision_weights.items() if k.startswith("pre_proj.")})
57text_encoder.load_state_dict({k[len("encoder."):]: v for k, v in text_weights.items() if k.startswith("encoder.")})
58text_pre_proj.load_state_dict({k[len("pre_proj."):]: v for k, v in text_weights.items() if k.startswith("pre_proj.")})
59
60vision_encoder.eval()
61text_encoder.eval()
62
63# ── Encode an image ──────────────────────────────────────────────────────────
64from torchvision import transforms
65from PIL import Image
66
67transform = transforms.Compose([
68 transforms.Resize(224), transforms.CenterCrop(224),
69 transforms.ToTensor(),
70 transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
71])
72
73image = Image.open("image.jpg").convert("RGB")
74pixel_values = transform(image).unsqueeze(0)
75
76with torch.no_grad():
77 image_features = vision_pre_proj(vision_encoder(pixel_values)) # (1, 512)
78
79# ── Encode a caption ─────────────────────────────────────────────────────────
80inputs = tokenize_with_eos_readout(tokenizer, "a photo of a cat")
81with torch.no_grad():
82 hidden = text_encoder(**inputs).last_hidden_state
83 text_hidden = last_unmasked_token(hidden, inputs["attention_mask"])
84 text_features = text_pre_proj(text_hidden) # (1, 512)| File | Contents |
|---|---|
vision_encoder.safetensors | Vision encoder (encoder.*), pre-projection head (pre_proj.*), and cross-modal projector MLP (projector.*) |
text_encoder.safetensors | Text encoder (encoder.*), pre-projection head (pre_proj.*), and cross-modal projector MLP (projector.*) |
config.json | Architecture and training hyperparameters |