import numpy as np class ScaledDotProductAttention(nn.Module): def __init__(self): super(ScaledDotProductAttention, self).__init__() def forward(self, Q, K, V): #Q, K, V of size [batch x sequence_length x dim] scores = torch.matmul(Q, K.transpose(-1, -2)) / np.sqrt(Q.shape[1]) attn = nn.Softmax(dim=-1)(scores) context = torch.matmul(attn, V) return context, attn #sanity checking q = torch.tensor([[[1.1,1.3],[0.9,0.8]]]).to(device) k = torch.tensor([[[0.9,1],[0.2,2.1]]]).to(device) v = torch.tensor([[[1.1,1.3],[0.9,0.8]]]).to(device) sample = ScaledDotProductAttention().to(device) sample(q,k,v)