(self, embed_dim, num_heads, dropout=0., bias=True, add_bias_kv=False, add_zero_attn=False, kdim=None, vdim=None)
| 362 | bias_v: Optional[torch.Tensor] |
| 363 | |
| 364 | def __init__(self, embed_dim, num_heads, dropout=0., bias=True, add_bias_kv=False, add_zero_attn=False, kdim=None, vdim=None): |
| 365 | super(MultiheadAttention, self).__init__() |
| 366 | self.embed_dim = embed_dim |
| 367 | self.kdim = kdim if kdim is not None else embed_dim |
| 368 | self.vdim = vdim if vdim is not None else embed_dim |
| 369 | self._qkv_same_embed_dim = self.kdim == embed_dim and self.vdim == embed_dim |
| 370 | |
| 371 | self.num_heads = num_heads |
| 372 | self.dropout = dropout |
| 373 | self.head_dim = embed_dim // num_heads |
| 374 | assert self.head_dim * num_heads == self.embed_dim, "embed_dim must be divisible by num_heads" |
| 375 | |
| 376 | if self._qkv_same_embed_dim is False: |
| 377 | self.q_proj_weight = Parameter(torch.Tensor(embed_dim, embed_dim)) |
| 378 | self.k_proj_weight = Parameter(torch.Tensor(embed_dim, self.kdim)) |
| 379 | self.v_proj_weight = Parameter(torch.Tensor(embed_dim, self.vdim)) |
| 380 | self.register_parameter('in_proj_weight', None) |
| 381 | else: |
| 382 | self.in_proj_weight = Parameter(torch.empty(3 * embed_dim, embed_dim)) |
| 383 | self.register_parameter('q_proj_weight', None) |
| 384 | self.register_parameter('k_proj_weight', None) |
| 385 | self.register_parameter('v_proj_weight', None) |
| 386 | |
| 387 | if bias: |
| 388 | self.in_proj_bias = Parameter(torch.empty(3 * embed_dim)) |
| 389 | else: |
| 390 | self.register_parameter('in_proj_bias', None) |
| 391 | self.out_proj = _LinearWithBias(embed_dim, embed_dim) |
| 392 | |
| 393 | if add_bias_kv: |
| 394 | self.bias_k = Parameter(torch.empty(1, 1, embed_dim)) |
| 395 | self.bias_v = Parameter(torch.empty(1, 1, embed_dim)) |
| 396 | else: |
| 397 | self.bias_k = self.bias_v = None |
| 398 | |
| 399 | self.add_zero_attn = add_zero_attn |
| 400 | |
| 401 | self._reset_parameters() |
| 402 | |
| 403 | def _reset_parameters(self): |
| 404 | if self._qkv_same_embed_dim: |
nothing calls this directly
no test coverage detected