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

TensorFlow中稀疏矩阵的传递与运算:稀疏矩阵与向量相乘方法

在TensorFlow中实现稀疏矩阵与向量的高效相乘

先把你的场景再理清楚:你有一个超大的矩阵,用稀疏格式[row, column, value]存储,还有一个向量,要做等效于np.dot(X, b)的乘法,而且绝对不能转成稠密矩阵——毕竟数据量太大,转了内存直接炸对吧?

别担心,TensorFlow专门有一套处理稀疏数据的工具,完全能满足你的需求,下面一步步给你讲怎么实现:

核心思路

TensorFlow里的tf.SparseTensor就是专门用来表示稀疏矩阵的,再配合tf.sparse.sparse_dense_matmul这个优化过的乘法函数,全程不用碰稠密矩阵,完美适配大规模数据场景。

具体实现步骤 & 代码

1. 拆分稀疏数据

首先把你的X_sparse拆成三个独立的列表:行索引、列索引、对应的非零值,这样才能喂给SparseTensor。

2. 创建SparseTensor对象

这里要注意,SparseTensor的索引必须是二维数组(每一行是一个(行,列)对),所以得把拆分后的行、列数组转置一下。另外还要指定稀疏矩阵的完整形状(比如你的例子里是7行5列)。

3. 执行乘法

把向量b转成TensorFlow的稠密张量,注意要变成二维的(比如(5,1)的列向量),因为sparse_dense_matmul要求输入都是二维的,最后再把结果转成一维的就行。

完整代码如下:

import tensorflow as tf
import numpy as np

# 你的稀疏矩阵数据
X_sparse = [ [1, 2, 1], [3, 0, 2], [3, 3, 3], [6, 1, 4], ]
# 你的向量b
b = [1,2,3,4,5]

# 拆分稀疏矩阵的行、列、值
rows = [item[0] for item in X_sparse]
cols = [item[1] for item in X_sparse]
vals = [item[2] for item in X_sparse]

# 处理成SparseTensor需要的索引格式:二维数组,每一行是(行,列)
sparse_indices = tf.convert_to_tensor([rows, cols], dtype=tf.int64)
sparse_indices = tf.transpose(sparse_indices)
# 非零值转成张量
sparse_values = tf.convert_to_tensor(vals, dtype=tf.float32)
# 指定稀疏矩阵的完整形状:7行5列
sparse_shape = (7, 5)

# 创建SparseTensor对象
X_tensor = tf.sparse.SparseTensor(indices=sparse_indices, values=sparse_values, dense_shape=sparse_shape)

# 把向量b转成二维张量(列向量)
b_tensor = tf.convert_to_tensor(b, dtype=tf.float32)
b_tensor = tf.expand_dims(b_tensor, axis=1)

# 执行稀疏-稠密矩阵乘法
result = tf.sparse.sparse_dense_matmul(X_tensor, b_tensor)

# 打印结果,和numpy的结果对比
print("TensorFlow计算结果:")
print(result.numpy().flatten())

# 用numpy验证正确性
X_dense = np.array([ [0, 0, 0, 0, 0], [0, 0, 1, 0, 0], [0, 0, 0, 0, 0], [2, 0, 0, 3, 0], [0, 0, 0, 0, 0], [0, 0, 0, 0, 0], [0, 4, 0, 0, 0] ])
print("NumPy验证结果:")
print(np.dot(X_dense, b))

关键注意点

  • tf.SparseTensor只存非零元素的索引和值,内存占用极低,完全适合超大矩阵。
  • tf.sparse.sparse_dense_matmul是TensorFlow专门优化的函数,内部不会把稀疏矩阵转成稠密格式,计算效率很高。
  • 索引格式一定要对:必须是二维的(n,2)数组,n是非零元素的个数,别搞错形状了。
  • 向量b要转成二维的,不然乘法会报错,最后用flatten()转成一维就和np.dot的结果格式一致了。

内容的提问来源于stack exchange,提问作者Tim

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.29 06:50:05