| 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) |