MathLLMs/MathCoder-VL-2B
730
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 