from typing import Literal, Optional import torch from transformers.cache_utils import Cache from transformers.configuration_utils import PretrainedConfig from transformers.models.qwen2.modeling_qwen2 import Qwen2RMSNorm from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import ( Qwen2_5_VLDecoderLayer, Qwen2_5_VLFlashAttention2, ) from .cross_attention import ( CrossAttention, CrossAttentionHandler, tie_qkvo_projections, ) from .configuration_qwen2_5vl_ca import Qwen2_5_VLCAConfig class QwenCrossAttention(CrossAttention): """A CrossAttention layer compatible with Qwen's projection conventions""" def __init__( self, config: Qwen2_5_VLCAConfig, layer_idx: int | None, ): super().__init__(config, layer_idx) # pyright: ignore[reportArgumentType] self.norm = Qwen2RMSNorm(config.hidden_size, eps=config.rms_norm_eps) assert config.rope_scaling is not None self.mrope_section = config.rope_scaling["mrope_section"] * 2 def init_from_config_proj( self, key: Literal["q", "o", "k", "v"], config: PretrainedConfig ) -> torch.nn.Linear: """Follows modeling_qwen2_5_vl.py initialization""" head_dim = config.hidden_size // config.num_attention_heads if key == "q": return torch.nn.Linear( config.hidden_size, config.num_attention_heads * head_dim, bias=True ) if key in {"k", "v"}: return torch.nn.Linear( config.hidden_size, config.num_key_value_heads * head_dim, bias=True ) if key == "o": return torch.nn.Linear( config.num_attention_heads * config.head_dim, config.hidden_size, bias=False ) raise NotImplementedError(f"Unknown key {key}") class Qwen2_5_VLAttention_CrossAttention(Qwen2_5_VLFlashAttention2): """ Qwen Attention with extra CrossAttention layer """ def __init__( self, config: Qwen2_5_VLCAConfig, layer_idx: Optional[int] = None, input_layernorm: torch.nn.Module | None = None, ): super().__init__(config, layer_idx) # pyright: ignore[reportArgumentType] self.cross_attn = QwenCrossAttention(config, layer_idx=layer_idx) self.cross_attention_handler: CrossAttentionHandler | None = None if getattr(config, "xa_share_qkvo", False): tie_qkvo_projections(self, self.cross_attn) @classmethod def from_qwen2_5_vl_attention( cls, attention: Qwen2_5_VLFlashAttention2, input_layernorm: torch.nn.Module | None ): """Init this layer from an existing Qwen Attention layer""" layer_idx = attention.layer_idx assert layer_idx is not None new_attention = cls(attention.config, layer_idx=layer_idx, input_layernorm=input_layernorm) # pyright: ignore new_attention.load_state_dict(attention.state_dict(), strict=False) return new_attention def forward( # pyright: ignore[reportIncompatibleMethodOverride] self, hidden_states: torch.Tensor, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_value: Optional[Cache] = None, output_attentions: bool = False, use_cache: bool = False, cache_position: Optional[torch.LongTensor] = None, position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None, ): attn_output, attn_weights, past_key_values = super().forward( hidden_states, attention_mask, position_ids, past_key_value, output_attentions, use_cache, cache_position, position_embeddings, ) if self.cross_attn is not None: ca_out = self.cross_attn( hidden_states=hidden_states, cross_attention_handler=self.cross_attention_handler, ) # ca_out is None when there is no handler (text-only or streaming non-first call) if ca_out is not None: attn_output = ca_out + attn_output return attn_output, attn_weights, past_key_values def maybe_replace_with_cross_attention_layers( m: torch.nn.Module, xa_layers: tuple[int, ...] | None, reindex: bool = False ): """Replace Attention layer by CrossAttention layer as needed""" if isinstance(m, Qwen2_5_VLDecoderLayer): layer_idx = m.self_attn.layer_idx assert layer_idx is not None if xa_layers is None or len(xa_layers) == 0 or layer_idx in xa_layers: m.self_attn = Qwen2_5_VLAttention_CrossAttention.from_qwen2_5_vl_attention( m.self_attn, input_layernorm=m.input_layernorm ) elif reindex: # shift left by number of cross-attention layers before this one logical_idx = layer_idx - sum(j < layer_idx for j in (xa_layers or ())) assert logical_idx >= 0 m.self_attn.layer_idx = logical_idx