attn_weights = torch.softmax(masked / d_out_kq**0.5, dim=1) print(attn_weights)