| 94 | import argparse |
| 95 | |
| 96 | class CLIP(nn.Module): |
| 97 | def __init__(self, args, layernorm_epsilon=1e-5): |
| 98 | super().__init__() |
| 99 | self.image_encoder = ImageEncoder(args, layernorm_epsilon=layernorm_epsilon) |
| 100 | text_args = argparse.Namespace(**vars(args)) |
| 101 | override_attrs = ['vocab_size', 'num_layers', 'hidden_size', 'num_attention_heads', 'layernorm_order', |
| 102 | 'max_sequence_length', 'inner_hidden_size', 'hidden_size_per_attention_head'] |
| 103 | for name in override_attrs: |
| 104 | text_attr = getattr(text_args, 'text_' + name, None) |
| 105 | if text_attr is not None: # else use encoder-config |
| 106 | setattr(text_args, name, text_attr) |
| 107 | self.text_encoder = TextEncoder(text_args, layernorm_epsilon=layernorm_epsilon) |
| 108 | self.logit_scale = nn.Parameter(torch.ones([]) * args.logit_scale_init_value) |
| 109 | |
| 110 | def encode_image(self, input_ids, position_ids, attention_mask=None, **kw_args): |
| 111 | return self.image_encoder(input_ids, position_ids, attention_mask, **kw_args) |
| 112 | |
| 113 | def encode_text(self, input_ids, position_ids, attention_mask, **kw_args): |
| 114 | return self.text_encoder(input_ids, position_ids, attention_mask, **kw_args) |
| 115 | |
| 116 | def reinit(self, mixin_names): # please use different mixin names for two encoders |
| 117 | self.image_encoder.reinit(mixin_names) |
| 118 | self.text_encoder.reinit(mixin_names) |
| 119 | |
| 120 | def forward(self, image_input_ids, image_position_ids, text_input_ids, text_position_ids, *, image_attention_mask=None, text_attention_mask=None, **kw_args): |
| 121 | image_embeds, *image_mems = self.encode_image(image_input_ids, image_position_ids, attention_mask=image_attention_mask, **kw_args) |
| 122 | text_embeds, *text_mems = self.encode_text(text_input_ids, text_position_ids, attention_mask=text_attention_mask, **kw_args) |
| 123 | |
| 124 | # normalized features |
| 125 | image_embeds = image_embeds / image_embeds.norm(dim=-1, keepdim=True) |
| 126 | text_embeds = text_embeds / text_embeds.norm(dim=-1, keepdim=True) |
| 127 | |
| 128 | # cosine similarity as logits |
| 129 | logit_scale = self.logit_scale.exp() |
| 130 | logits_per_text = torch.matmul(text_embeds, image_embeds.t()) * logit_scale |
| 131 | logits_per_image = logits_per_text.T |
| 132 | return image_embeds, text_embeds, logits_per_text, logits_per_image |
| 133 | |
| 134 | @classmethod |
| 135 | def add_model_specific_args(cls, parser): |
| 136 | group = parser.add_argument_group('SiameseModel', 'CLIP') |
| 137 | group.add_argument("--text-layernorm-order", type=str, default=None) |
| 138 | group.add_argument("--text-num-layers", type=int, default=None) |
| 139 | group.add_argument("--text-hidden-size", type=int, default=None) |
| 140 | group.add_argument("--text-num-attention-heads", type=int, default=None) |
| 141 | group.add_argument("--text-max-sequence-length", type=int, default=None) |
| 142 | group.add_argument("--text-inner-hidden-size", type=int, default=None) |
| 143 | group.add_argument("--text-hidden-size-per-attention-head", type=int, default=None) |
| 144 | group.add_argument("--logit-scale-init-value", type=float, default=None) |
| 145 | return parser |
| 146 | |
| 147 | @classmethod |
| 148 | def from_pretrained(cls, args, name, *, path=None, url=None): |
| 149 | model_path = auto_create(name, path=path, url=url) |
| 150 | args = update_args_with_file(args, path=os.path.join(model_path, 'model_config.json')) |
| 151 | model = get_model(args, cls) |
| 152 | load_checkpoint(model, args, load_path=model_path) |
| 153 | return model, args |
no outgoing calls
no test coverage detected