如何在TensorFlow2.0中实现稀疏矩阵与3D张量的张量乘法并提取对角线?
解决思路与实现步骤
你的核心问题是稀疏矩阵无法直接和3D张量做常规矩阵乘法,且不能转稠密矩阵。我们可以跳过生成中间巨大张量C的步骤,直接利用稀疏矩阵的非零元素特性计算目标对角线结果,既省内存又高效。
核心逻辑拆解
你需要的最终结果等价于:result[i,k] = sum_j (A[i,j] * B[j,i,k])
完全不需要先计算完整的C矩阵,直接针对每个非零的A[i,j],提取对应位置的B[j,i,k]相乘后按i分组求和即可。
TensorFlow 实现代码
import tensorflow as tf # 假设A是你的稀疏矩阵(tf.SparseTensor类型) # B是形状为(32600, 60400, 64)的稠密张量 # 1. 提取稀疏矩阵A的非零元素索引和值 a_indices = A.indices # 形状:(nnz, 2),每一行是[i, j] a_values = A.values # 形状:(nnz,) # 2. 从B中取出对应A非零位置的B[j, i, :] # 把A的索引[i,j]转为[j,i],对应B的索引格式 b_indices = tf.gather(a_indices, [1, 0], axis=1) # 提取B中对应位置的64维向量 b_selected = tf.gather_nd(B, b_indices) # 形状:(nnz, 64) # 3. A的非零值与对应B值相乘 multiplied = tf.expand_dims(a_values, axis=1) * b_selected # 形状:(nnz, 64) # 4. 按i维度分组求和,得到每个i对应的结果 i_indices = a_indices[:, 0] result = tf.math.unsorted_segment_sum( multiplied, segment_ids=i_indices, num_segments=60400 # 对应A的第一维度大小 ) # result形状为(60400, 64),就是你要的最终结果
关键说明
- 常规方法失效原因:
tf.sparse.sparse_dense_matmul仅支持2D矩阵运算,3D张量无法直接输入;tf.einsum不支持稀疏张量作为输入,因此会报类型不匹配错误。 - 本方法优势:只处理A的非零元素,完全避开了生成(60400,60400,64)这种超大张量,内存占用仅和A的非零元素数量相关,适合大规模稀疏场景。
内容的提问来源于stack exchange,提问作者LLLLLL
相关产品推荐
相关产品推荐

