MCPcopy Create free account
hub / github.com/FunAudioLLM/Fun-Audio-Chat / FunAudioChatAudioAttention

Class FunAudioChatAudioAttention

funaudiochat/modeling_funaudiochat.py:100–171  ·  view source on GitHub ↗

Multi-headed attention from 'Attention Is All You Need' paper

Source from the content-addressed store, hash-verified

98
99
100class 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,

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected