(self, args, keep_lang=False)
| 7 | |
| 8 | class ImageEncoder(torch.nn.Module): |
| 9 | def __init__(self, args, keep_lang=False): |
| 10 | super().__init__() |
| 11 | |
| 12 | self.model, self.train_preprocess, self.val_preprocess = clip.load( |
| 13 | args.model, args.device, jit=False |
| 14 | ) |
| 15 | |
| 16 | self.cache_dir = args.cache_dir |
| 17 | |
| 18 | if not keep_lang and hasattr(self.model, "transformer"): |
| 19 | delattr(self.model, "transformer") |
| 20 | |
| 21 | def forward(self, images): |
| 22 | assert self.model is not None |