Skip to main content

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).