Multi-headed attention from 'Attention Is All You Need' paper
| 98 | |
| 99 | |
| 100 | class FunAudioChatAudioAttention(nn.Module): |
| 101 | """Multi-headed attention from 'Attention Is All You Need' paper""" |
| 102 | |
| 103 | def __init__( |
| 104 | self, |
| 105 | config: FunAudioChatAudioEncoderConfig, |
| 106 | ): |
| 107 | super().__init__() |
| 108 | self.embed_dim = config.d_model |
| 109 | self.num_heads = config.encoder_attention_heads |
| 110 | self.dropout = config.attention_dropout |
| 111 | self.head_dim = self.embed_dim // self.num_heads |
| 112 | self.num_key_value_groups = 1 # needed for eager attention |
| 113 | self.config = config |
| 114 | |
| 115 | if (self.head_dim * self.num_heads) != self.embed_dim: |
| 116 | raise ValueError( |
| 117 | f"embed_dim must be divisible by num_heads (got `embed_dim`: {self.embed_dim}" |
| 118 | f" and `num_heads`: {self.num_heads})." |
| 119 | ) |
| 120 | self.scaling = self.head_dim**-0.5 |
| 121 | self.attention_dropout = 0.0 |
| 122 | self.is_decoder = False |
| 123 | self.is_causal = False |
| 124 | |
| 125 | self.k_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=False) |
| 126 | self.v_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=True) |
| 127 | self.q_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=True) |
| 128 | self.out_proj = nn.Linear(self.embed_dim, self.embed_dim, bias=True) |
| 129 | |
| 130 | def forward( |
| 131 | self, |
| 132 | hidden_states: torch.Tensor, |
| 133 | cu_seqlens: Optional[torch.Tensor] = None, |
| 134 | attention_mask: Optional[torch.Tensor] = None, |
| 135 | **kwargs, |
| 136 | ) -> torch.Tensor: |
| 137 | seq_length, _ = hidden_states.size() |
| 138 | |
| 139 | query_states = self.q_proj(hidden_states).reshape(seq_length, self.num_heads, -1) |
| 140 | key_states = self.k_proj(hidden_states).reshape(seq_length, self.num_heads, -1) |
| 141 | value_states = self.v_proj(hidden_states).reshape(seq_length, self.num_heads, -1) |
| 142 | |
| 143 | query_states = query_states.transpose(0, 1).unsqueeze(0) |
| 144 | key_states = key_states.transpose(0, 1).unsqueeze(0) |
| 145 | value_states = value_states.transpose(0, 1).unsqueeze(0) |
| 146 | max_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).max() |
| 147 | |
| 148 | attention_interface = eager_attention_forward |
| 149 | if self.config._attn_implementation != "eager": |
| 150 | attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation] |
| 151 | |
| 152 | attn_output, _ = attention_interface( |
| 153 | self, |
| 154 | query_states, |
| 155 | key_states, |
| 156 | value_states, |
| 157 | attention_mask=attention_mask, |