CoolFace
Modelpublic

MathLLMs/MathCoder-VL-2B

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
7likes30downloads
modeling_internvl_chat.py359 linesDownload Raw Back to root
1# --------------------------------------------------------2# InternVL3# Copyright (c) 2023 OpenGVLab4# Licensed under The MIT License [see LICENSE for details]5# --------------------------------------------------------6import warnings7from typing import Any, List, Optional, Tuple, Union8 9import torch.utils.checkpoint10from peft import LoraConfig, get_peft_model11from torch import nn12from torch.nn import CrossEntropyLoss13from transformers import (AutoModel, GenerationConfig, LlamaForCausalLM,14                          LlamaTokenizer)15from transformers.modeling_outputs import CausalLMOutputWithPast16from transformers.modeling_utils import PreTrainedModel17from transformers.utils import ModelOutput, logging18 19from .configuration_internvl_chat import InternVLChatConfig20from .modeling_intern_vit import InternVisionModel21from .modeling_internlm2 import InternLM2ForCausalLM22 23logger = logging.get_logger(__name__)24 25 26class InternVLChatModel(PreTrainedModel):27    config_class = InternVLChatConfig28    main_input_name = 'pixel_values'29    _no_split_modules = ['InternVisionEncoderLayer', 'LlamaDecoderLayer', 'InternLM2DecoderLayer']30 31    def __init__(self, config: InternVLChatConfig, vision_model=None, language_model=None):32        super().__init__(config)33 34        image_size = config.force_image_size or config.vision_config.image_size35        patch_size = config.vision_config.patch_size36        self.patch_size = patch_size37        self.select_layer = config.select_layer38        self.template = config.template39        self.num_image_token = int((image_size // patch_size) ** 2 * (config.downsample_ratio ** 2))40        self.downsample_ratio = config.downsample_ratio41        self.ps_version = config.ps_version42 43        logger.info(f'num_image_token: {self.num_image_token}')44        logger.info(f'ps_version: {self.ps_version}')45        if vision_model is not None:46            self.vision_model = vision_model47        else:48            self.vision_model = InternVisionModel(config.vision_config)49        if language_model is not None:50            self.language_model = language_model51        else:52            if config.llm_config.architectures[0] == 'LlamaForCausalLM':53                self.language_model = LlamaForCausalLM(config.llm_config)54            elif config.llm_config.architectures[0] == 'InternLM2ForCausalLM':55                self.language_model = InternLM2ForCausalLM(config.llm_config)56            else:57                raise NotImplementedError(f'{config.llm_config.architectures[0]} is not implemented.')58 59        vit_hidden_size = config.vision_config.hidden_size60        llm_hidden_size = config.llm_config.hidden_size61 62        self.mlp1 = nn.Sequential(63            nn.LayerNorm(vit_hidden_size * int(1 / self.downsample_ratio) ** 2),64            nn.Linear(vit_hidden_size * int(1 / self.downsample_ratio) ** 2, llm_hidden_size),65            nn.GELU(),66            nn.Linear(llm_hidden_size, llm_hidden_size)67        )68 69        # if config.force_image_size != config.vision_config.image_size:70        #     self.vision_model.resize_pos_embeddings(71        #         old_size=config.vision_config.image_size,72        #         new_size=config.force_image_size,73        #         patch_size=config.vision_config.patch_size74        #     )75 76        self.img_context_token_id = None77        self.neftune_alpha = None78 79        if config.use_backbone_lora:80            self.wrap_backbone_lora(r=config.use_backbone_lora, lora_alpha=2 * config.use_backbone_lora)81 82        if config.use_llm_lora:83            self.wrap_llm_lora(r=config.use_llm_lora, lora_alpha=2 * config.use_llm_lora)84 85    def wrap_backbone_lora(self, r=128, lora_alpha=256, lora_dropout=0.05):86        lora_config = LoraConfig(87            r=r,88            target_modules=['attn.qkv', 'attn.proj', 'mlp.fc1', 'mlp.fc2'],89            lora_alpha=lora_alpha,90            lora_dropout=lora_dropout,91        )92        self.vision_model = get_peft_model(self.vision_model, lora_config)93        self.vision_model.print_trainable_parameters()94 95    def wrap_llm_lora(self, r=128, lora_alpha=256, lora_dropout=0.05):96        lora_config = LoraConfig(97            r=r,98            target_modules=['self_attn.q_proj', 'self_attn.k_proj', 'self_attn.v_proj', 'self_attn.o_proj',99                            'mlp.gate_proj', 'mlp.down_proj', 'mlp.up_proj'],100            lora_alpha=lora_alpha,101            lora_dropout=lora_dropout,102            task_type='CAUSAL_LM'103        )104        self.language_model = get_peft_model(self.language_model, lora_config)105        self.language_model.enable_input_require_grads()106        self.language_model.print_trainable_parameters()107 108    def forward(109            self,110            pixel_values: torch.FloatTensor,111            input_ids: torch.LongTensor = None,112            attention_mask: Optional[torch.Tensor] = None,113            position_ids: Optional[torch.LongTensor] = None,114            image_flags: Optional[torch.LongTensor] = None,115            past_key_values: Optional[List[torch.FloatTensor]] = None,116            labels: Optional[torch.LongTensor] = None,117            use_cache: Optional[bool] = None,118            output_attentions: Optional[bool] = None,119            output_hidden_states: Optional[bool] = None,120            return_dict: Optional[bool] = None,121    ) -> Union[Tuple, CausalLMOutputWithPast]:122        return_dict = return_dict if return_dict is not None else self.config.use_return_dict123 124        image_flags = image_flags.squeeze(-1)125        input_embeds = self.language_model.get_input_embeddings()(input_ids)126 127        vit_embeds = self.extract_feature(pixel_values)128        vit_embeds = vit_embeds[image_flags == 1]129        vit_batch_size = pixel_values.shape[0]130 131        B, N, C = input_embeds.shape132        input_embeds = input_embeds.reshape(B * N, C)133 134        if torch.distributed.get_rank() == 0:135            print(f'dynamic ViT batch size: {vit_batch_size}, images per sample: {vit_batch_size / B}, dynamic token length: {N}')136 137        input_ids = input_ids.reshape(B * N)138        selected = (input_ids == self.img_context_token_id)139        try:140            input_embeds[selected] = input_embeds[selected] * 0.0 + vit_embeds.reshape(-1, C)141        except Exception as e:142            vit_embeds = vit_embeds.reshape(-1, C)143            print(f'warning: {e}, input_embeds[selected].shape={input_embeds[selected].shape}, '144                  f'vit_embeds.shape={vit_embeds.shape}')145            n_token = selected.sum()146            input_embeds[selected] = input_embeds[selected] * 0.0 + vit_embeds[:n_token]147 148        input_embeds = input_embeds.reshape(B, N, C)149 150        outputs = self.language_model(151            inputs_embeds=input_embeds,152            attention_mask=attention_mask,153            position_ids=position_ids,154            past_key_values=past_key_values,155            use_cache=use_cache,156            output_attentions=output_attentions,157            output_hidden_states=output_hidden_states,158            return_dict=return_dict,159        )160        logits = outputs.logits161 162        loss = None163        if labels is not None:164            # Shift so that tokens < n predict n165            shift_logits = logits[..., :-1, :].contiguous()166            shift_labels = labels[..., 1:].contiguous()167            # Flatten the tokens168            loss_fct = CrossEntropyLoss()169            shift_logits = shift_logits.view(-1, self.language_model.config.vocab_size)170            shift_labels = shift_labels.view(-1)171            # Enable model parallelism172            shift_labels = shift_labels.to(shift_logits.device)173            loss = loss_fct(shift_logits, shift_labels)174 175        if not return_dict:176            output = (logits,) + outputs[1:]177            return (loss,) + output if loss is not None else output178 179        return CausalLMOutputWithPast(180            loss=loss,181            logits=logits,182            past_key_values=outputs.past_key_values,183            hidden_states=outputs.hidden_states,184            attentions=outputs.attentions,185        )186 187    def pixel_shuffle(self, x, scale_factor=0.5):188        n, w, h, c = x.size()189        # N, W, H, C --> N, W, H * scale, C // scale190        x = x.view(n, w, int(h * scale_factor), int(c / scale_factor))191        # N, W, H * scale, C // scale --> N, H * scale, W, C // scale192        x = x.permute(0, 2, 1, 3).contiguous()193        # N, H * scale, W, C // scale --> N, H * scale, W * scale, C // (scale ** 2)194        x = x.view(n, int(h * scale_factor), int(w * scale_factor),195                   int(c / (scale_factor * scale_factor)))196        if self.ps_version == 'v1':197            warnings.warn("In ps_version 'v1', the height and width have not been swapped back, "198                          'which results in a transposed image.')199        else:200            x = x.permute(0, 2, 1, 3).contiguous()201        return x202 203    def noised_embed(self, vit_embeds, noise_alpha=5):204        dims = torch.tensor(vit_embeds.size(1) * vit_embeds.size(2))205        mag_norm = noise_alpha / torch.sqrt(dims)206        noise = torch.zeros_like(vit_embeds).uniform_(-mag_norm, mag_norm)207        return vit_embeds + noise208 209    def extract_feature(self, pixel_values):210        if self.select_layer == -1:211            vit_embeds = self.vision_model(212                pixel_values=pixel_values,213                output_hidden_states=False,214                return_dict=True).last_hidden_state215        else:216            vit_embeds = self.vision_model(217                pixel_values=pixel_values,218                output_hidden_states=True,219                return_dict=True).hidden_states[self.select_layer]220        vit_embeds = vit_embeds[:, 1:, :]221 222        if self.training and self.neftune_alpha is not None:223            vit_embeds = self.noised_embed(vit_embeds, self.neftune_alpha)224 225        h = w = int(vit_embeds.shape[1] ** 0.5)226        vit_embeds = vit_embeds.reshape(vit_embeds.shape[0], h, w, -1)227        vit_embeds = self.pixel_shuffle(vit_embeds, scale_factor=self.downsample_ratio)228        vit_embeds = vit_embeds.reshape(vit_embeds.shape[0], -1, vit_embeds.shape[-1])229        vit_embeds = self.mlp1(vit_embeds)230        return vit_embeds231 232    def batch_chat(self, tokenizer, pixel_values, image_counts, questions, generation_config, history=None,233                         return_history=False, IMG_START_TOKEN='<img>', IMG_END_TOKEN='</img>',234                         IMG_CONTEXT_TOKEN='<IMG_CONTEXT>'):235        if history is not None or return_history:236            print('Now multi-turn chat is not supported in batch_chat.')237            raise NotImplementedError238        img_context_token_id = tokenizer.convert_tokens_to_ids(IMG_CONTEXT_TOKEN)239        self.img_context_token_id = img_context_token_id240 241        from .conversation import get_conv_template242 243        queries = []244        image_bs = pixel_values.shape[0]245        # print(f'dynamic ViT batch size: {image_bs}, image_counts: {image_counts}')246        for idx, image_count in enumerate(image_counts):247            image_token = IMG_START_TOKEN + IMG_CONTEXT_TOKEN * self.num_image_token * image_count + IMG_END_TOKEN248            question = image_token + '\n' + questions[idx]249            template = get_conv_template(self.template)250            template.append_message(template.roles[0], question)251            template.append_message(template.roles[1], None)252            query = template.get_prompt()253            queries.append(query)254        tokenizer.padding_side = 'left'255        model_inputs = tokenizer(queries, return_tensors='pt', padding=True)256        input_ids = model_inputs['input_ids'].cuda()257        attention_mask = model_inputs['attention_mask'].cuda()258        eos_token_id = tokenizer.convert_tokens_to_ids(template.sep)259        generation_config['eos_token_id'] = eos_token_id260 261        generation_output = self.generate(262            pixel_values=pixel_values,263            input_ids=input_ids,264            attention_mask=attention_mask,265            **generation_config266        )267        responses = tokenizer.batch_decode(generation_output, skip_special_tokens=True)268        responses = [response.split(template.sep)[0].strip() for response in responses]269        return responses270 271    def chat(self, tokenizer, pixel_values, question, generation_config, history=None, return_history=False,272             IMG_START_TOKEN='<img>', IMG_END_TOKEN='</img>', IMG_CONTEXT_TOKEN='<IMG_CONTEXT>'):273 274        img_context_token_id = tokenizer.convert_tokens_to_ids(IMG_CONTEXT_TOKEN)275        self.img_context_token_id = img_context_token_id276 277        from .conversation import get_conv_template278 279        template = get_conv_template(self.template)280        image_bs = pixel_values.shape[0]281        print(f'dynamic ViT batch size: {image_bs}')282        if history is None:283            history = []284            image_tokens = IMG_START_TOKEN + IMG_CONTEXT_TOKEN * self.num_image_token * image_bs + IMG_END_TOKEN285            question = image_tokens + '\n' + question286        else:287            for (old_question, old_answer) in history:288                template.append_message(template.roles[0], old_question)289                template.append_message(template.roles[1], old_answer)290        template.append_message(template.roles[0], question)291        template.append_message(template.roles[1], None)292        query = template.get_prompt()293        model_inputs = tokenizer(query, return_tensors='pt')294        input_ids = model_inputs['input_ids'].cuda()295        attention_mask = model_inputs['attention_mask'].cuda()296        eos_token_id = tokenizer.convert_tokens_to_ids(template.sep)297        generation_config['eos_token_id'] = eos_token_id298 299        generation_output = self.generate(300            pixel_values=pixel_values,301            input_ids=input_ids,302            attention_mask=attention_mask,303            **generation_config304        )305        response = tokenizer.batch_decode(generation_output, skip_special_tokens=True)[0]306        response = response.split(template.sep)[0].strip()307        history.append((question, response))308        if return_history:309            return response, history310        else:311            # query_to_print = query.replace(image_tokens, '<image>')312            # print(query_to_print, response)313            return response314        return response315 316    @torch.no_grad()317    def generate(318            self,319            pixel_values: Optional[torch.FloatTensor] = None,320            input_ids: Optional[torch.FloatTensor] = None,321            attention_mask: Optional[torch.LongTensor] = None,322            visual_features: Optional[torch.FloatTensor] = None,323            generation_config: Optional[GenerationConfig] = None,324            output_hidden_states: Optional[bool] = None,325            return_dict: Optional[bool] = None,326            **generate_kwargs,327    ) -> torch.LongTensor:328 329        assert self.img_context_token_id is not None330        if pixel_values is not None:331            if visual_features is not None:332                vit_embeds = visual_features333            else:334                vit_embeds = self.extract_feature(pixel_values)335            input_embeds = self.language_model.get_input_embeddings()(input_ids)336            B, N, C = input_embeds.shape337            input_embeds = input_embeds.reshape(B * N, C)338 339            input_ids = input_ids.reshape(B * N)340            selected = (input_ids == self.img_context_token_id)341            assert selected.sum() != 0342            input_embeds[selected] = vit_embeds.reshape(-1, C).to(input_embeds.device)343 344            input_embeds = input_embeds.reshape(B, N, C)345        else:346            input_embeds = self.language_model.get_input_embeddings()(input_ids)347 348        outputs = self.language_model.generate(349            inputs_embeds=input_embeds,350            attention_mask=attention_mask,351            generation_config=generation_config,352            output_hidden_states=output_hidden_states,353            return_dict=return_dict,354            use_cache=True,355            **generate_kwargs,356        )357 358        return outputs359