"""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"]