pcuenq/nvidia-nano-clone
016
1import os2import warnings3from typing import List, Optional, Tuple, Union4 5import torch6import transformers7from torch import nn8from torch.nn import CrossEntropyLoss9from transformers import AutoModel, AutoModelForCausalLM, GenerationConfig10from transformers.modeling_outputs import CausalLMOutputWithPast11from transformers.modeling_utils import PreTrainedModel12from transformers.utils import logging13 14from .configuration import NemotronH_Nano_VL_V2_Config15from .modeling_nemotron_h import NemotronHForCausalLM16from .evs import EfficientVideoSampling17 18logger = logging.get_logger(__name__)19 20 21"""22The following code is adapted from the23https://huggingface.co/OpenGVLab/InternVL2-Llama3-76B/blob/main/modeling_internvl_chat.py repository24 25The chat function is adapted to handle NVLM 1-D tile-tagging design for dynamic high-resolution images.26"""27 28 29class SquaredReLU(nn.Module):30 def forward(self, x):31 return torch.pow(torch.nn.functional.relu(x), 2)32 33 34class RMSNorm(nn.Module):35 def __init__(self, hidden_size, eps=1e-5):36 super().__init__()37 self.weight = nn.Parameter(torch.ones(hidden_size))38 self.eps = eps39 40 def forward(self, hidden_states):41 input_dtype = hidden_states.dtype42 hidden_states = hidden_states.to(torch.float32)43 variance = hidden_states.pow(2).mean(-1, keepdim=True)44 hidden_states = hidden_states * torch.rsqrt(variance + self.eps)45 return (self.weight.to(torch.float32) * hidden_states).to(input_dtype)46 47 48def version_cmp(v1, v2, op='eq'):49 import operator50 51 from packaging import version52 op_func = getattr(operator, op)53 return op_func(version.parse(v1), version.parse(v2))54 55 56class NemotronH_Nano_VL_V2(PreTrainedModel):57 config_class = NemotronH_Nano_VL_V2_Config58 main_input_name = 'pixel_values'59 _supports_flash_attn_2 = True60 _no_split_modules = ['NemotronHBlock']61 62 def __init__(self, config: NemotronH_Nano_VL_V2_Config):63 super().__init__(config)64 65 assert version_cmp(transformers.__version__, '4.36.2', 'ge')66 image_size = config.force_image_size67 patch_size = config.patch_size68 self.patch_size = patch_size69 self.template = config.template70 self.num_image_token = int((image_size // patch_size) ** 2 * (config.downsample_ratio ** 2))71 self.downsample_ratio = config.downsample_ratio72 self.ps_version = config.ps_version73 self.image_tag_type = config.image_tag_type74 self.img_context_token_id = config.img_context_token_id75 self.video_context_token_id = config.video_context_token_id76 77 logger.info(f'num_image_token: {self.num_image_token}')78 logger.info(f'ps_version: {self.ps_version}')79 80 self.language_model = AutoModelForCausalLM.from_config(config.llm_config, trust_remote_code=True)81 self.vision_model = AutoModel.from_config(config.vision_config, trust_remote_code=True)82 self.vision_model.model._initialize_weights = self.vision_model.model._init_weights # WAR for transformers issue 38358 83 self.vision_model.radio_model.make_preprocessor_external()84 self.vision_model = self.vision_model.to(self.language_model.config.torch_dtype)85 86 self.drop_vision_class_token = True87 88 # Construct the vision projection.89 # Default90 vit_hidden_size = config.vit_hidden_size91 vision_projection_hidden_size = config.projector_hidden_size92 llm_hidden_size = config.llm_config.hidden_size93 94 self.video_pruning_rate = config.video_pruning_rate95 96 self.mlp1 = nn.Sequential(97 RMSNorm(vit_hidden_size * int(1 / self.downsample_ratio) ** 2, eps=1e-5),98 nn.Linear(vit_hidden_size * int(1 / self.downsample_ratio) ** 2, vision_projection_hidden_size, bias=False),99 SquaredReLU(),100 nn.Linear(vision_projection_hidden_size, llm_hidden_size, bias=False)101 )102 self.mlp1 = self.mlp1.to(self.language_model.config.torch_dtype)103 104 def forward(105 self,106 pixel_values: torch.FloatTensor,107 input_ids: torch.LongTensor = None,108 attention_mask: Optional[torch.Tensor] = None,109 position_ids: Optional[torch.LongTensor] = None,110 image_flags: Optional[torch.LongTensor] = None,111 past_key_values: Optional[List[torch.FloatTensor]] = None,112 labels: Optional[torch.LongTensor] = None,113 inputs_embeds = None,114 use_cache: Optional[bool] = None,115 output_attentions: Optional[bool] = None,116 output_hidden_states: Optional[bool] = None,117 return_dict: Optional[bool] = None,118 ) -> Union[Tuple, CausalLMOutputWithPast]:119 return_dict = return_dict if return_dict is not None else self.config.use_return_dict120 121 if inputs_embeds is None:122 inputs_embeds = self.language_model.get_input_embeddings()(input_ids)123 124 image_flags = image_flags.squeeze(-1)125 126 B, N, C = inputs_embeds.shape127 inputs_embeds = inputs_embeds.reshape(B * N, C)128 129 input_ids = input_ids.reshape(B * N)130 selected = (input_ids == self.img_context_token_id)131 132 vit_batch_size = pixel_values.shape[0]133 vit_embeds = self.extract_feature(pixel_values)134 135 del pixel_values136 137 if torch.distributed.get_rank() == 0:138 print(f'dynamic ViT batch size: {vit_batch_size}, images per sample: {vit_batch_size / B}, dynamic token length: {N}')139 140 vit_embeds = vit_embeds[image_flags == 1]141 try:142 inputs_embeds[selected] = inputs_embeds[selected] * 0.0 + vit_embeds.reshape(-1, C)143 except Exception as e:144 vit_embeds = vit_embeds.reshape(-1, C)145 print(f'warning: {e}, inputs_embeds[selected].shape={inputs_embeds[selected].shape}, '146 f'vit_embeds.shape={vit_embeds.shape}')147 n_token = selected.sum()148 inputs_embeds[selected] = inputs_embeds[selected] * 0.0 + vit_embeds[:n_token]149 150 del vit_embeds151 152 inputs_embeds = inputs_embeds.reshape(B, N, C)153 154 outputs = self.language_model(155 inputs_embeds=inputs_embeds,156 attention_mask=attention_mask,157 position_ids=position_ids,158 past_key_values=past_key_values,159 use_cache=use_cache,160 output_attentions=output_attentions,161 output_hidden_states=output_hidden_states,162 return_dict=return_dict,163 )164 logits = outputs.logits165 166 loss = None167 if labels is not None:168 # Shift so that tokens < n predict n169 shift_logits = logits[..., :-1, :].contiguous()170 shift_labels = labels[..., 1:].contiguous()171 # Flatten the tokens172 loss_fct = CrossEntropyLoss()173 shift_logits = shift_logits.view(-1, self.language_model.config.vocab_size)174 shift_labels = shift_labels.view(-1)175 # Enable model parallelism176 shift_labels = shift_labels.to(shift_logits.device)177 loss = loss_fct(shift_logits, shift_labels)178 179 if not return_dict:180 output = (logits,) + outputs[1:]181 return (loss,) + output if loss is not None else output182 183 return CausalLMOutputWithPast(184 loss=loss,185 logits=logits,186 past_key_values=outputs.past_key_values,187 hidden_states=outputs.hidden_states,188 attentions=outputs.attentions,189 )190 191 def pixel_shuffle(self, x, scale_factor=0.5):192 n, w, h, c = x.size()193 # N, W, H, C --> N, W, H * scale, C // scale194 x = x.view(n, w, int(h * scale_factor), int(c / scale_factor))195 # N, W, H * scale, C // scale --> N, H * scale, W, C // scale196 x = x.permute(0, 2, 1, 3).contiguous()197 # N, H * scale, W, C // scale --> N, H * scale, W * scale, C // (scale ** 2)198 x = x.view(n, int(h * scale_factor), int(w * scale_factor),199 int(c / (scale_factor * scale_factor)))200 if self.ps_version == 'v1':201 warnings.warn("In ps_version 'v1', the height and width have not been swapped back, "202 'which results in a transposed image.')203 else:204 x = x.permute(0, 2, 1, 3).contiguous()205 return x206 207 def extract_feature(self, pixel_values):208 vit_embeds = self.vision_model(pixel_values).features209 vit_embeds = vit_embeds.to(dtype=torch.bfloat16)210 h = w = int(vit_embeds.shape[1] ** 0.5)211 vit_embeds = vit_embeds.reshape(vit_embeds.shape[0], h, w, -1)212 vit_embeds = self.pixel_shuffle(vit_embeds, scale_factor=self.downsample_ratio)213 vit_embeds = vit_embeds.reshape(vit_embeds.shape[0], -1, vit_embeds.shape[-1])214 vit_embeds = self.mlp1(vit_embeds)215 return vit_embeds216 217 @torch.no_grad()218 def generate(219 self,220 pixel_values: Optional[torch.FloatTensor] = None,221 pixel_values_videos: Optional[torch.FloatTensor] = None,222 input_ids: Optional[torch.FloatTensor] = None,223 attention_mask: Optional[torch.LongTensor] = None,224 generation_config: Optional[GenerationConfig] = None,225 output_hidden_states: Optional[bool] = None,226 return_dict: Optional[bool] = None,227 **generate_kwargs,228 ) -> torch.LongTensor:229 assert self.img_context_token_id is not None230 if pixel_values is not None or pixel_values_videos is not None:231 image_vit_embeds, video_vit_embeds = None, None232 if pixel_values is not None:233 pixel_values = pixel_values.to(dtype=self.vision_model.config.torch_dtype)234 image_vit_embeds = self.extract_feature(pixel_values)235 if pixel_values_videos is not None:236 pixel_values_videos = pixel_values_videos.to(dtype=self.vision_model.config.torch_dtype)237 video_vit_embeds = self.extract_feature(pixel_values_videos)238 inputs_embeds = self.language_model.get_input_embeddings()(input_ids)239 B, N, C = inputs_embeds.shape240 inputs_embeds = inputs_embeds.reshape(B * N, C)241 input_ids_copy = input_ids.reshape(B * N)242 if image_vit_embeds is not None:243 image_mask = (input_ids_copy == self.img_context_token_id)244 assert image_mask.sum() != 0245 inputs_embeds[image_mask] = image_vit_embeds.reshape(-1, C).to(inputs_embeds.device, inputs_embeds.dtype)246 if video_vit_embeds is not None:247 if B > 1:248 raise NotImplementedError("Video is not supported for batch size > 1")249 video_mask = (input_ids_copy == self.video_context_token_id)250 assert video_mask.sum() != 0251 inputs_embeds[video_mask] = video_vit_embeds.reshape(-1, C).to(inputs_embeds.device, inputs_embeds.dtype)252 if video_vit_embeds is not None and self.video_pruning_rate > 0: # EVS253 h = w = int(video_vit_embeds.shape[1] ** 0.5) # assumption here (and everywhere else) is that shape is square254 evs_mask = EfficientVideoSampling.compute_retention_mask(255 video_embeds=video_vit_embeds,256 thw=(video_vit_embeds.shape[0], h, w),257 spatial_merge_size=1, # we already work on vision embeddings, so no downsampling to follow258 q=self.video_pruning_rate,259 )260 print(f"pruning rate: {self.video_pruning_rate}, EVS mask: {evs_mask.sum().item()} tokens retained out of {evs_mask.numel()} total video tokens ({evs_mask.sum().item() / evs_mask.numel() * 100:.2f}%)")261 262 retention_mask = torch.ones_like(input_ids_copy, dtype=torch.bool)263 retention_mask[video_mask] = evs_mask.view(-1)264 inputs_embeds = inputs_embeds[retention_mask].unsqueeze(0) # adding batch=1265 if attention_mask is not None:266 attention_mask = attention_mask[:, retention_mask].contiguous()267 if input_ids is not None:268 input_ids = input_ids[:, retention_mask].contiguous()269 else:270 inputs_embeds = inputs_embeds.reshape(B, N, C)271 else:272 inputs_embeds = self.language_model.get_input_embeddings()(input_ids)273 # print(f"DEBUG: input_ids shape: {input_ids.shape}")274 # print(f"DEBUG: input text: {self._tokenizer.decode(input_ids[0])}")275 outputs = self.language_model.generate(276 input_ids=input_ids,277 inputs_embeds=inputs_embeds,278 attention_mask=attention_mask,279 generation_config=generation_config,280 output_hidden_states=output_hidden_states,281 use_cache=True,282 # return_dict_in_generate=True,283 # output_scores=True,284 **generate_kwargs,285 )286 287 return outputs288 