scipy.interpolate.RBFInterpolator中epsilon参数的行为疑问及与Rbf的差异
scipy.interpolate.RBFInterpolator中epsilon参数的行为疑问及与Rbf的差异
我一直在尝试将一些使用scipy.interpolate.Rbf的代码迁移到scipy.interpolate.RBFInterpolator。不过我发现后者的epsilon参数行为似乎有所不同——实际上在我的测试中,至少使用multiquadric核时,我将这个参数调整几个数量级,输出结果都没有明显变化。而从现有的scipy文档中,我也搞不清楚它的工作原理(文档很好地描述了RBFInterpolator的平滑处理,但似乎只有Rbf的文档明确说明了epsilon作为尺度参数是如何进入核函数的)。
为了展示这个现象,我准备了以下测试代码——严格来说不是最小可复现示例,因为我用了ROOT来可视化输出。
import sys import numpy as np import scipy import ROOT as rt def true_func(x,y,sigma,zscale): # first Gaussian, at center g1 = zscale * np.exp(-0.5 * (np.square(x) + np.square(y)) / np.square(sigma)) # second Gaussian, offset height_scale = 0.5 xp = x - 3. * sigma g2 = height_scale * zscale * np.exp(-0.5 * (np.square(xp) + np.square(y)) / np.square(sigma)) # add a couple sharper peaks positions = [(0,-2 * sigma), (0, -1 * sigma), (0, 2 * sigma)] spikes = 0 pow = 1.1 sig = sigma / 10 height_scale = 2. for pos in positions: xp = x - pos[0] yp = y - pos[1] spikes += height_scale * zscale * np.exp(-0.5 * (np.power(np.abs(xp),pow) + np.power(np.abs(yp),pow)) / np.power(sig,pow)) return g1 + g2 + spikes def test(new=False): N = 15 xscale = 100 xlin = np.linspace(-xscale*N,xscale*N,2 * N) ylin = np.linspace(-xscale*N,xscale*N,2 * N) x,y = np.meshgrid(xlin,ylin) # generate our z values rng = np.random.default_rng() zscale = 10 sigma = xscale * N / 4 z = true_func(x,y,sigma,zscale) z += 0.1 * zscale * rng.uniform(size=z.shape) xf = x.flatten() yf = y.flatten() zf = z.flatten() # Create two interpolators with different values of epsilon, # keep everything else the same between them. basis = 'multiquadric' rbf_dict = {} epsilon_vals = [0.1, 1000] if(new): for epsilon in epsilon_vals: rbf_dict[epsilon] = scipy.interpolate.RBFInterpolator( np.vstack((xf,yf)).T, zf, kernel=basis, epsilon=epsilon ) else: for epsilon in epsilon_vals: rbf_dict[epsilon] = scipy.interpolate.Rbf( xf,yf, zf, kernel=basis, epsilon=epsilon ) # now evaluate the two interpolators on the grid points = np.stack((x.ravel(), y.ravel()), axis=-1) if(new): evals = {key:val(points) for key,val in rbf_dict.items()} else: evals = {key:val(x,y) for key,val in rbf_dict.items()} diffs = {} for i,(key,val) in enumerate(evals.items()): if(i == 0): continue diffs[key] = (val - evals[epsilon_vals[0]]) print(np.max(diffs[key])) # now plot things dims = (1600,1200) c = rt.TCanvas('c1','c1',*dims) c.Divide(2,2) c.cd(1) true_graph = rt.TGraph2D(len(zf),xf,yf,zf) true_graph.SetName('true_graph') true_graph.Draw('SURF2Z') if(new): true_graph.SetTitle("scipy.interpolate.RBFInterpolator Test") else: true_graph.SetTitle("scipy.interpolate.Rbf Test") true_graph.GetXaxis().SetTitle("x") true_graph.GetYaxis().SetTitle("y") true_graph.GetZaxis().SetTitle("z") true_graph.SetNpx(80) true_graph.SetNpy(80) # now draw the two interpolations with the largest difference in epsilon interp1 = rt.TGraph2D(len(zf),xf,yf,evals[epsilon_vals[0]]) interp1.SetName('interp1') interp2 = rt.TGraph2D(len(zf),xf,yf,evals[epsilon_vals[-1]]) interp2.SetName('interp2') interp1.SetLineColor(rt.kRed) interp2.SetLineColor(rt.kGreen) # interp2.SetLineWidth(2) interp2.SetLineStyle(rt.kDotted) for g in (interp1, interp2): g.SetNpx(80) g.SetNpy(80) interp1.Draw('SAME SURF1') interp2.Draw('SAME SURF1') c.cd(2) diff_graph = rt.TGraph2D(len(zf),xf,yf,diffs[epsilon_vals[-1]]) diff_graph.SetName('diff_graph') diff_graph.SetTitle('Difference between interpolations, epsilon #in [{}, {}]'.format(epsilon_vals[0],epsilon_vals[-1])) diff_graph.Draw('SURF1') rt.gPad.SetLogz() c.cd(3) interp1.Draw('CONTZ') interp1.SetTitle('Interpolation with epsilon = {}'.format(epsilon_vals[0])) c.cd(4) interp2.Draw('CONTZ') interp2.SetTitle('Interpolation with epsilon = {}'.format(epsilon_vals[-1])) c.Draw() if(new): c.SaveAs('c_new.pdf') else: c.SaveAs('c_old.pdf') return def main(args): test(new=False) test(new=True) if(__name__=='__main__'): main(sys.argv)
这段代码会生成两组输出图:
- 使用
scipy.interpolate.Rbf的输出 - 使用
scipy.interpolate.RBFInterpolator的输出
可能我遇到的是RBF方法其他变化带来的结果,但两种方法的结果差异如此之大,而且使用RBFInterpolator时,不同epsilon值的插值结果几乎没有区别,这让我感到很奇怪。我也大致浏览了scipy的相关源代码,但目前还没搞清楚问题所在。
希望能得到大家的帮助或建议,谢谢!
备注:内容来源于stack exchange,提问作者JTO
相关产品推荐
相关产品推荐

