TensorFlow中CSR稀疏矩阵表示与转置函数使用疑问
解答
TensorFlow中tf.raw_ops稀疏矩阵系列接口使用的CSR格式,和标准三数组CSR规范仅存在上层封装差异,底层存储逻辑完全一致。
- 你直接
tf.print输出的不是CSR的原始存储内容,而是TensorFlow自定义的CSRSparseMatrix不透明资源对象。这类对象不会直接暴露内部存储数组,直接打印只会展示对象类型、内存地址、矩阵形状这类元信息,和你预期的三数组输出不符是正常现象。 - TensorFlow的CSR稀疏矩阵底层完全遵循标准CSR存储规范,三个核心数组和通用定义一一对应:
values:对应标准CSR的A数组,按行优先顺序存储所有非零元素值row_pointers:对应标准CSR的IA数组,存储每行首个非零元素在values数组中的偏移位置col_indices:对应标准CSR的JA数组,存储每个非零元素对应的列索引
- 要查看底层三数组结构,不能直接打印封装对象,需要调用专用解析接口
tf.raw_ops.CSRSparseMatrixComponents提取内部字段,示例代码如下:
import tensorflow as tf import numpy as np dense_input=np.array([[1,2,0], [1,2,0],[1,2,0]]).astype(np.float32) csr=tf.raw_ops.DenseToCSRSparseMatrix( dense_input=dense_input, indices=[[0,0],[0,1],[1,0],[1,1],[2,0],[2,1]], name=None ) # 提取CSR底层存储数组,二维矩阵批次索引固定传0 row_ptrs, col_inds, vals = tf.raw_ops.CSRSparseMatrixComponents( csr_sparse_matrix=csr, index=0, type=tf.float32 ) print("行偏移数组(IA/row_pointers):", row_ptrs.numpy()) print("列索引数组(JA/col_indices):", col_inds.numpy()) print("非零值数组(A/values):", vals.numpy())
运行上述代码就能得到和标准CSR完全一致的三个数组结构。
- 补充说明:TensorFlow的CSR实现原生支持批量稀疏矩阵存储,因此接口设计上增加了批次索引参数。处理二维稀疏矩阵时批次索引固定传0即可;处理三维批量稀疏矩阵时,依次传入0、1…的批次索引,就能拿到每个样本对应的CSR三数组。
- 你调用
SparseMatrixTranspose得到的返回值同样是封装后的CSRSparseMatrix对象,用上述相同方法提取组件,即可验证转置后的CSR数组和你自行实现的基础转置逻辑结果对齐。
内容的提问来源于stack exchange,提问作者yys_c
相关产品推荐
相关产品推荐

