Q is (B, Tq, D), K is (B, Tk, D) and V is (B, Tk, Dv). Return
softmax(Q @ K.transpose(-2, -1) / sqrt(D)) @ V
of shape (B, Tq, Dv). Softmax goes over the key axis, the last one, and the
scaling is by the square root of the key dimension -- without it the dot products
grow with D and push softmax into a region where its gradients vanish.
No loops; this is three batched operations.
Input
Q =
tensor([[[0.0000, 0.0833, 0.1667, 0.2500],
[0.3333, 0.4167, 0.5000, 0.5833],
[0.6667, 0.7500, 0.8333, 0.9167]]])
K =
tensor([[[0.0000, 0.1250, 0.2500, 0.3750],
[0.5000, 0.6250, 0.7500, 0.8750]]])
V =
tensor([[[0., 1., 2.],
[3., 4., 5.]]])
Output
tensor([[[1.5936, 2.5936, 3.5936],
[1.8379, 2.8379, 3.8379],
[2.0646, 3.0646, 4.0646]]])