Uses the OpenCLIP transformer encoder for text
| 132 | |
| 133 | |
| 134 | class FrozenOpenCLIPEmbedder(AbstractEncoder): |
| 135 | """ |
| 136 | Uses the OpenCLIP transformer encoder for text |
| 137 | """ |
| 138 | LAYERS = [ |
| 139 | #"pooled", |
| 140 | "last", |
| 141 | "penultimate" |
| 142 | ] |
| 143 | def __init__(self, arch="ViT-H-14", |
| 144 | version="/opt/data/private/QiuKunpeng/Diffusion/ControlNet/models/CLIP-ViT-H-14-laion2B-s32B-b79K/open_clip_pytorch_model.bin", |
| 145 | device="cuda", max_length=77, |
| 146 | freeze=True, layer="last"): |
| 147 | super().__init__() |
| 148 | assert layer in self.LAYERS |
| 149 | model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'), pretrained=version) |
| 150 | del model.visual |
| 151 | self.model = model |
| 152 | |
| 153 | self.device = device |
| 154 | self.max_length = max_length |
| 155 | if freeze: |
| 156 | self.freeze() |
| 157 | self.layer = layer |
| 158 | if self.layer == "last": |
| 159 | self.layer_idx = 0 |
| 160 | elif self.layer == "penultimate": |
| 161 | self.layer_idx = 1 |
| 162 | else: |
| 163 | raise NotImplementedError() |
| 164 | |
| 165 | def freeze(self): |
| 166 | self.model = self.model.eval() |
| 167 | for param in self.parameters(): |
| 168 | param.requires_grad = False |
| 169 | |
| 170 | def forward(self, text): |
| 171 | tokens = open_clip.tokenize(text) |
| 172 | z = self.encode_with_transformer(tokens.to(self.device)) |
| 173 | return z |
| 174 | |
| 175 | def encode_with_transformer(self, text): |
| 176 | x = self.model.token_embedding(text) # [batch_size, n_ctx, d_model] |
| 177 | x = x + self.model.positional_embedding |
| 178 | x = x.permute(1, 0, 2) # NLD -> LND |
| 179 | x = self.text_transformer_forward(x, attn_mask=self.model.attn_mask) |
| 180 | x = x.permute(1, 0, 2) # LND -> NLD |
| 181 | x = self.model.ln_final(x) |
| 182 | return x |
| 183 | |
| 184 | def text_transformer_forward(self, x: torch.Tensor, attn_mask = None): |
| 185 | for i, r in enumerate(self.model.transformer.resblocks): |
| 186 | if i == len(self.model.transformer.resblocks) - self.layer_idx: |
| 187 | break |
| 188 | if self.model.transformer.grad_checkpointing and not torch.jit.is_scripting(): |
| 189 | x = checkpoint(r, x, attn_mask) |
| 190 | else: |
| 191 | x = r(x, attn_mask=attn_mask) |
nothing calls this directly
no outgoing calls
no test coverage detected