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

TensorFlow稀疏张量共轭梯度为何比Scipy实现慢5倍?

为什么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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 00:46:00