Grouped Query Attention

Grouped-Query Attention (GQA, Ainslie et al. 2023).

Projects queries to num_heads heads but keys and values to num_kv_heads heads. Each key/value head is shared across num_heads / num_kv_heads query heads, interpolating between multi-head attention (num_kv_heads == num_heads) and multi-query attention (num_kv_heads == 1).

Reference:

Ainslie et al. (2023), “GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints”, arXiv:2305.13245.

class src.model.attention.grouped_query_attention.GroupedQueryAttention(*args: Any, **kwargs: Any)[source]

Bases: Module

Grouped-query attention with configurable key-value heads.

Parameters:

config – Model configuration object with attributes hidden_size, num_heads, num_kv_heads, dropout, use_bitnet, and optionally mode ("encoder" or "decoder").

Raises:
  • ValueError – If hidden_size is not divisible by num_heads.

  • ValueError – If num_kv_heads is not in [1, num_heads].

  • ValueError – If num_heads is not divisible by num_kv_heads.

forward(x: torch.Tensor, logical_layer_idx: int | None = None) torch.Tensor[source]