CoolFace
Apppublic

Aluode/PerceptionLabPortable

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
modular_aya_vision.py298 linesDownload Raw Back to aya_vision
1# coding=utf-82# Copyright 2025 the Cohere Inc. team. All rights reserved.3#4# Licensed under the Apache License, Version 2.0 (the "License");5# you may not use this file except in compliance with the License.6# You may obtain a copy of the License at7#8#     http://www.apache.org/licenses/LICENSE-2.09#10# Unless required by applicable law or agreed to in writing, software11# distributed under the License is distributed on an "AS IS" BASIS,12# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.13# See the License for the specific language governing permissions and14# limitations under the License.15"""PyTorch AyaVision model."""16 17from typing import Optional, Union18 19import torch20from torch import nn21 22from transformers.models.llava.modeling_llava import (23    LlavaCausalLMOutputWithPast,24    LlavaForConditionalGeneration,25    LlavaModel,26    LlavaModelOutputWithPast,27    LlavaPreTrainedModel,28    TransformersKwargs,29)30 31from ...activations import ACT2FN32from ...cache_utils import Cache33from ...processing_utils import Unpack34from ...utils import auto_docstring, logging35from ...utils.generic import check_model_inputs36from .configuration_aya_vision import AyaVisionConfig37 38 39logger = logging.get_logger(__name__)40 41 42class AyaVisionMultiModalProjector(nn.Module):43    def __init__(self, config: AyaVisionConfig):44        super().__init__()45        self.config = config46        self.downsample_factor = config.downsample_factor47        self.alignment_intermediate_size = getattr(48            config, "alignment_intermediate_size", config.text_config.hidden_size49        )50        self.layernorm = nn.LayerNorm(51            config.vision_config.hidden_size * (config.downsample_factor**2), eps=config.adapter_layer_norm_eps52        )53 54        self.linear_1 = nn.Linear(55            config.vision_config.hidden_size * (config.downsample_factor**2),56            self.alignment_intermediate_size,57            bias=True,58        )59 60        self.act = ACT2FN["silu"]  # SwiGLU uses SiLU activation61        # For SwiGLU, project down to half size since we split intermediate dim62        self.linear_2 = nn.Linear(self.alignment_intermediate_size // 2, config.text_config.hidden_size, bias=True)63 64    def forward(self, image_features):65        image_features = self.pixel_shuffle(image_features)66        image_features = self.layernorm(image_features)67        hidden_states = self.linear_1(image_features)68 69        # Split along last dimension and apply SwiGLU70        x, gate = hidden_states.chunk(2, dim=-1)71        hidden_states = self.act(gate) * x72 73        hidden_states = self.linear_2(hidden_states)74        return hidden_states75 76    def pixel_shuffle(self, image_features):  # B, S, D77        batch_size, seq_length, feature_dim = image_features.shape78        height = width = int(seq_length**0.5)79        image_features = image_features.reshape(image_features.shape[0], width, height, -1)80        channels = image_features.shape[-1]81        image_features = image_features.reshape(82            batch_size, width, int(height / self.downsample_factor), int(channels * self.downsample_factor)83        )84        image_features = image_features.permute(0, 2, 1, 3)85        image_features = image_features.reshape(86            batch_size, int(height / self.downsample_factor), int(width / self.downsample_factor), -187        )88        image_features = image_features.permute(0, 2, 1, 3)89        return image_features90 91 92class AyaVisionPreTrainedModel(LlavaPreTrainedModel):93    _can_compile_fullgraph = False94    _can_record_outputs = {95        "hidden_states": "DecoderLayer",96        "attentions": "Attention",97    }98 99 100class AyaVisionCausalLMOutputWithPast(LlavaCausalLMOutputWithPast):101    pass102 103 104class AyaVisionModelOutputWithPast(LlavaModelOutputWithPast):105    pass106 107 108class AyaVisionModel(LlavaModel):109    # Unlike LLaVA, the model doesn't have to deal with Pixtral-style image states110    def get_image_features(111        self,112        pixel_values: torch.FloatTensor,113        vision_feature_layer: Optional[Union[int, list[int]]] = None,114        vision_feature_select_strategy: Optional[str] = None,115        **kwargs,116    ):117        """118        Obtains image last hidden states from the vision tower and apply multimodal projection.119 120        Args:121            pixel_values (`torch.FloatTensor]` of shape `(batch_size, channels, height, width)`):122               The tensors corresponding to the input images.123            vision_feature_layer (`Union[int, list[int]]`, *optional*):124                The index of the layer to select the vision feature. If multiple indices are provided,125                the vision feature of the corresponding indices will be concatenated to form the126                vision features.127            vision_feature_select_strategy (`str`, *optional*):128                The feature selection strategy used to select the vision feature from the vision backbone.129                Can be one of `"default"` or `"full"`130        Returns:131            image_features (`torch.Tensor`): Image feature tensor of shape `(num_images, image_length, embed_dim)`).132        """133        vision_feature_layer = (134            vision_feature_layer if vision_feature_layer is not None else self.config.vision_feature_layer135        )136        vision_feature_select_strategy = (137            vision_feature_select_strategy138            if vision_feature_select_strategy is not None139            else self.config.vision_feature_select_strategy140        )141 142        if vision_feature_select_strategy not in ["default", "full"]:143            raise ValueError(f"Unexpected select feature strategy: {self.config.vision_feature_select_strategy}")144 145        kwargs = {k: v for k, v in kwargs.items() if v is not None}146        # this is not memory efficient at all (output_hidden_states=True) will save all the hidden states.147        image_outputs = self.vision_tower(pixel_values, output_hidden_states=True, **kwargs)148 149        # If we have one vision feature layer, return the corresponding hidden states,150        # otherwise, select the hidden states of each feature layer and concatenate them151        if isinstance(vision_feature_layer, int):152            selected_image_feature = image_outputs.hidden_states[vision_feature_layer]153            if vision_feature_select_strategy == "default":154                selected_image_feature = selected_image_feature[:, 1:]155        else:156            hs_pool = [image_outputs.hidden_states[layer_idx] for layer_idx in vision_feature_layer]157            # For default; crop CLS from each hidden state in the hidden state pool158            if vision_feature_select_strategy == "default":159                hs_pool = [hs[:, 1:] for hs in hs_pool]160            selected_image_feature = torch.cat(hs_pool, dim=-1)161 162        image_features = self.multi_modal_projector(selected_image_feature)163        return image_features164 165    @check_model_inputs()166    @auto_docstring167    def forward(168        self,169        input_ids: Optional[torch.LongTensor] = None,170        pixel_values: Optional[torch.FloatTensor] = None,171        attention_mask: Optional[torch.Tensor] = None,172        position_ids: Optional[torch.LongTensor] = None,173        past_key_values: Optional[Cache] = None,174        inputs_embeds: Optional[torch.FloatTensor] = None,175        vision_feature_layer: Optional[Union[int, list[int]]] = None,176        vision_feature_select_strategy: Optional[str] = None,177        use_cache: Optional[bool] = None,178        cache_position: Optional[torch.LongTensor] = None,179        **kwargs: Unpack[TransformersKwargs],180    ) -> Union[tuple, AyaVisionModelOutputWithPast]:181        vision_feature_layer = (182            vision_feature_layer if vision_feature_layer is not None else self.config.vision_feature_layer183        )184        vision_feature_select_strategy = (185            vision_feature_select_strategy186            if vision_feature_select_strategy is not None187            else self.config.vision_feature_select_strategy188        )189 190        if (input_ids is None) ^ (inputs_embeds is not None):191            raise ValueError("You must specify exactly one of input_ids or inputs_embeds")192 193        if inputs_embeds is None:194            inputs_embeds = self.get_input_embeddings()(input_ids)195 196        if pixel_values is not None:197            image_features = self.get_image_features(198                pixel_values=pixel_values,199                vision_feature_layer=vision_feature_layer,200                vision_feature_select_strategy=vision_feature_select_strategy,201            )202            image_features = image_features.to(inputs_embeds.device, inputs_embeds.dtype)203            special_image_mask = self.get_placeholder_mask(204                input_ids, inputs_embeds=inputs_embeds, image_features=image_features205            )206            inputs_embeds = inputs_embeds.masked_scatter(special_image_mask, image_features)207 208        outputs = self.language_model(209            attention_mask=attention_mask,210            position_ids=position_ids,211            past_key_values=past_key_values,212            inputs_embeds=inputs_embeds,213            use_cache=use_cache,214            cache_position=cache_position,215            **kwargs,216        )217 218        return AyaVisionModelOutputWithPast(219            last_hidden_state=outputs.last_hidden_state,220            past_key_values=outputs.past_key_values,221            hidden_states=outputs.hidden_states,222            attentions=outputs.attentions,223            image_hidden_states=image_features if pixel_values is not None else None,224        )225 226 227class AyaVisionForConditionalGeneration(LlavaForConditionalGeneration):228    def forward(229        self,230        input_ids: Optional[torch.LongTensor] = None,231        pixel_values: Optional[torch.FloatTensor] = None,232        attention_mask: Optional[torch.Tensor] = None,233        position_ids: Optional[torch.LongTensor] = None,234        past_key_values: Optional[Cache] = None,235        inputs_embeds: Optional[torch.FloatTensor] = None,236        vision_feature_layer: Optional[Union[int, list[int]]] = None,237        vision_feature_select_strategy: Optional[str] = None,238        labels: Optional[torch.LongTensor] = None,239        cache_position: Optional[torch.LongTensor] = None,240        logits_to_keep: Union[int, torch.Tensor] = 0,241        image_sizes: Optional[torch.Tensor] = None,242        **kwargs: Unpack[TransformersKwargs],243    ) -> Union[tuple, AyaVisionCausalLMOutputWithPast]:244        r"""245        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):246            Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,247            config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored248            (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.249 250        Example:251 252        ```python253        >>> from transformers import AutoProcessor, AyaVisionForConditionalGeneration254        >>> import torch255 256        >>> torch_device = "cuda:0"257        >>> processor = AutoProcessor.from_pretrained("CohereForAI/aya-vision-8b", use_fast=True)258        >>> model = AyaVisionForConditionalGeneration.from_pretrained("CohereForAI/aya-vision-8b", device_map=torch_device)259 260        >>> messages = [261        ...     {262        ...         "role": "user",263        ...         "content": [264        ...             {265        ...                 "type": "image",266        ...                 "url": "https://pbs.twimg.com/media/Fx7YvfQWYAIp6rZ?format=jpg&name=medium",267        ...             },268        ...             {"type": "text", "text": "चित्र में लिखा पाठ क्या कहता है?"},269        ...         ],270        ...     }271        ... ]272 273        >>> inputs = processor.apply_chat_template(274        ...     messages, padding=True, add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt", device=torch_device275        ... ).to(model.device)276 277        >>> gen_tokens = model.generate(**inputs, max_new_tokens=300, do_sample=True, temperature=0.3)278        >>> processor.tokenizer.decode(gen_tokens[0][inputs.input_ids.shape[1]:], skip_special_tokens=True)279        ```"""280        super().forward(281            input_ids=input_ids,282            pixel_values=pixel_values,283            attention_mask=attention_mask,284            position_ids=position_ids,285            past_key_values=past_key_values,286            inputs_embeds=inputs_embeds,287            vision_feature_layer=vision_feature_layer,288            vision_feature_select_strategy=vision_feature_select_strategy,289            labels=labels,290            cache_position=cache_position,291            logits_to_keep=logits_to_keep,292            image_sizes=image_sizes,293            **kwargs,294        )295 296 297__all__ = ["AyaVisionForConditionalGeneration", "AyaVisionPreTrainedModel", "AyaVisionModel"]298 
Aluode/PerceptionLabPortable · CoolFace