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

SciPy三维曲线拟合报错:函数调用结果非有效浮点数组

SciPy curve_fit拟合三维函数的错误修正

问题描述

尝试用SciPy的curve_fit拟合三维函数f(x,y,z)=xyz+1,编写了拟合函数:

def func1(data, a, b):
    return data[:,0]*data[:,1]*data[:,2]*a + b

因为curve_fit仅接受单个变量输入,计划将数据拆分为x=data[:,0]、y=data[:,1]、z=data[:,2],其余代码如下:

N = 50
L = 1
line = np.linspace(0, L, N, dtype=float)
X, Y, Z = np.meshgrid(line, line, line)

def test(x, y, z):
    return x*y*z + 1

K = test(X, Y, Z)
guess = (1, 1)
params, pcov = sp.optimize.curve_fit(func1, K[:,:3], K[:,3], guess)
print(params)

运行时报错:

---------------------------------------------------------------------------
ValueError                                Traceback (most recent call last)
ValueError: object too deep for desired array
---------------------------------------------------------------------------
error                                     Traceback (most recent call last)
~\AppData\Local\Temp/ipykernel_17432/2610089433.py in <module>
      7 K = test(X, Y, Z)
      8 guess = (1, 1)
----> 9 params, pcov = sp.optimize.curve_fit(func1, K[:,:3], K[:,3], guess)
     10 print(params)

d:\Anaconda\lib\site-packages\scipy\optimize\_minpack_py.py in curve_fit(f, xdata, ydata, p0, sigma, absolute_sigma, check_finite, bounds, method, jac, full_output, **kwargs)
    832             raise TypeError(f"The number of func parameters={n} must not"
    833                             f" exceed the number of data points={ydata.size}")
--> 834         res = leastsq(func, p0, Dfun=jac, full_output=1, **kwargs)
    835         popt, pcov, infodict, errmsg, ier = res
    836         ysize = len(infodict['fvec'])

d:\Anaconda\lib\site-packages\scipy\optimize\_minpack_py.py in leastsq(func, x0, args, Dfun, full_output, col_deriv, ftol, xtol, gtol, maxfev, epsfcn, factor, diag)
    421         if maxfev == 0:
    422             maxfev = 200*(n + 1)
--> 423         retval = _minpack._lmdif(func, x0, args, full_output, ftol, xtol,
    424                                  gtol, maxfev, epsfcn, factor, diag)
    425     else:

error: Result from function call is not a proper array of floats.

疑问:是否是数据处理方式与参考示例不同?如何正确使用func1拟合test函数?

错误原因与修正方案

核心错误分析

  • 数据维度错误:meshgrid生成的X,Y,Z是三维数组(形状为(50,50,50)),调用test(X,Y,Z)得到的K也是三维数组。K[:,:3]和K[:,3]的切片逻辑完全错误,导致输入curve_fit的xdata和ydata维度混乱,不符合函数要求的格式。
  • 输入格式不匹配:curve_fit要求xdata为二维数组时,每行对应一组自变量(x,y,z);ydata为一维数组,对应每组自变量的函数值。原代码未将三维网格数据转换为符合要求的扁平格式。

修正后的代码

import numpy as np
from scipy import optimize as sp_opt

def func1(data, a, b):
    return data[:,0] * data[:,1] * data[:,2] * a + b

N = 50
L = 1
line = np.linspace(0, L, N, dtype=float)
X, Y, Z = np.meshgrid(line, line, line)

def test(x, y, z):
    return x*y*z + 1

# 生成真实函数值
K = test(X, Y, Z)

# 将三维网格数据扁平化,组合成(n_samples, 3)的自变量数组
xdata = np.stack([X.ravel(), Y.ravel(), Z.ravel()], axis=1)
# 将函数值扁平化,得到一维数组
ydata = K.ravel()

guess = (1, 1)
params, pcov = sp_opt.curve_fit(func1, xdata, ydata, guess)
print(params)

关键步骤解释

  • 扁平化数据:用ravel()将三维的X,Y,Z和K转为一维数组,再通过np.stack将三个自变量数组组合成二维的xdata,每行对应一组(x,y,z)数据。
  • 匹配输入格式:处理后xdata形状为(N**3, 3),ydata形状为(N**3,),完全符合curve_fit的输入要求,此时func1可正确提取每组的x,y,z值计算。

内容的提问来源于stack exchange,提问作者Michael Adrian Javier

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 10:35:25