MCPcopy Create free account
hub / github.com/UX-Decoder/Semantic-SAM / forward

Method forward

semantic_sam/modules/attention.py:426–487  ·  view source on GitHub ↗

r""" Args: query, key, value: map a query and a set of key-value pairs to an output. See "Attention Is All You Need" for more details. key_padding_mask: if provided, specified padding elements in the key will be ignored by the attention. When given a binar

(self, query: Tensor, key: Tensor, value: Tensor, key_padding_mask: Optional[Tensor] = None,
                need_weights: bool = True, attn_mask: Optional[Tensor] = None)

Source from the content-addressed store, hash-verified

424 super(MultiheadAttention, self).__setstate__(state)
425
426 def forward(self, query: Tensor, key: Tensor, value: Tensor, key_padding_mask: Optional[Tensor] = None,
427 need_weights: bool = True, attn_mask: Optional[Tensor] = None) -> Tuple[Tensor, Optional[Tensor]]:
428 r"""
429 Args:
430 query, key, value: map a query and a set of key-value pairs to an output.
431 See "Attention Is All You Need" for more details.
432 key_padding_mask: if provided, specified padding elements in the key will
433 be ignored by the attention. When given a binary mask and a value is True,
434 the corresponding value on the attention layer will be ignored. When given
435 a byte mask and a value is non-zero, the corresponding value on the attention
436 layer will be ignored
437 need_weights: output attn_output_weights.
438 attn_mask: 2D or 3D mask that prevents attention to certain positions. A 2D mask will be broadcasted for all
439 the batches while a 3D mask allows to specify a different mask for the entries of each batch.
440
441 Shapes for inputs:
442 - query: :math:`(L, N, E)` where L is the target sequence length, N is the batch size, E is
443 the embedding dimension.
444 - key: :math:`(S, N, E)`, where S is the source sequence length, N is the batch size, E is
445 the embedding dimension.
446 - value: :math:`(S, N, E)` where S is the source sequence length, N is the batch size, E is
447 the embedding dimension.
448 - key_padding_mask: :math:`(N, S)` where N is the batch size, S is the source sequence length.
449 If a ByteTensor is provided, the non-zero positions will be ignored while the position
450 with the zero positions will be unchanged. If a BoolTensor is provided, the positions with the
451 value of ``True`` will be ignored while the position with the value of ``False`` will be unchanged.
452 - attn_mask: if a 2D mask: :math:`(L, S)` where L is the target sequence length, S is the
453 source sequence length.
454
455 If a 3D mask: :math:`(N\cdot\text{num\_heads}, L, S)` where N is the batch size, L is the target sequence
456 length, S is the source sequence length. ``attn_mask`` ensure that position i is allowed to attend
457 the unmasked positions. If a ByteTensor is provided, the non-zero positions are not allowed to attend
458 while the zero positions will be unchanged. If a BoolTensor is provided, positions with ``True``
459 is not allowed to attend while ``False`` values will be unchanged. If a FloatTensor
460 is provided, it will be added to the attention weight.
461
462 Shapes for outputs:
463 - attn_output: :math:`(L, N, E)` where L is the target sequence length, N is the batch size,
464 E is the embedding dimension.
465 - attn_output_weights: :math:`(N, L, S)` where N is the batch size,
466 L is the target sequence length, S is the source sequence length.
467 """
468 if not self._qkv_same_embed_dim:
469 return multi_head_attention_forward(
470 query, key, value, self.embed_dim, self.num_heads,
471 self.in_proj_weight, self.in_proj_bias,
472 self.bias_k, self.bias_v, self.add_zero_attn,
473 self.dropout, self.out_proj.weight, self.out_proj.bias,
474 training=self.training,
475 key_padding_mask=key_padding_mask, need_weights=need_weights,
476 attn_mask=attn_mask, use_separate_proj_weight=True,
477 q_proj_weight=self.q_proj_weight, k_proj_weight=self.k_proj_weight,
478 v_proj_weight=self.v_proj_weight)
479 else:
480 return multi_head_attention_forward(
481 query, key, value, self.embed_dim, self.num_heads,
482 self.in_proj_weight, self.in_proj_bias,
483 self.bias_k, self.bias_v, self.add_zero_attn,

Callers

nothing calls this directly

Calls 1

Tested by

no test coverage detected