MultiheadAttention
MultiheadAttention(embed_dim, num_heads, dropout=0.0, bias=True, add_bias_kv=False, add_zero_attn=False, kdim=None, vdim=None, batch_first=False, device=None, dtype=None)
MultiHead(Q, K, V) = Concat(head_1, ..., head_h) W_o with
head_i = Attention(Q W_q_i, K W_k_i, V W_v_i).
Args:
embed_dim: the model width; every head gets embed_dim // num_heads of it.
num_heads: the number of parallel attention heads.
dropout: accepted for signature compatibility; this runtime
serves inference, so a non-zero value only takes effect
under training and raises there.
bias: add biases to the input and output projections.
add_bias_kv: append one learned key and value row
(bias_k / bias_v) to every sequence.
add_zero_attn: append one all-zero key and value row to every
sequence.
kdim: the key width when it differs from embed_dim.
vdim: the value width when it differs from embed_dim.
batch_first: inputs and outputs as [batch, seq, feature]
instead of [seq, batch, feature].
device: where the parameters live; None means the cpu.
dtype: the parameter dtype; None means float32.
__init__
__init__(self, embed_dim: 'int', num_heads: 'int', dropout: 'float' = 0.0, bias: 'bool' = True, add_bias_kv: 'bool' = False, add_zero_attn: 'bool' = False, kdim: 'int | None' = None, vdim: 'int | None' = None, batch_first: 'bool' = False, device: 'Device | str | None' = None, dtype: 'DtypeLike | None' = None) -> 'None'
Initialize self. See help(type(self)) for accurate signature.
extra_repr
extra_repr(self) -> 'str'
extra_repr() -> str
One line of per-class detail for :meth:__repr__; a layer prints
its geometry here (in_features=64, out_features=256).
forward
forward(self, query: 'Tensor', key: 'Tensor', value: 'Tensor', key_padding_mask: 'Tensor | None' = None, need_weights: 'bool' = True, attn_mask: 'Tensor | None' = None, average_attn_weights: 'bool' = True, is_causal: 'bool' = False) -> 'tuple[Tensor, Tensor | None]'
forward(query, key, value, key_padding_mask=None, need_weights=True, attn_mask=None, average_attn_weights=True, is_causal=False) -> (attn_output, attn_output_weights)
Args:
query: [L, N, E_q] ([N, L, E_q] with batch_first),
or unbatched [L, E_q].
key: [S, N, E_k] ([N, S, E_k] with batch_first).
value: [S, N, E_v] ([N, S, E_v] with batch_first).
key_padding_mask: [N, S]; a boolean mask marks the keys
to IGNORE with True, a floating mask adds to the
attention logits.
need_weights: also return the attention weights; the
explicit softmax path runs then, the fused kernel
otherwise.
attn_mask: [L, S] or [N * num_heads, L, S]; boolean
marks the positions to ignore with True, floating adds
to the logits.
average_attn_weights: average the returned weights over the
heads ([N, L, S]) instead of [N, num_heads, L, S].
is_causal: apply the causal mask (with attn_mask absent).
Returns:
(attn_output, attn_output_weights): the output shaped like
query with the last dim embed_dim, and the weights or
None.
reset_parameters
reset_parameters(self) -> 'None'
reset_parameters() -> None
Fill this module's own projections the way the reference module
does: a uniform Glorot draw for the input projections, zeros for
their biases, a normal Glorot draw for bias_k / bias_v. The
out_proj layer keeps its declared slots (bind them with
load_state_dict).