TensorFlow中2D稀疏张量与3D稠密矩阵点积实现及替代方案问询
首先直接给结论:完全可以用2D稀疏张量替代你原来的2D稠密张量,得到等价的计算结果——前提是稀疏张量的非零元素位置和值要和原稠密张量对应上(后面会说明你给出的sparse2d和原dense2d的差异)。
一、对应代码实现
先修正一下你给出的sparse2d:原dense2d的非零元素在(0,1)位置值为1.0,(0,2)位置值为2.0,所以正确的稀疏张量定义应该是:
sparse2d = tf.SparseTensor( indices=[[0, 1], [0, 2]], # 对应原dense2d的非零位置 values=[1.0, 2.0], dense_shape=[3, 3] )
如果要用这个稀疏张量替代原dense2d,实现和原代码等价的3D张量与2D稀疏张量的点积,因为tf.tensordot不直接支持SparseTensor,而tf.sparse_tensor_dense_matmul只支持2D张量,我们可以用reshape + 稀疏-稠密矩阵乘法 + 还原shape的方式实现:
import tensorflow as tf # 定义3D占位符张量 shape = [2, 4, 3] dense3d = tf.placeholder("float", shape=shape) # 定义等价于原dense2d的稀疏张量 sparse2d = tf.SparseTensor( indices=[[0, 1], [0, 2]], values=[1.0, 2.0], dense_shape=[3, 3] ) # 将3D张量reshape为2D,执行稀疏-稠密矩阵乘法,再reshape回3D dense3d_reshaped = tf.reshape(dense3d, [-1, shape[-1]]) # shape变为[8,3] res_reshaped = tf.sparse_tensor_dense_matmul(dense3d_reshaped, sparse2d) # shape变为[8,3] res = tf.reshape(res_reshaped, shape) # 还原为原shape[2,4,3]
如果你坚持用你给出的sparse2d(indices=[[0,0], [1,2]]),只需要替换上面的sparse2d定义即可,计算逻辑完全一致。
二、关于tf.sparse_tensor_dense_matmul不支持高秩张量的替代方案
除了上面的reshape方案,还有两种常用的替代思路:
1. 使用tf.map_fn遍历3D张量的每个切片
把3D张量的每个2D切片(shape=[4,3])单独拿出来,和稀疏张量做乘法,再把结果拼接回去:
# 对dense3d第一个维度的每个切片执行稀疏乘法 res = tf.map_fn( lambda x: tf.sparse_tensor_dense_matmul(x, sparse2d), dense3d, dtype=tf.float32 )
这种方案更直观,适合理解,但在大张量场景下,reshape方案的性能通常更好,因为它利用了矩阵乘法的批量优化。
2. 使用tf.einsum结合稀疏张量转稠密(不推荐)
如果你不介意临时把稀疏张量转成稠密张量,可以用tf.einsum实现点积,但这样就失去了稀疏张量节省内存的优势:
# 把稀疏张量转成稠密张量 sparse2d_dense = tf.sparse.to_dense(sparse2d) # 用einsum实现等价的点积逻辑 res = tf.einsum('ijk,kl->ijl', dense3d, sparse2d_dense)
这种方案只适合稀疏度很低的场景,否则浪费内存,不如直接用原稠密张量方案。
补充说明
原代码中res.set_shape(shape)是为了显式指定结果的shape,上面的reshape和map_fn方案都会自动保留正确的shape,所以不需要额外调用set_shape。
内容的提问来源于stack exchange,提问作者alpaca

