TensorFlow稀疏张量共轭梯度为何比Scipy实现慢5倍?
问题背景
我用TensorFlow实现了共轭梯度法对稀疏矩阵求逆,测试矩阵是有限元得到的质量矩阵与刚度矩阵之和,条件良好。在Colab环境下,和Scipy同方法同数据对比,计算结果一致,但TensorFlow速度慢5倍——Scipy耗时0.27s,TensorFlow耗时1.37s。需要处理10万×10万级的大矩阵,无法转换为稠密矩阵,想知道TensorFlow版本算法缓慢的原因。
测试代码如下:
import tensorflow as tf import numpy as np from scipy.sparse import coo_matrix,linalg import os import sys os.environ['TF_CPP_MIN_LOG_LEVEL'] = '2' from time import time from scipy.spatial import Delaunay def create_mesh(Lx=1,Ly=1,Nx=100,Ny=100): mesh0=dict() dx = Lx/Nx dy = Ly/Ny XX,YY=np.meshgrid(np.arange(0,Lx+dx,dx),np.arange(0,Ly+dy,dy)) points=np.vstack((XX.ravel(),YY.ravel())).T #np.random.shuffle(points) tri = Delaunay(points) mesh0['Pts']=np.copy(points).astype(np.float32) mesh0['Tria']=np.copy(tri.simplices).astype(int) return(mesh0) def eval_connectivity(mesh0): print('computing mesh connectivity') npt=mesh0['Pts'].shape[0] connectivity = {} for jpt in range(npt): connectivity[jpt] = [] for Tria in mesh0['Tria']: for ilpt in range(3): iglobalPt=Tria[ilpt] for jlpt in range(1+ilpt,3): jglobalPt=Tria[jlpt] connectivity[iglobalPt].append(jglobalPt) connectivity[jglobalPt].append(iglobalPt) for key,value in connectivity.items(): connectivity[key]=np.unique(np.array(value,dtype=int)) return(connectivity) def eval_local_mass(mesh0,iTri): lmass = np.zeros(shape=(3,3),dtype=np.float32) Tria=mesh0['Tria'][iTri] v10 = mesh0['Pts'][Tria[1],:]-mesh0['Pts'][Tria[0],:] v20 = mesh0['Pts'][Tria[2],:]-mesh0['Pts'][Tria[0],:] N12 = np.cross(v10,v20) Tsurf = 0.5*np.linalg.norm(N12) for ipt in range(3): lmass[ipt,ipt]=1.0/12.0 for jpt in range(1+ipt,3): lmass[ipt,jpt] = 1.0/24.0 lmass[jpt,ipt] = lmass[ipt,jpt] lmass = 2.0*Tsurf*lmass return(lmass) def eval_local_stiffness(mesh0,iTri): Tria = mesh0['Tria'][iTri] v10 = mesh0['Pts'][Tria[1],:]-mesh0['Pts'][Tria[0],:] v20 = mesh0['Pts'][Tria[2],:]-mesh0['Pts'][Tria[0],:] N12 = np.cross(v10,v20) Tsurf = 0.5*np.linalg.norm(N12) covbT = np.zeros(shape=(3,3),dtype=np.float32) covbT[0,:2] = v10 covbT[1,:2] = v20 covbT[2,2] = N12/(2*Tsurf) contrb = np.linalg.inv(covbT) v1 = contrb[:,0] v2 = contrb[:,1] a = np.dot(v1,v1) b = np.dot(v1,v2) c = np.dot(v2,v2) gij_c = np.array([[a,b],[b,c]],dtype=np.float32) lgrad = np.array([[-1.0,1.0,0.0], [-1.0,0.0,1.0] ],dtype=np.float32) lstif = Tsurf*np.matmul( np.matmul(lgrad.T,gij_c), lgrad ) return(lstif) def compute_vectors_sparse_matrices(mesh0): npt = mesh0['Pts'].shape[0] connect = eval_connectivity(mesh0) nzero = 0 for key,value in connect.items(): nzero += (1+value.shape[0]) I = np.zeros(shape=(nzero),dtype=int) J = np.zeros(shape=(nzero),dtype=int) VM = np.zeros(shape=(nzero),dtype=np.float32) VS = np.zeros(shape=(nzero),dtype=np.float32) k0 = np.zeros(shape=(npt+1),dtype=int) k0[0] = 0 k = -1 for jpt in range(npt): loc_con = connect[jpt].tolist()[:] loc_con.append(jpt) loc_con = np.sort(loc_con) k0[jpt+1]=k0[jpt]+loc_con.shape[0] for jloc in range(loc_con.shape[0]): k=k+1 I[k]= jpt J[k]= loc_con[jloc] for iTr, Tria in enumerate(mesh0['Tria']): lstiff = eval_local_stiffness(mesh0,iTr) lmass = eval_local_mass(mesh0,iTr) for iEntry,irow in enumerate(Tria): loc_con = connect[irow].tolist()[:] loc_con.append(irow) loc_con = np.sort(loc_con) for jEntry,jcol in enumerate(Tria): indexEntry = k0[irow]+np.where(loc_con==jcol)[0] VM[indexEntry] = VM[indexEntry]+lmass[iEntry,jEntry] VS[indexEntry] = VS[indexEntry]+lstiff[iEntry,jEntry] return(I,J,VM,VS) def compute_global_sparse_matrices(mesh0): I,J,VM,VS = compute_vectors_sparse_matrices(mesh0) npt = mesh0['Pts'].shape[0] MASS = coo_matrix((VM,(I,J)),shape=(npt,npt)) STIFF = coo_matrix((VS,(I,J)),shape=(npt,npt)) return(MASS,STIFF) def compute_global_sparse_tensors(mesh0): I,J,VM,VS = compute_vectors_sparse_matrices(mesh0) npt = mesh0['Pts'].shape[0] indices = np.hstack([I[:,np.newaxis], J[:,np.newaxis]]) MASS = tf.sparse.SparseTensor(indices=indices, values=VM.astype(np.float32), dense_shape=[npt, npt]) STIFF = tf.sparse.SparseTensor(indices=indices, values=VS.astype(np.float32), dense_shape=[npt, npt]) return(MASS,STIFF) def compute_matrices_scipy(mesh0): MASS,STIFF = compute_global_sparse_matrices(mesh0) return(MASS,STIFF) def compute_matrices_tensorflow(mesh0): MASS,STIFF = compute_global_sparse_tensors(mesh0) return(MASS,STIFF) def conjgrad_scipy(A,b,x0,niter=100,toll=1.e-5): x = np.copy(x0) r = b - A * x p = np.copy(r) rsold = np.dot(r,r) for it in range(niter): Ap = A * p alpha = rsold /np.dot(p,Ap) x += alpha * p r -= alpha * Ap rsnew = np.dot(r,r) if (np.sqrt(rsnew) < toll): break p = r + (rsnew / rsold) * p rsold = rsnew return([x,it,np.sqrt(rsnew)]) def conjgrad_tensorflow(A,b,x0,niter=100,toll=1.e-5): x = x0 r = b - tf.sparse.sparse_dense_matmul(A,x) p = r rsold = tf.reduce_sum(tf.multiply(r, r)) for it in range(niter): Ap = tf.sparse.sparse_dense_matmul(A,p) alpha = rsold /tf.reduce_sum(tf.multiply(p, Ap)) x += alpha * p r -= alpha * Ap rsnew = tf.reduce_sum(tf.multiply(r, r)) if (tf.sqrt(rsnew) < toll): break p = r + (rsnew / rsold) * p rsold = rsnew return([x,it,tf.sqrt(rsnew)]) mesh = create_mesh(Lx=10,Ly=10,Nx=100,Ny=100) x0 = tf.constant( (mesh['Pts'][:,0]<5 ).astype(np.float32) ) nit_time = 10 dcoef = 1.0 maxit = x0.shape[0]//2 stoll = 1.e-6 print('nb of nodes:\t{}'.format(mesh['Pts'].shape[0])) print('nb of trias:\t{}'.format(mesh['Tria'].shape[0])) t0 = time() MASS0,STIFF0 = compute_matrices_scipy(mesh) elapsed_scipy=time()-t0 print('Matrices; elapsed: {:3.5f} s'.format(elapsed_scipy)) A = MASS0+dcoef*STIFF0 x = np.copy(np.squeeze(x0.numpy()) ) t0 = time() for jt in range(nit_time): b = MASS0*x x1,it,tol=conjgrad_scipy(A,b,x,niter=maxit,toll=stoll) x=np.copy(x1) print('time {}; iters {}; resid: {:3.2f}'.format(1+jt,it,tol) ) elapsed_scipy=time()-t0 print('elapsed, scipy: {:3.5f} s'.format(elapsed_scipy)) t0 = time() MASS,STIFF =compute_matrices_tensorflow(mesh) elapsed=time()-t0 print('Matrices; elapsed: {:3.5f} s'.format(elapsed)) x = None x1 = None A = tf.sparse.add(MASS,tf.sparse.map_values(tf.multiply, STIFF, dcoef)) x = tf.expand_dims(tf.identity(x0),axis=1) t0 = time() for jt in range(nit_time): b = tf.sparse.sparse_dense_matmul(MASS,x) x1,it,tol=conjgrad_tensorflow(A,b,x,niter=maxit,toll=stoll) x = x1 print('time {}; iters {}; resid: {:3.2f}'.format(1+jt,it,tol) ) elapsed_tf=time()-t0 print('elapsed, tf: {:3.2f} s'.format(elapsed_tf)) print('elapsed times:') print('scipy: {:3.2f} s\ttf: {:3.2f} s'.format(elapsed_scipy,elapsed_tf))
速度差异的核心原因
稀疏矩阵计算的优化深度不同
Scipy的稀疏矩阵操作依赖经过数十年打磨的底层线性代数库(如SuperLU、MKL),针对有限元这类局部稀疏、带状结构的矩阵做了专门的内存布局和计算优化,CPU向量指令(AVX2、AVX-512)利用率极高。而TensorFlow的tf.sparse模块主要为深度学习场景设计,通用稀疏线性代数的优化优先级较低,tf.sparse.sparse_dense_matmul没有针对这类矩阵结构做特殊优化。Eager模式的额外开销
当前的TensorFlow共轭梯度实现是在Eager模式下逐步执行的,每一次迭代中的张量操作都会触发独立的计算调度,累积了大量的运行时开销。而Scipy的操作直接调用底层优化函数,没有TensorFlow的图构建、设备同步等额外成本。设备调度的不合理性
Colab默认会将TensorFlow运算分配到GPU,但GPU对稀疏矩阵乘法的加速效果远不如稠密矩阵,反而可能因为CPU-GPU之间的数据传输拖慢整体速度。Scipy默认在CPU上运行,完全利用了CPU的缓存和向量优化。
优化方案
1. 用@tf.function编译计算图
将共轭梯度函数转换为TensorFlow计算图,消除Eager模式的逐次调用开销,同时开启XLA编译进一步优化:
@tf.function(jit_compile=True) def conjgrad_tensorflow(A,b,x0,niter=100,toll=1.e-5): x = x0 r = b - tf.sparse.sparse_dense_matmul(A,x) p = r rsold = tf.reduce_sum(tf.multiply(r, r)) for it in tf.range(niter): Ap = tf.sparse.sparse_dense_matmul(A,p) alpha = rsold / tf.reduce_sum(tf.multiply(p, Ap)) x += alpha * p r -= alpha * Ap rsnew = tf.reduce_sum(tf.multiply(r, r)) if tf.sqrt(rsnew) < toll: break p = r + (rsnew / rsold) * p rsold = rsnew return [x, it, tf.sqrt(rsnew)]
注意将Python的range替换为tf.range,确保图模式兼容性。
2. 转换为CSR稀疏格式
TensorFlow 2.10+支持将SparseTensor转换为CSR格式,这是稀疏矩阵乘法更高效的存储格式:
# 在创建矩阵后转换为CSR A = tf.sparse.add(MASS,tf.sparse.map_values(tf.multiply, STIFF, dcoef)) A_csr = tf.sparse.to_csr(A) # 函数内使用CSR格式进行乘法 Ap = tf.sparse.sparse_dense_matmul(A_csr, p)
3. 强制在CPU上运行
如果GPU对稀疏计算没有优势,强制TensorFlow在CPU上执行,避免数据传输开销:
with tf.device('/CPU:0'): MASS,STIFF = compute_matrices_tensorflow(mesh) A = tf.sparse.add(MASS,tf.sparse.map_values(tf.multiply, STIFF, dcoef)) A_csr = tf.sparse.to_csr(A) x = tf.expand_dims(tf.identity(x0),axis=1) t0 = time() for jt in range(nit_time): b = tf.sparse.sparse_dense_matmul(MASS,x) x1,it,tol=conjgrad_tensorflow(A_csr,b,x,niter=maxit,toll=stoll) x = x1 print('time {}; iters {}; resid: {:3.2f}'.format(1+jt,it,tol) ) elapsed_tf=time()-t0
4. 使用官方实现的共轭梯度函数
TensorFlow提供了tf.linalg.experimental.conjugate_gradient,经过官方优化,性能优于自定义实现:
def conjgrad_tensorflow(A,b,x0,niter=100,toll=1.e-5): result = tf.linalg.experimental.conjugate_gradient( A, b, x0=x0, max_iter=niter, tol=toll ) return [result.x, result.num_iterations, tf.sqrt(result.objective_value)]
内容的提问来源于stack exchange,提问作者Cesare

