使用tf.tensordot 无法(以有效的方式)完成您想要执行的操作。但是,该操作有一个专用函数tf.linalg.matvec,它可以开箱即用地处理批次。你也可以用tf.einsum做同样的事情,比如tf.einsum('bmn,bn->bm', my_tensors, my_vectors)。
关于tf.tensordot,通常它计算两个给定张量的“全部与全部”乘积,但匹配和减少一些轴。当没有给出轴时(您必须显式传递 axes=[[], []] 来执行此操作),它会创建一个张量,其中两个输入的维度连接在一起。所以,如果你有my_tensors 形状为(b, m, n) 和my_vectors 形状为(b, n) 并且你这样做:
res = tf.tensordot(my_tensors, my_vectors, axes=[[], []])
你得到res,形状为(b, m, n, b, n),这样res[p, q, r, s, t] == my_tensors[p, q, r] * my_vectors[s, t]。
axes 参数用于指定输入张量中“匹配”的维度。沿匹配轴的值相乘和相加(如点积),因此这些匹配的维度会从输出中减少。 axes 可以采用两种不同的形式:
- 如果是单个整数
N,则第一个参数的最后一个 N 维度与 b 的第一个 N 维度匹配。在您的示例中,这对应于 my_tensor 和 my_vector 中具有 n 元素的维度。
- 如果是列表,则必须包含两个子列表
axes_a 和axes_b,每个子列表具有相同数量的整数N。在这种形式中,您明确指出给定值的哪些维度是匹配的。因此,在您的示例中,您可以传递 axes=[[1], [0]],这意味着“将第一个参数 (my_tensor) 的维度 1 与第二个参数 (my_vector) 的维度 0 匹配”。
如果您现在有形状为(b, m, n) 的my_tensors 和形状为(b, n) 的my_vectors,那么您需要将第一个的维度2 与第二个的维度1 匹配,所以你可以通过axes=[[2], [1]]。然而,这会给你一个结果res,形状为(b, m, b),这样res[i, :, j] 是矩阵my_tensors[i] 和向量my_vectors[j] 的乘积。然后,您可以只获取您想要的结果(i == j 的结果),或多或少有些复杂,例如 tf.transpose(tf.linalg.diag_part(tf.transpose(res, [1, 0, 2]))),但您需要做的计算量远远超过获得相同结果所需的量。