Buckets:
| import torch | |
| import open_clip | |
| import utils | |
| class ImageEncoder(torch.nn.Module): | |
| def __init__(self, args, keep_lang=False): | |
| super().__init__() | |
| print(f'Loading {args.model} pre-trained weights.') | |
| if '__pretrained__' in args.model: | |
| name, pretrained = args.model.split('__pretrained__') | |
| else: | |
| name = args.model | |
| pretrained = 'openai' | |
| self.model, self.train_preprocess, self.val_preprocess = open_clip.create_model_and_transforms( | |
| name, pretrained=pretrained, cache_dir=args.openclip_cachedir) | |
| self.cache_dir = args.cache_dir | |
| if not keep_lang and hasattr(self.model, 'transformer'): | |
| delattr(self.model, 'transformer') | |
| def forward(self, images): | |
| assert self.model is not None | |
| return self.model.encode_image(images) | |
| def __call__(self, inputs): | |
| return self.forward(inputs) | |
| def save(self, filename): | |
| print(f'Saving image encoder to {filename}') | |
| utils.torch_save(self, filename) | |
| def load(cls, model_name, filename): | |
| print(f'Loading image encoder from {filename}') | |
| state_dict = torch.load(filename) | |
| return cls.load(model_name, state_dict) | |
| def load_from_state_dict(cls, model_name, state_dict): | |
| self.model, self.train_preprocess, self.val_preprocess = open_clip.create_model_and_transforms( | |
| name, pretrained=pretrained, cache_dir=args.openclip_cachedir) | |
| self.model.load_from_state_dict(state_dict) | |
| class ClassificationHead(torch.nn.Linear): | |
| def __init__(self, normalize, weights, biases=None): | |
| output_size, input_size = weights.shape | |
| super().__init__(input_size, output_size) | |
| self.normalize = normalize | |
| if weights is not None: | |
| self.weight = torch.nn.Parameter(weights.clone()) | |
| if biases is not None: | |
| self.bias = torch.nn.Parameter(biases.clone()) | |
| else: | |
| self.bias = torch.nn.Parameter(torch.zeros_like(self.bias)) | |
| def forward(self, inputs): | |
| if self.normalize: | |
| inputs = inputs / inputs.norm(dim=-1, keepdim=True) | |
| return super().forward(inputs) | |
| def __call__(self, inputs): | |
| return self.forward(inputs) | |
| def save(self, filename): | |
| print(f'Saving classification head to {filename}') | |
| utils.torch_save(self, filename) | |
| def load(cls, filename): | |
| print(f'Loading classification head from {filename}') | |
| return utils.torch_load(filename) | |
| class ImageClassifier(torch.nn.Module): | |
| def __init__(self, image_encoder, classification_head): | |
| super().__init__() | |
| self.image_encoder = image_encoder | |
| self.classification_head = classification_head | |
| if self.image_encoder is not None: | |
| if hasattr(self.image_encoder, 'train_preprocess'): | |
| self.train_preprocess = self.image_encoder.train_preprocess | |
| self.val_preprocess = self.image_encoder.val_preprocess | |
| elif hasattr(self.image_encoder.model, 'train_preprocess'): | |
| self.train_preprocess = self.image_encoder.model.train_preprocess | |
| self.val_preprocess = self.image_encoder.model.val_preprocess | |
| def freeze_head(self): | |
| self.classification_head.weight.requires_grad_(False) | |
| self.classification_head.bias.requires_grad_(False) | |
| def forward(self, inputs): | |
| features = self.image_encoder(inputs) | |
| outputs = self.classification_head(features) | |
| return outputs | |
| def __call__(self, inputs): | |
| return self.forward(inputs) | |
| def save(self, filename): | |
| print(f'Saving image classifier to {filename}') | |
| utils.torch_save(self, filename) | |
| def load(cls, filename): | |
| print(f'Loading image classifier from {filename}') | |
| return utils.torch_load(filename) | |
| class ImageClassifier_debug(torch.nn.Module): | |
| def __init__(self, image_encoder, image_encoder2, classification_head): | |
| super().__init__() | |
| self.image_encoder = image_encoder | |
| self.image_encoder2 = image_encoder2 | |
| self.classification_head = classification_head | |
| if self.image_encoder is not None: | |
| self.train_preprocess = self.image_encoder.train_preprocess | |
| self.val_preprocess = self.image_encoder.val_preprocess | |
| def freeze_head(self): | |
| self.classification_head.weight.requires_grad_(False) | |
| self.classification_head.bias.requires_grad_(False) | |
| def forward(self, inputs): | |
| features = self.image_encoder(inputs) | |
| features2 = self.image_encoder2(inputs) | |
| outputs = self.classification_head(features + features2) | |
| return outputs | |
| def __call__(self, inputs): | |
| return self.forward(inputs) | |
| def save(self, filename): | |
| print(f'Saving image classifier to {filename}') | |
| utils.torch_save(self, filename) | |
| def load(cls, filename): | |
| print(f'Loading image classifier from {filename}') | |
| return utils.torch_load(filename) | |
| class MultiHeadImageClassifier(torch.nn.Module): | |
| def __init__(self, image_encoder, classification_heads): | |
| super().__init__() | |
| self.image_encoder = image_encoder | |
| self.classification_heads = torch.nn.ModuleList(classification_heads) | |
| if self.image_encoder is not None: | |
| self.train_preprocess = self.image_encoder.train_preprocess | |
| self.val_preprocess = self.image_encoder.val_preprocess | |
| def freeze_head(self): | |
| for idx in range(len(self.classification_heads)): | |
| self.classification_heads[idx].weight.requires_grad_(False) | |
| self.classification_heads[idx].bias.requires_grad_(False) | |
| def forward(self, inputs, head_idx): | |
| features = self.image_encoder(inputs) | |
| outputs = self.classification_heads[head_idx](features) | |
| return outputs | |
| def __call__(self, inputs, head_idx): | |
| return self.forward(inputs, head_idx) | |
| def save(self, filename): | |
| print(f'Saving image classifier to {filename}') | |
| utils.torch_save(self, filename) | |
| def load(cls, filename): | |
| print(f'Loading image classifier from {filename}') | |
| return utils.torch_load(filename) | |
Xet Storage Details
- Size:
- 6.36 kB
- Xet hash:
- b6df25f64f8256065c073f770d68fc305029e12eebecf3f22182e7a0d323a032
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.