from functools import partial from typing import Any, Sequence from typing import cast as type_cast import torch from transformers.cache_utils import DynamicCache from transformers.generation.utils import GenerateOutput from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import ( Qwen2_5_VLCausalLMOutputWithPast, Qwen2_5_VLForConditionalGeneration, ) from .cross_attention import CrossAttentionHandler, tie_qkvo_projections from .image_encoder import Qwen25VLEncoder from .configuration_qwen2_5vl_ca import Qwen2_5_VLCAConfig from .language_qwen2_5vl_ca import ( Qwen2_5_VLAttention_CrossAttention, QwenCrossAttention, maybe_replace_with_cross_attention_layers, ) class V2Qwen2_5VL(Qwen2_5_VLForConditionalGeneration): # pyright: ignore[reportIncompatibleMethodOverride] config_class = Qwen2_5_VLCAConfig def __init__(self, config: Qwen2_5_VLCAConfig, **kwargs: Any) -> None: del kwargs super().__init__(config) # Wrap the Qwen visual encoder so its output matches our CA interface self.image_prefix = Qwen25VLEncoder(self.visual) # type: ignore[assignment] self.visual = None self.model.apply( partial(maybe_replace_with_cross_attention_layers, xa_layers=self.config.xa_layers) ) # The cross-attention layers are swapped in after the base post_init, so register # the shared-weight alias keys and (re-)tie now that the CA modules exist. if config.xa_share_qkvo: shared_keys: list[str] = [] for i, layer in enumerate(self.model.layers): if isinstance(layer.self_attn, Qwen2_5_VLAttention_CrossAttention): for proj, biased in ( ("q_proj", True), ("k_proj", True), ("v_proj", True), ("o_proj", False), ): prefix = f"model.layers.{i}.self_attn.cross_attn.{proj}" shared_keys.append(f"{prefix}.weight") if biased: shared_keys.append(f"{prefix}.bias") self._tied_weights_keys = list(self._tied_weights_keys or []) + shared_keys self.tie_weights() def _tie_weights(self) -> None: if not getattr(self.config, "xa_share_qkvo", False): return for layer in self.model.layers: if isinstance(layer.self_attn, Qwen2_5_VLAttention_CrossAttention): tie_qkvo_projections(layer.self_attn, layer.self_attn.cross_attn) def get_device(self) -> str: """Return the device type of the model""" return next(self.parameters()).device.type @property def token_dim(self) -> int: """Returns the number of dimensions for the token representation""" return self.config.hidden_size def _update_model_kwargs_for_generation( self, outputs: Any, model_kwargs: dict[str, Any], is_encoder_decoder: bool = False, num_new_tokens: int = 1, ): """Override to handle multi-turn generation and propagate updated attention masks""" if (am := outputs.get("updated_attention_mask", None)) is not None: model_kwargs["attention_mask"] = am if "updated_cache_position" in outputs: model_kwargs["cache_position"] = outputs.get("updated_cache_position") else: start = 0 if (kv := model_kwargs.get("past_key_values", None)) is not None: start = kv._seen_tokens - am.shape[1] model_kwargs["cache_position"] = torch.arange( start, start + am.shape[1], dtype=model_kwargs["cache_position"].dtype, device=model_kwargs["cache_position"].device, ) # Call parent to get default updates model_kwargs = super()._update_model_kwargs_for_generation( outputs, model_kwargs, is_encoder_decoder, num_new_tokens ) # Used by prepare_inputs_for_generation model_kwargs["__is_first_gen_call__"] = False return model_kwargs def prepare_inputs_for_generation( # pyright: ignore[reportIncompatibleMethodOverride] self, input_ids: torch.Tensor, past_key_values: DynamicCache | None = None, **kwargs: Any, ): """Override to avoid Qwen erasing pixel_values on subsequent generation calls""" backup = None __is_first_gen_call__ = kwargs.get("__is_first_gen_call__", True) if __is_first_gen_call__: backup = kwargs.get("pixel_values", None) if past_key_values is not None and ( kwargs.get("cache_position") is None or type_cast(torch.Tensor, kwargs.get("cache_position")).shape[0] == 0 ): # We're continuing from a cached state past_length = past_key_values._seen_tokens kwargs["cache_position"] = torch.arange( past_length, past_length + (input_ids.shape[1] if __is_first_gen_call__ else 1), dtype=torch.long, device=input_ids.device, ) out = super().prepare_inputs_for_generation( input_ids, past_key_values=past_key_values, **kwargs, ) if backup is not None: out["pixel_values"] = backup return out def prepare_multimodal_inputs( self, input_ids: torch.Tensor | None = None, inputs_embeds: torch.Tensor | None = None, attention_mask: torch.Tensor | None = None, image_embeds_insertion_points: list[torch.Tensor] | None = None, labels: torch.Tensor | None = None, pixel_values: torch.Tensor | list[torch.Tensor] | None = None, pre_image_tokens: list[int] | None = None, post_image_tokens: list[int] | None = None, **_kwargs: Any, ) -> dict: """Get a batch data mixing text and image data""" del _kwargs processed_inputs: dict = { "input_ids": input_ids, "inputs_embeds": inputs_embeds, "labels": labels, "attention_mask": attention_mask, "image_embeds_insertion_points": image_embeds_insertion_points, } if pixel_values is not None: processed_inputs.update(self.image_prefix(pixel_values)) image_embeds = processed_inputs.get("image_embeds") assert image_embeds is not None assert (isinstance(image_embeds, torch.Tensor) and image_embeds.ndim == 3) or ( isinstance(image_embeds, list) and all(_x.ndim == 2 for _x in image_embeds) ) # Add kwargs necessary to compute cu_seqlens windows for CA processed_inputs["ca_windows_info"] = { "num_post_image_tokens": 0 if post_image_tokens is None else len(post_image_tokens), "num_pre_image_tokens": 0 if pre_image_tokens is None else len(pre_image_tokens), } return processed_inputs def forward( # type: ignore[override] # pylint: disable=W0221 self, input_ids: torch.Tensor | None = None, inputs_embeds: torch.Tensor | None = None, attention_mask: torch.Tensor | None = None, pixel_values: torch.Tensor | list[torch.Tensor] | None = None, return_loss: bool = True, labels: torch.Tensor | None = None, image_embeds_insertion_points: list[torch.Tensor] | None = None, pre_image_tokens: list[int] | None = None, post_image_tokens: list[int] | None = None, **kwargs: Any, ) -> tuple | Qwen2_5_VLCausalLMOutputWithPast: """Multi-modal forward pass""" if self.training: assert return_loss is True, ( "Qwen2.5VL always computes its own labels/losses in train mode" ) if inputs_embeds is None: assert input_ids is not None inputs_embeds = type_cast(torch.Tensor, self.model.embed_tokens(input_ids)) # Case 1: First generation call — compute image embeddings and set up CA handler if kwargs.pop("__is_first_gen_call__", True): processed_inputs = self.prepare_multimodal_inputs( input_ids=input_ids, inputs_embeds=inputs_embeds, attention_mask=attention_mask, image_embeds_insertion_points=image_embeds_insertion_points, pixel_values=pixel_values, labels=labels, pre_image_tokens=pre_image_tokens, post_image_tokens=post_image_tokens, ) image_embeds = processed_inputs.get("image_embeds", None) inst_points = processed_inputs.get("image_embeds_insertion_points", None) # Only build a handler when images are actually present cross_attention_handler: CrossAttentionHandler | None = None if image_embeds is not None and len(image_embeds) > 0: cross_attention_handler = CrossAttentionHandler( inputs_embeds=torch.zeros_like(inputs_embeds), image_embeds=image_embeds, image_embeds_insertion_points=inst_points, ca_windows_info=processed_inputs.pop("ca_windows_info", None), training=self.training, ) self.update_cross_attention_states(cross_attention_handler) # Run Qwen with the attention layers replaced to use cross-attention assert inputs_embeds is not None, "Could not compute input embeddings!" out = super().forward( inputs_embeds=inputs_embeds, # type: ignore[arg-type] attention_mask=attention_mask, pixel_values=None, **kwargs, ) return out @property def default_generation_eos_token_id(self) -> int | Sequence[int] | None: return self.generation_config.eos_token_id if self.generation_config is not None else None @torch.no_grad() def generate_from_image( # pyright: ignore[reportInconsistentOverload] self, reset_streaming: bool = True, temperature: float | None = 0.0, eos_token_id: int | Sequence[int] | None = None, **kwargs: Any, ) -> GenerateOutput | torch.LongTensor: """Custom generate function""" # init self-attention KVCache if kwargs.get("past_key_values", None) is None: kwargs["past_key_values"] = DynamicCache() if eos_token_id is None: eos_token_id = self.default_generation_eos_token_id # To avoid generate warning if kwargs.get("pad_token_id", None) is None: kwargs["pad_token_id"] = kwargs.get("eos_token_id", None) if isinstance(kwargs["pad_token_id"], (list, tuple)): kwargs["pad_token_id"] = kwargs["pad_token_id"][0] if "pre_image_tokens" not in kwargs: kwargs["pre_image_tokens"] = list(self.config.pre_image_tokens) if "post_image_tokens" not in kwargs: kwargs["post_image_tokens"] = list(self.config.post_image_tokens) if not kwargs.get("do_sample", False): temperature = None kwargs.pop("top_p", None) kwargs.pop("top_k", None) # Generate self.start_ca_streaming_states() outputs = self.generate( use_cache=True, eos_token_id=eos_token_id, temperature=temperature, **kwargs, ) if reset_streaming: self.reset_ca_streaming_states() return outputs def update_cross_attention_states(self, handler: CrossAttentionHandler | None): """Push the new handler into all CA attention layers""" def __update__(m: torch.nn.Module): nonlocal handler if isinstance(m, Qwen2_5_VLAttention_CrossAttention): m.cross_attention_handler = handler self.apply(__update__) def reset_ca_streaming_states(self) -> None: def __reset__(m: torch.nn.Module): if isinstance(m, QwenCrossAttention): m._set_streaming(False, ()) m.reset_streaming() if hasattr(m, "cross_attention_handler"): del m.cross_attention_handler m.cross_attention_handler = None self.apply(__reset__) def start_ca_streaming_states(self) -> None: def __start__(m: torch.nn.Module): if isinstance(m, QwenCrossAttention): m._set_streaming(True, ()) self.apply(__start__)