CoolFace
Apppublic

XaviXva/Video-LLaVA

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
0likes
mae_encoder.py81 linesDownload Raw Back to multimodal_encoder
1import torch2import torch.nn as nn3 4from transformers import ViTMAEForPreTraining, AutoConfig, AutoImageProcessor5 6 7class MAEVisionTower(nn.Module):8    def __init__(self, vision_tower, args, cache_dir='./cache_dir', delay_load=False):9        super().__init__()10 11        self.is_loaded = False12        self.cache_dir = cache_dir13        self.vision_tower_name = vision_tower14        self.select_layer = args.mm_vision_select_layer15        self.select_feature = getattr(args, 'mm_vision_select_feature', 'patch')16 17        if not delay_load:18            self.load_model()19        else:20            self.cfg_only = AutoConfig.from_pretrained(self.vision_tower_name, cache_dir=self.cache_dir)21 22    def load_model(self):23        self.image_processor = AutoImageProcessor.from_pretrained(self.vision_tower_name, cache_dir=self.cache_dir)24        vision_tower = ViTMAEForPreTraining.from_pretrained(self.vision_tower_name, cache_dir=self.cache_dir)25        self.vision_tower = vision_tower.vit26        self.vision_tower.requires_grad_(False)27 28        self.is_loaded = True29 30    def feature_select(self, image_forward_outs):31        image_features = image_forward_outs.hidden_states[self.select_layer]32        if self.select_feature == 'patch':33            image_features = image_features[:, 1:]34        elif self.select_feature == 'cls_patch':35            image_features = image_features36        else:37            raise ValueError(f'Unexpected select feature: {self.select_feature}')38        # print(image_features.shape)39        return image_features40 41    @torch.no_grad()42    def forward(self, images):43        if type(images) is list:44            image_features = []45            for image in images:46                image_forward_out = self.vision_tower(image.to(device=self.device, dtype=self.dtype).unsqueeze(0), output_hidden_states=True)47                image_feature = self.feature_select(image_forward_out).to(image.dtype)48                image_features.append(image_feature)49        else:50            image_forward_outs = self.vision_tower(images.to(device=self.device, dtype=self.dtype), output_hidden_states=True)51            image_features = self.feature_select(image_forward_outs).to(images.dtype)52 53        return image_features54 55    @property56    def dummy_feature(self):57        return torch.zeros(1, self.hidden_size, device=self.device, dtype=self.dtype)58 59    @property60    def dtype(self):61        return self.vision_tower.dtype62 63    @property64    def device(self):65        return self.vision_tower.device66 67    @property68    def config(self):69        if self.is_loaded:70            return self.vision_tower.config71        else:72            return self.cfg_only73 74    @property75    def hidden_size(self):76        return self.config.hidden_size77 78    @property79    def num_patches(self):80        return (self.config.image_size // self.config.patch_size) ** 281