class MultiHeadAttention(nn.Module): def __init__(self, d_model, num_heads): super().__init__() self.d_model = d_model self.num_heads = num_heads self.d_k = d_model // num_heads self.W_q = nn.Linear(d_model, d_model) self.W_k = nn.Linear(d_model, d_model) self.W_v = nn.Linear(d_model, d_model) self.W_o = nn.Linear(d_model, d_model) def forward(self, x): batch_size, seq_len, d_model = x.size() # Linear transformations Q = self.W_q(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) K = self.W_k(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) V = self.W_v(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2) # Scaled dot-product attention scores = torch.matmul(Q, K.transpose(-2, -1)) / (self.d_k ** 0.5) attention_weights = F.softmax(scores, dim=-1) attended_values = torch.matmul(attention_weights, V) # Concatenate heads attended_values = attended_values.transpose(1, 2).contiguous().view( batch_size, seq_len, d_model) # Final projection output = self.W_o(attended_values) return output, attention_weights # Usage attention = MultiHeadAttention(d_model=64, num_heads=8) output, weights = attention(embeddings.unsqueeze(0)) print(f"Output shape: {output.shape}") # [1, 6, 64]