| 110 | |
| 111 | |
| 112 | class CLIP(nn.Module): |
| 113 | def __init__(self,image_encode, text_encode, use_allgather): |
| 114 | super().__init__() |
| 115 | self.use_allgather = use_allgather |
| 116 | self.visual =image_encode |
| 117 | self.text_encoder = text_encode |
| 118 | self.logit_scale = nn.Parameter(torch.ones([1])) |
| 119 | # self.logit_scale = nn.Parameter(torch.ones([])) |
| 120 | nn.init.constant_(self.logit_scale, np.log(1/0.07)) |
| 121 | #nn.init.normal_(self.text_projection, std=self.transformer.width ** -0.5) |
| 122 | |
| 123 | def text_parameters(self): |
| 124 | param = [self.logit_scale] |
| 125 | if self.text_encoder.text_encode_type == 'Transformer': |
| 126 | param.append(self.text_encoder.positional_embedding) |
| 127 | elif self.text_encoder.text_encode_type == 'Bert': |
| 128 | # print('Bert', self.text_encoder.text_transformer.cls.predictions, flush=True) |
| 129 | # param.extend([self.text_encoder.text_transformer.cls.predictions.decoder.weight, |
| 130 | # self.text_encoder.text_transformer.cls.predictions.bias]) |
| 131 | param.extend([self.text_encoder.text_transformer.cls.predictions.bias]) |
| 132 | return param |
| 133 | |
| 134 | def text_modules(self): |
| 135 | if self.text_encoder.text_encode_type == 'Transformer': |
| 136 | return [self.text_encoder.transformer, self.text_encoder.text_projection, self.text_encoder.token_embedding, self.text_encoder.ln_final] |
| 137 | elif self.text_encoder.text_encode_type == 'Bert': |
| 138 | # print('Bert', self.text_encoder.text_transformer, flush=True) |
| 139 | return [self.text_encoder.text_transformer.bert, self.text_encoder.text_projection, |
| 140 | self.text_encoder.text_transformer.cls.predictions.transform] |
| 141 | # self.text_encoder.text_transformer.cls.predictions.decoder, # decoder: bias |
| 142 | else: |
| 143 | import ipdb |
| 144 | ipdb.set_trace() |
| 145 | return [self.text_encoder.text_transformer, self.text_encoder.text_projection] |
| 146 | |
| 147 | def visual_parameters(self): |
| 148 | return [] |
| 149 | |
| 150 | def visual_modules(self): |
| 151 | return [self.visual] |
| 152 | |
| 153 | @property |
| 154 | def dtype(self): |
| 155 | try: |
| 156 | return self.visual.conv1.weight.dtype |
| 157 | except: |
| 158 | try: |
| 159 | return self.visual.head.weight.dtype |
| 160 | except: |
| 161 | try: |
| 162 | return self.visual.stem[0].weight.dtype |
| 163 | except: |
| 164 | return self.text_encoder.text_projection.weight.dtype |
| 165 | |
| 166 | def encode_image(self, image): |
| 167 | return self.visual(image.type(self.dtype)) |
| 168 | |
| 169 | def sample_captions(self, texts): |
nothing calls this directly
no outgoing calls
no test coverage detected