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