import torch import torch.nn.functional as F # Parameters batch_size = 2 seq_len = 4 # Sequence length embed_dim = 8 # Embedding size of input tokens num_heads = 2 # Number of attention heads head_dim = embed_dim // num_heads # Example data: batch of sequences, each token is a vector x = torch.randn(batch_size, seq_len, embed_dim) # Trainable projection matrices for queries, keys, values W_q = torch.randn(embed_dim, embed_dim) W_k = torch.randn(embed_dim, embed_dim) W_v = torch.randn(embed_dim, embed_dim) # Output projection after concatenation W_o = torch.randn(embed_dim, embed_dim) def split_heads(tensor, num_heads): # Reshape last dimension for num_heads, transpose to (batch, heads, seq, head_dim) batch, seq, embed = tensor.size() tensor = tensor.view(batch, seq, num_heads, head_dim) return tensor.transpose(1, 2) def combine_heads(tensor): # Reverse: (batch, heads, seq, head_dim) -> (batch, seq, embed_dim) batch, num_heads, seq, head_dim = tensor.size() tensor = tensor.transpose(1, 2).contiguous() return tensor.view(batch, seq, num_heads * head_dim) # 1. Project input vectors for queries, keys, and values Q = x @ W_q # (batch, seq, embed_dim) K = x @ W_k V = x @ W_v # 2. Split into heads Q = split_heads(Q, num_heads) # (batch, num_heads, seq, head_dim) K = split_heads(K, num_heads) V = split_heads(V, num_heads) # 3. Scaled dot-product attention for each head scores = Q @ K.transpose(-2, -1) / (head_dim ** 0.5) # (batch, num_heads, seq, seq) weights = F.softmax(scores, dim=-1) heads = weights @ V # (batch, num_heads, seq, head_dim) # 4. Concatenate heads and final projection concat = combine_heads(heads) # (batch, seq, embed_dim) output = concat @ W_o # (batch, seq, embed_dim) print("Output shape:", output.shape) # Output: torch.Size([2, 4, 8])