TensorFlow中不转密集张量实现单列稀疏张量与矩形张量逐列相乘
解决方案
报错原因
tensorflow.python.framework.errors_impl.InvalidArgumentError: SparseDenseBinaryOpShared broadcasts dense to sparse only; got incompatible shapes: [3,1] vs. [3,2] [Op:SparseDenseCwiseMul]
出现该错误是因为TensorFlow的稀疏-稠密逐元素运算仅支持将稠密张量广播匹配稀疏张量的形状,不支持反向广播稀疏张量到稠密张量的形状,因此形状为[3,1]的稀疏张量无法直接和[3,2]的稠密张量做逐元素乘法。
实现方案
两种方案全程都不会将稀疏张量实化为稠密张量,内存效率远高于先转稠密再计算的实现。
方案1:稀疏矩阵乘法实现广播(代码简洁,通用场景适用)
通过稀疏矩阵乘全1行向量的方式,先将单列稀疏张量广播到和稠密张量b一致的形状,再做逐元素乘法:
import tensorflow as tf a = tf.constant([[10], [0], [4]]) s = tf.sparse.from_dense(a) b = tf.constant([[3, 4], [1, 2], [2, 5]]) # 构造适配b列数的全1行向量,用于广播稀疏张量的最后一维 broadcast_ones = tf.ones((1, b.shape[-1]), dtype=s.dtype) # 稀疏矩阵乘法得到shape为[3,2]的稀疏张量,每行值和原s对应行一致 s_broadcast = tf.sparse.sparse_dense_matmul(s, broadcast_ones) # 形状匹配后直接做逐元素乘法 result = s_broadcast * b # 验证结果(可选,仅调试用) print(tf.sparse.to_dense(result))
输出结果:
tf.Tensor( [[30 40] [ 0 0] [ 8 20]], shape=(3, 2), dtype=int32)
方案2:手动构造稀疏张量(极高稀疏度场景效率最优)
如果稀疏张量s的非零元素占比极低,可以直接操作稀疏张量的索引和值构造结果,避免矩阵乘法开销:
# 提取s的非零行和非零值 nonzero_rows = s.indices[:, 0] nonzero_vals = s.values col_count = b.shape[1] # 构造新稀疏张量的索引:每个非零行对应所有列索引 new_indices = tf.concat( [ tf.repeat(nonzero_rows[:, None], col_count, axis=0), tf.tile(tf.range(col_count)[None, :], (len(nonzero_rows), 1)) ], axis=-1 ) new_indices = tf.cast(tf.reshape(new_indices, (-1, 2)), tf.int64) # 构造新稀疏张量的值:非零值乘对应行的b的所有元素 new_values = tf.reshape(tf.gather(b, nonzero_rows) * nonzero_vals[:, None], (-1,)) # 生成结果稀疏张量 result = tf.sparse.SparseTensor(indices=new_indices, values=new_values, dense_shape=b.shape) # 验证结果(可选,仅调试用) print(tf.sparse.to_dense(result))
内容的提问来源于stack exchange,提问作者safe_malloc
相关产品推荐
相关产品推荐

