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

使用2D高斯函数拟合2D直方图时出现拟合偏差问题求助

2D高斯函数拟合2D直方图失败问题

尝试用2D高斯函数拟合2D直方图时遇到困境:拟合人工生成的含噪2D高斯数据结果正常,但拟合分箱后的2D直方图时,即便初始参数很优,拟合质量仍极差——拟合结果未居中,sigma参数也与预期不符。此外实际应用中拟合出的高斯函数振幅会趋近于0(或设置的下界)。怀疑scipy.optimize.curve_fit是否适用于此类拟合,或是代码存在错误?

示例代码如下:

import numpy as np
from scipy.optimize import curve_fit
import matplotlib.pyplot as plt

def twoD_Gaussian(xy, amplitude, xo, yo, sigma_x, sigma_y, rho):
    x, y = xy
    xo = float(xo)
    yo = float(yo)
    a = 1 / sigma_x**2
    b = rho / sigma_x / sigma_y
    c = 1 / sigma_y**2
    g = amplitude * np.exp( - (1/(2*(1-rho**2)))*(a*((x-xo)**2) - 2*b*(x-xo)*(y-yo) + c*((y-yo)**2)))
    return g.ravel()

# Produce a number of points in x-y from 1 distribution. 
mean = [0,0]
cov = [[3,1.5],[1.5,1]] 
N = int(1e6)
x,y = np.random.multivariate_normal(mean,cov,N).T

nbr_bins = 100 #int(2 * N ** (1/3))
print("numbins", nbr_bins)

# Produce 2D histogram
H,xedges,yedges = np.histogram2d(x,y,bins=nbr_bins)
bin_centers_x = (xedges[:-1]+xedges[1:])/2.0
bin_centers_y = (yedges[:-1]+yedges[1:])/2.0
X,Y = np.meshgrid(bin_centers_x, bin_centers_y)

data = twoD_Gaussian((X, Y), H.max(), mean[0], mean[1], np.sqrt(cov[0][0]), np.sqrt(cov[1][1]), cov[0][1]/(np.sqrt(cov[0][0])*np.sqrt(cov[1][1])))
data_noisy = data + 0.2*np.random.normal(size=data.shape)

# Initial Guess
p0 = (H.max(),mean[0],mean[1],1.7,1,0.9)

# Curve Fit parameters with histo
coeff, var_matrix = curve_fit(twoD_Gaussian,(X,Y),H.ravel(),p0=p0)

# Curve fit on noisy data
popt, pcov = curve_fit(twoD_Gaussian, (X, Y), data_noisy, p0=p0)

print('hist fit', coeff)
print('noisy data fit', popt)

data_fitted_hist = twoD_Gaussian((X, Y), *coeff)
data_fitted_noise = twoD_Gaussian((X, Y), *popt)

# Calculate the extent of the plots
extent = [X.min(), X.max(), Y.min(), Y.max()]

# Plot the hist fit
plt.hist2d(x, y, bins = nbr_bins, cmap = 'terrain_r')
plt.contour(X, Y, data_fitted_hist.reshape(nbr_bins, nbr_bins), 5)
plt.scatter(0, 0, marker='x', s=100, color='r')
plt.xlim(extent[0], extent[1])  # Set the x-axis limits
plt.ylim(extent[2], extent[3])  # Set the y-axis limits
plt.gca().set_aspect('equal', adjustable='box')  # Set aspect ratio to be equal
plt.show()

# Plot the noisy data fit
fig, ax = plt.subplots(1, 1)
ax.scatter(0, 0, marker='x', s=100, color='w')
ax.imshow(data_noisy.reshape(nbr_bins, nbr_bins), cmap=plt.cm.jet, origin='lower', extent=extent)
ax.contour(X, Y, data_fitted_noise.reshape(nbr_bins, nbr_bins), 8, colors='w')
ax.set_xlim(extent[0], extent[1])  # Set the x-axis limits
ax.set_ylim(extent[2], extent[3])  # Set the y-axis limits
ax.set_aspect('equal', adjustable='box')  # Set aspect ratio to be equal
plt.show()

拟合结果图显示,基于直方图的拟合曲线未居中,与预期偏差明显。


问题分析与解决

1. 核心错误:直方图与网格的维度对应错位

np.histogram2d返回的H中,H[i,j]对应x在第i个区间、y在第j个区间的计数,但np.meshgrid(bin_centers_x, bin_centers_y)默认生成的网格,X按行排列x中心、Y按列排列y中心,导致坐标与直方图计数的对应关系完全错位,这是拟合结果偏移的直接原因。

修复方式二选一:

  • 转置直方图后再传入拟合:
    coeff, var_matrix = curve_fit(twoD_Gaussian,(X,Y),H.T.ravel(),p0=p0)
    
  • 调整meshgrid的索引模式为ij,让网格维度与直方图完全匹配:
    X,Y = np.meshgrid(bin_centers_x, bin_centers_y, indexing='ij')
    

2. curve_fit的参数约束与权重优化

curve_fit完全适用于2D直方图的高斯拟合,但需要解决两个关键问题:

  • 参数边界约束:实际应用中振幅趋近于0,是因为参数无约束收敛到不合理区间。通过bounds设置合理范围:
    bounds = (
        [0, -np.inf, -np.inf, 1e-3, 1e-3, -0.999],  # 下界:振幅≥0,sigma>0,rho不超出(-1,1)
        [np.inf, np.inf, np.inf, np.inf, np.inf, 0.999]  # 上界
    )
    coeff, var_matrix = curve_fit(twoD_Gaussian,(X,Y),H.T.ravel(),p0=p0,bounds=bounds)
    
  • 直方图噪声权重:直方图计数服从泊松分布,高计数的bins可靠性更高。传入sigma=np.sqrt(H.T.ravel())让拟合优先关注这些区域:
    coeff, var_matrix = curve_fit(twoD_Gaussian,(X,Y),H.T.ravel(),p0=p0,bounds=bounds,sigma=np.sqrt(H.T.ravel()),absolute_sigma=True)
    

3. 初始参数优化

初始参数的精度会影响拟合稳定性,将rho的初始值设为真实值能提升收敛效果:

true_rho = cov[0][1]/(np.sqrt(cov[0][0])*np.sqrt(cov[1][1]))
p0 = (H.max(),mean[0],mean[1],np.sqrt(cov[0][0]),np.sqrt(cov[1][1]),true_rho)

修复后验证

完成上述调整后,拟合出的参数会接近真实值,拟合曲线也会正确居中,直方图拟合结果将与含噪数据拟合结果一致。

内容的提问来源于stack exchange,提问作者Owen Syrett

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 05:45:55