CoolFace
Modelpublic

bilzepython/minimind-3v

sourceHugging Faceapache-2.0updated 6mo agoView on Hugging Face
0likes9downloads
model_vlm.py156 linesDownload Raw Back to root
1import os2import torch3import warnings4from .model_minimind import *5from typing import Optional, Tuple, List, Union6from torch import nn7from transformers import Siglip2ImageProcessor, Siglip2VisionModel8from transformers.modeling_outputs import MoeCausalLMOutputWithPast9 10warnings.filterwarnings('ignore')11 12 13class VLMConfig(MiniMindConfig):14    model_type = "minimind-v"15 16    def __init__(self, image_special_token='<|image_pad|>', image_ids=[12], **kwargs):17        self.image_special_token = image_special_token18        self.image_ids = image_ids19        self.image_hidden_size = kwargs.get("image_hidden_size", 768)20        self.image_token_len = kwargs.get("image_token_len", 64)21        super().__init__(**kwargs)22 23class MMVisionProjector(nn.Module):24    def __init__(self, in_dim, out_dim, source_tokens=256, target_tokens=64):25        super().__init__()26        self.target_tokens = target_tokens27        self.merge = source_tokens // target_tokens28        self.mlp = nn.Sequential(29            nn.Linear(in_dim * self.merge, out_dim),30            nn.GELU(),31            nn.Linear(out_dim, out_dim),32        )33    def forward(self, x):34        b, n, d = x.shape35        x = x.reshape(b, self.target_tokens, d * self.merge)36        return self.mlp(x)37 38# 继承自语言模型39class MiniMindVLM(MiniMindForCausalLM):40    config_class = VLMConfig41 42    def __init__(self, config: VLMConfig = None, vision_model_path="./model/siglip2-base-p16-ve"):43        self.config = config or VLMConfig()44        super().__init__(self.config)45        self.vision_encoder, self.processor = self.__class__.get_vision_model(vision_model_path)46        self.vision_proj = MMVisionProjector(self.config.image_hidden_size, self.config.hidden_size, target_tokens=self.config.image_token_len)47 48    @staticmethod49    def get_vision_model(model_path: str):50        from transformers import logging as hf_logging51        hf_logging.set_verbosity_error()52        if not os.path.exists(model_path):53            return None, None54        model = Siglip2VisionModel.from_pretrained(model_path)55        processor = Siglip2ImageProcessor.from_pretrained(model_path)56        # 冻结 vision_encoder 的所有参数57        for param in model.parameters():58            param.requires_grad = False59        return model.eval(), processor60 61    @staticmethod62    def image2tensor(image, processor):63        if image.mode in ['RGBA', 'LA']: image = image.convert('RGB')64        inputs = processor(images=image, return_tensors="pt")65        return inputs66 67    @staticmethod68    def get_image_embeddings(image_inputs, vision_model):69        if hasattr(image_inputs, 'keys'):70            image_inputs = {k: v.squeeze(1) if v.ndim > 2 and v.shape[1] == 1 else v for k, v in image_inputs.items()}71        with torch.no_grad():72            outputs = vision_model(**image_inputs)73        return outputs.last_hidden_state74 75    @torch.compiler.disable76    def count_vision_proj(self, tokens, h, vision_tensors=None, seqlen=512):77        if vision_tensors is None or not self.config.image_ids:78            return h79        marker, vf = self.config.image_ids[0], vision_tensors80        if vf.dim() == 3:81            vf = vf.unsqueeze(1)82        out = []83        for b in range(h.size(0)):84            hb, seq, k, i = h[b], tokens[b].tolist(), 0, 085            while i < len(seq):86                if seq[i] == marker:87                    start = i88                    while i < len(seq) and seq[i] == marker:89                        i += 190                    if k < vf.size(1):91                        hb = torch.cat((hb[:start], vf[b][k][:i - start], hb[i:]), dim=0)[:seqlen]92                        k += 193                else:94                    i += 195            out.append(hb)96        return torch.stack(out)97 98    def forward(self,99                input_ids: Optional[torch.Tensor] = None,100                attention_mask: Optional[torch.Tensor] = None,101                past_key_values: Optional[List[Tuple[torch.Tensor, torch.Tensor]]] = None,102                use_cache: bool = False,103                logits_to_keep: Union[int, torch.Tensor] = 0,104                labels: Optional[torch.Tensor] = None,105                pixel_values: Optional[torch.FloatTensor] = None,106                **args):107        batch_size, seq_length = input_ids.shape108        if hasattr(past_key_values, 'layers'): past_key_values = None109        past_key_values = past_key_values or [None] * len(self.model.layers)110        start_pos = past_key_values[0][0].shape[1] if past_key_values[0] is not None else 0111 112        hidden_states = self.model.dropout(self.model.embed_tokens(input_ids))113 114        if pixel_values is not None and start_pos == 0:115            if hasattr(pixel_values, 'keys'):116                img_emb = MiniMindVLM.get_image_embeddings(pixel_values, self.vision_encoder)117                vision_tensors = self.vision_proj(img_emb)118            else:119                if len(pixel_values.shape) == 6:120                    pixel_values = pixel_values.squeeze(2)121                bs, num, c, im_h, im_w = pixel_values.shape122                stack_dim = 1 if bs > 1 else 0123                vision_tensors = torch.stack([self.vision_proj(MiniMindVLM.get_image_embeddings(pixel_values[:, i, :, :, :], self.vision_encoder)) for i in range(num)], dim=stack_dim)124            hidden_states = self.count_vision_proj(tokens=input_ids, h=hidden_states, vision_tensors=vision_tensors, seqlen=input_ids.shape[1])125 126        position_embeddings = (127            self.model.freqs_cos[start_pos:start_pos + seq_length],128            self.model.freqs_sin[start_pos:start_pos + seq_length]129        )130 131        presents = []132        for layer_idx, (layer, past_key_value) in enumerate(zip(self.model.layers, past_key_values)):133            hidden_states, present = layer(134                hidden_states,135                position_embeddings,136                past_key_value=past_key_value,137                use_cache=use_cache,138                attention_mask=attention_mask139            )140            presents.append(present)141 142        hidden_states = self.model.norm(hidden_states)143 144        aux_loss = sum([l.mlp.aux_loss for l in self.model.layers if isinstance(l.mlp, MOEFeedForward)], hidden_states.new_zeros(1).squeeze())145        slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep146        logits = self.lm_head(hidden_states[:, slice_indices, :])147 148        loss = None149        if labels is not None:150            shift_logits = logits[..., :-1, :].contiguous()151            shift_labels = labels[..., 1:].contiguous()152            loss = F.cross_entropy(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1), ignore_index=-100)153 154        output = MoeCausalLMOutputWithPast(loss=loss, aux_loss=aux_loss, logits=logits, past_key_values=presents, hidden_states=hidden_states)155        return output156