如何加速多波段图像的多维高斯分布计算?
多波段图像多维高斯分布计算的性能优化
问题背景
需要为多波段图像计算多维高斯分布,将每个像素的多波段取值视为向量,通过计算V^T * inv(covariance) * V实现。但使用np.apply_along_axis处理(3,1000,1000)尺寸图像时速度过慢,需优化。
原实现代码:
import numpy as np import numpy.linalg as lg X,Y=np.meshgrid(np.arange(-10,11,1/10),np.arange(-10,11,1/10),indexing='xy') Z=np.stack((X,Y)) cov=np.array([[10,-0.4],[-0.4,1]]) W=lg.inv(cov) def ee(x): global W x=x.reshape((-1,1)) return lg.matmul(lg.matmul(x.transpose(),W),x)[0,0] tt=np.apply_along_axis(ee,0,Z) p=np.exp(-0.5*tt)/(np.power(2*np.pi,cov.shape[0]/2)*np.power(lg.det(cov),0.5)) import matplotlib.pyplot as plt plt.imshow(p)
优化方案
1. 向量化计算替代循环
np.apply_along_axis本质是Python层循环,效率极低。通过调整数组维度,利用numpy底层优化的向量化操作一次性完成所有像素的二次型计算:
import numpy as np import numpy.linalg as lg import matplotlib.pyplot as plt # 生成数据 X,Y=np.meshgrid(np.arange(-10,11,1/10),np.arange(-10,11,1/10),indexing='xy') Z=np.stack((X,Y)) # 形状(2, 210, 210) cov=np.array([[10,-0.4],[-0.4,1]]) W=lg.inv(cov) # 将图像转为二维数组:每行对应一个像素的多波段向量 Z_reshaped = Z.reshape(Z.shape[0], -1).T # 形状(210*210, 2) # 批量计算所有像素的二次型:两种等价方式选其一 tt = np.einsum('ni,ij,nj->n', Z_reshaped, W, Z_reshaped) # 或:tt = np.sum((Z_reshaped @ W) * Z_reshaped, axis=1) # 计算高斯概率密度并还原图像形状 dim = cov.shape[0] norm_factor = np.power(2*np.pi, dim/2) * np.sqrt(lg.det(cov)) p = np.exp(-0.5*tt).reshape(Z.shape[1], Z.shape[2]) / norm_factor plt.imshow(p) plt.show()
2. Cholesky分解优化数值稳定性与效率
若协方差矩阵为正定矩阵,用Cholesky分解替代直接求逆,计算更快且数值更稳定:
import numpy as np import numpy.linalg as lg import matplotlib.pyplot as plt X,Y=np.meshgrid(np.arange(-10,11,1/10),np.arange(-10,11,1/10),indexing='xy') Z=np.stack((X,Y)) cov=np.array([[10,-0.4],[-0.4,1]]) # Cholesky分解:cov = L @ L.T,inv(cov) = L.T^{-1} @ L^{-1} L = lg.cholesky(cov) L_inv = lg.inv(L) Z_reshaped = Z.reshape(Z.shape[0], -1).T # (N, 2) # 二次型等价于变换后向量的平方和 Z_transformed = Z_reshaped @ L_inv.T tt = np.sum(Z_transformed ** 2, axis=1) dim = cov.shape[0] norm_factor = np.power(2*np.pi, dim/2) * np.prod(np.diag(L)) # det(cov) = (det(L))²,sqrt(det(cov))=prod(diag(L)) p = np.exp(-0.5*tt).reshape(Z.shape[1], Z.shape[2]) / norm_factor plt.imshow(p) plt.show()
3. GPU加速(可选)
若有NVIDIA GPU,用CuPy替代numpy,代码几乎完全兼容,可获得数量级的速度提升:
import cupy as cp import cupy.linalg as lg import matplotlib.pyplot as plt # 用CuPy生成数据 X,Y=cp.meshgrid(cp.arange(-10,11,1/10),cp.arange(-10,11,1/10),indexing='xy') Z=cp.stack((X,Y)) cov=cp.array([[10,-0.4],[-0.4,1]]) W=lg.inv(cov) Z_reshaped = Z.reshape(Z.shape[0], -1).T tt = cp.einsum('ni,ij,nj->n', Z_reshaped, W, Z_reshaped) dim = cov.shape[0] norm_factor = cp.power(2*cp.pi, dim/2) * cp.sqrt(lg.det(cov)) p = cp.exp(-0.5*tt).reshape(Z.shape[1], Z.shape[2]) / norm_factor # 转回numpy显示图像 plt.imshow(cp.asnumpy(p)) plt.show()
内容的提问来源于stack exchange,提问作者mohammad sadeg
相关产品推荐
相关产品推荐

