MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / CLIP

Class CLIP

SwissArmyTransformer/sat/model/official/clip_model.py:96–153  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

94import argparse
95
96class 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

Callers 2

transform_param.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected