You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.24 04:15:07