MCPcopy Create free account
hub / github.com/MoonInTheRiver/DiffSinger / MultiheadAttention

Class MultiheadAttention

modules/commons/common_layers.py:166–464  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

164
165
166class MultiheadAttention(nn.Module):
167 def __init__(self, embed_dim, num_heads, kdim=None, vdim=None, dropout=0., bias=True,
168 add_bias_kv=False, add_zero_attn=False, self_attention=False,
169 encoder_decoder_attention=False):
170 super().__init__()
171 self.embed_dim = embed_dim
172 self.kdim = kdim if kdim is not None else embed_dim
173 self.vdim = vdim if vdim is not None else embed_dim
174 self.qkv_same_dim = self.kdim == embed_dim and self.vdim == embed_dim
175
176 self.num_heads = num_heads
177 self.dropout = dropout
178 self.head_dim = embed_dim // num_heads
179 assert self.head_dim * num_heads == self.embed_dim, "embed_dim must be divisible by num_heads"
180 self.scaling = self.head_dim ** -0.5
181
182 self.self_attention = self_attention
183 self.encoder_decoder_attention = encoder_decoder_attention
184
185 assert not self.self_attention or self.qkv_same_dim, 'Self-attention requires query, key and ' \
186 'value to be of the same size'
187
188 if self.qkv_same_dim:
189 self.in_proj_weight = Parameter(torch.Tensor(3 * embed_dim, embed_dim))
190 else:
191 self.k_proj_weight = Parameter(torch.Tensor(embed_dim, self.kdim))
192 self.v_proj_weight = Parameter(torch.Tensor(embed_dim, self.vdim))
193 self.q_proj_weight = Parameter(torch.Tensor(embed_dim, embed_dim))
194
195 if bias:
196 self.in_proj_bias = Parameter(torch.Tensor(3 * embed_dim))
197 else:
198 self.register_parameter('in_proj_bias', None)
199
200 self.out_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
201
202 if add_bias_kv:
203 self.bias_k = Parameter(torch.Tensor(1, 1, embed_dim))
204 self.bias_v = Parameter(torch.Tensor(1, 1, embed_dim))
205 else:
206 self.bias_k = self.bias_v = None
207
208 self.add_zero_attn = add_zero_attn
209
210 self.reset_parameters()
211
212 self.enable_torch_version = False
213 if hasattr(F, "multi_head_attention_forward"):
214 self.enable_torch_version = True
215 else:
216 self.enable_torch_version = False
217 self.last_attn_probs = None
218
219 def reset_parameters(self):
220 if self.qkv_same_dim:
221 nn.init.xavier_uniform_(self.in_proj_weight)
222 else:
223 nn.init.xavier_uniform_(self.k_proj_weight)

Callers 2

__init__Method · 0.85
__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected