Image-Text-to-Text
Transformers
Safetensors
English
qwen2_5vl_ca
feature-extraction
conversational
custom_code
File size: 5,085 Bytes
5f7de59
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
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