CoolFace
Modelpublic

pcuenq/nvidia-nano-clone

sourceHugging Faceotherupdated 11mo agoView on Hugging Face
0likes16downloads
modeling.py288 linesDownload Raw Back to root
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