def scaled_dot_product_attention(self, Q, K, V, valid_lens): d_k = Q.shape[-1] scores = np.matmul(Q, K.transpose(0, 2, 1)) / np.sqrt(d_k) if valid_lens is not None: mask = np.arange(scores.shape[-1]) < valid_lens[:, None] scores = np.where(mask[:, None, :], scores, -np.inf) attention_weights = np.exp(scores - np.max(scores, axis=-1, keepdims=True)) attention_weights /= attention_weights.sum(axis=-1, keepdims=True) return np.matmul(attention_weights, V)