Aluode/PerceptionLabPortable
0
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 