RemoteCLIP / model /remoteclip.py
zhangrenchao's picture
Update RemoteCLIP model package
2ca760b verified
Raw
History Blame Contribute Delete
4.79 kB
"""RemoteCLIP dual encoder with ViT and causal text Transformer."""
import math
import torch
from torch import nn
from torch.nn import functional as F
class VisionTransformer(nn.Module):
def __init__(self, image_size, patch_size, width, layers, heads, output_dim):
super().__init__()
if image_size % patch_size:
raise ValueError("image_size must be divisible by patch_size")
patches = (image_size // patch_size) ** 2
self.patch_embed = nn.Conv2d(3, width, patch_size, patch_size, bias=False)
self.class_embedding = nn.Parameter(torch.empty(1, 1, width))
self.position_embedding = nn.Parameter(torch.empty(1, patches + 1, width))
layer = nn.TransformerEncoderLayer(
width, heads, width * 4, activation="gelu", batch_first=True,
norm_first=True, dropout=0.0,
)
self.transformer = nn.TransformerEncoder(layer, layers)
self.norm = nn.LayerNorm(width)
self.projection = nn.Parameter(torch.empty(width, output_dim))
nn.init.normal_(self.class_embedding, std=width ** -0.5)
nn.init.normal_(self.position_embedding, std=width ** -0.5)
nn.init.normal_(self.projection, std=width ** -0.5)
def forward(self, images):
tokens = self.patch_embed(images).flatten(2).transpose(1, 2)
cls = self.class_embedding.expand(images.shape[0], -1, -1)
tokens = torch.cat((cls, tokens), dim=1) + self.position_embedding
return self.norm(self.transformer(tokens)[:, 0]) @ self.projection
class RemoteCLIP(nn.Module):
"""CLIP-compatible encoders; EOT is the largest token id in each sequence."""
def __init__(
self,
vocabulary_size=49408,
context_length=77,
eot_token_id=49407,
image_size=224,
patch_size=32,
embed_dim=64,
vision_width=64,
vision_layers=2,
vision_heads=4,
text_width=64,
text_layers=2,
text_heads=4,
):
super().__init__()
self.context_length = context_length
self.eot_token_id = eot_token_id
self.visual = VisionTransformer(
image_size, patch_size, vision_width, vision_layers, vision_heads, embed_dim
)
self.token_embedding = nn.Embedding(vocabulary_size, text_width, padding_idx=0)
self.position_embedding = nn.Parameter(torch.empty(context_length, text_width))
text_layer = nn.TransformerEncoderLayer(
text_width, text_heads, text_width * 4, activation="gelu",
batch_first=True, norm_first=True, dropout=0.0,
)
self.text_transformer = nn.TransformerEncoder(text_layer, text_layers)
self.text_norm = nn.LayerNorm(text_width)
self.text_projection = nn.Parameter(torch.empty(text_width, embed_dim))
self.logit_scale = nn.Parameter(torch.tensor(math.log(1 / 0.07)))
nn.init.normal_(self.position_embedding, std=0.01)
nn.init.normal_(self.text_projection, std=text_width ** -0.5)
def encode_image(self, images):
if images.ndim != 4 or images.shape[1:] != (3, 224, 224):
raise ValueError("images must have paper-compatible shape [B,3,224,224]")
return F.normalize(self.visual(images), dim=-1)
def encode_text(self, tokens):
if tokens.ndim != 2 or tokens.shape[1] != self.context_length:
raise ValueError(f"tokens must have shape [B,{self.context_length}]")
causal_mask = torch.full(
(self.context_length, self.context_length), float("-inf"), device=tokens.device
).triu_(1)
features = self.token_embedding(tokens) + self.position_embedding
features = self.text_norm(self.text_transformer(features, mask=causal_mask))
eot_positions = tokens.eq(self.eot_token_id).to(torch.int64).argmax(dim=-1)
pooled = features[torch.arange(tokens.shape[0], device=tokens.device), eot_positions]
return F.normalize(pooled @ self.text_projection, dim=-1)
def forward(self, images, tokens):
return self.encode_image(images), self.encode_text(tokens), self.logit_scale.exp().clamp(max=100)
def multi_positive_clip_loss(image_features, text_features, pair_ids, logit_scale):
"""Symmetric CLIP loss where all samples sharing pair_id are positives."""
logits = logit_scale * image_features @ text_features.t()
positives = pair_ids[:, None].eq(pair_ids[None, :])
log_i = F.log_softmax(logits, dim=1)
log_t = F.log_softmax(logits.t(), dim=1)
loss_i = -(log_i.masked_fill(~positives, 0).sum(1) / positives.sum(1))
loss_t = -(log_t.masked_fill(~positives.t(), 0).sum(1) / positives.t().sum(1))
return (loss_i.mean() + loss_t.mean()) / 2
__all__ = ["RemoteCLIP", "multi_positive_clip_loss"]