Image-Text-to-Text
Transformers
Safetensors
English
qwen2_5vl_ca
feature-extraction
conversational
custom_code
CASA-Qwen2_5-VL-3B / language_qwen2_5vl_ca.py
ameroyer's picture nielsr's picture
nielsr HF Staff
Super-squash branch 'main' using huggingface_hub
5f7de59
Raw History Blame
5.09 kB
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