def forward(self, q, k, v, mask=None): # Q, K, V shape: [batch, seq_len, heads, head_dim] attn_scores = torch.einsum('bqhd, bkhd -> bhqk', q, k) / (self.head_dim ** 0.5) if mask is not None: attn_scores = attn_scores.masked_fill(mask == 0, float('-inf')) attn_probs = torch.softmax(attn_scores, dim=-1) output = torch.einsum('bhqk, bkhd -> bqhd', attn_probs, v) return output