Python中tanh函数拟合数据集效果极差的优化方法咨询
I've been using this code to fit a tanh function to datasets successfully before, but it's failing completely with my current dataset—the fitted curve doesn't match the data points at all. I've tried several approaches but can't figure out what's wrong. Here's my code:
import matplotlib.pyplot as plt from scipy.optimize import curve_fit import numpy as np xdata = [26, -30, -60, 80, 110, 60, 40] ydata = [53, 12.6, 15.8, 100, 102.1, 90.4, 70.2] def tanh(x, a, b, c, d ): return a * np.tanh( b * x + c ) + d p0 = [max(ydata), np.median(xdata),1,min(ydata)] popt, pcov = curve_fit(tanh, xdata, ydata, p0, method='dogbox') xModel = np.linspace(min(xdata), max(xdata)) yModel = tanh(xModel, *popt) def Plot(graphWidth, graphHeight): fig = plt.figure(figsize=(16,9)) ax = fig.add_subplot(111) ax.plot(xdata, ydata, 'D', label="data") ax.plot(xModel, yModel, label="fit") ax.legend() plt.show() graphWidth = 800 graphHeight = 600 Plot(graphWidth, graphHeight)
How can I improve the fitting result?
I've run into similar issues with nonlinear curve fitting before—curve_fit is extremely sensitive to initial parameter guesses, especially for functions like tanh that have saturation behavior. Let's break down the problem and fix it:
1. The Root Cause: Bad Initial Parameter Guess (p0)
Your original p0 is set to [max(ydata), np.median(xdata), 1, min(ydata)], which doesn't align with how the tanh function works:
- The tanh function itself ranges from
-1to1, soashould represent half the total range of your y-data, not the maximum y-value. bcontrols the steepness of the tanh curve. Your originalbis set to the median of x-data (10), which makesb*x + cway too large for most x-values—tanh will immediately hit saturation at-1or1, leaving the optimizer stuck in a bad local minimum.cshifts the tanh curve along the x-axis; your guess of1doesn't reflect the data's inflection point (which looks like it's near x=0).
2. Better Initial Parameter Guess
Let's calculate a more logical p0 based on your data:
a: Half the range of y-data →(max(ydata) - min(ydata)) / 2≈ (102.1 - 12.6)/2 ≈ 44.75d: The midpoint of y-data →(max(ydata) + min(ydata)) / 2≈ (102.1 + 12.6)/2 ≈ 57.35b: A small value to match the x-data's range (your x spans from -60 to 110, so a value like 0.05 keeps the tanh curve from saturating too quickly)c: Adjusts the x-position of the inflection point. Looking at your data, y jumps from ~15 to ~70 between x=-60 and x=40, so settingcto 0 (centering the inflection at x=0) is a reasonable starting point.
3. Updated Code with Fixes
Here's the modified code with better p0 and a few minor tweaks:
import matplotlib.pyplot as plt from scipy.optimize import curve_fit import numpy as np # Convert to numpy arrays for smoother calculations xdata = np.array([26, -30, -60, 80, 110, 60, 40]) ydata = np.array([53, 12.6, 15.8, 100, 102.1, 90.4, 70.2]) def tanh(x, a, b, c, d ): return a * np.tanh(b * x + c) + d # Improved initial parameter guess ymax, ymin = ydata.max(), ydata.min() p0 = [(ymax - ymin)/2, 0.05, 0, (ymax + ymin)/2] # Use default 'lm' method (works well for smooth nonlinear functions) # Allow more iterations if needed with maxfev popt, pcov = curve_fit(tanh, xdata, ydata, p0, maxfev=10000) # Generate more points for a smoother fitted curve xModel = np.linspace(xdata.min(), xdata.max(), 100) yModel = tanh(xModel, *popt) def Plot(): fig = plt.figure(figsize=(16,9)) ax = fig.add_subplot(111) ax.plot(xdata, ydata, 'D', label="data") # Display fitted parameters in the legend ax.plot(xModel, yModel, label=f"fit: a={popt[0]:.2f}, b={popt[1]:.4f}, c={popt[2]:.2f}, d={popt[3]:.2f}") ax.legend() plt.xlabel('X') plt.ylabel('Y') plt.show() Plot()
4. Additional Tips for Better Fitting
- Convert data to numpy arrays: This makes calculations faster and avoids potential issues with list operations.
- Try different optimization methods: If 'lm' fails, you can go back to 'dogbox' or try 'trf'—but good initial guesses usually make 'lm' work best.
- Normalize data (optional): If your x/y values are large, scaling them to the [-1,1] range can help the optimizer converge faster. For example:
Just remember to reverse the normalization for your final fit curve.x_norm = (xdata - xdata.mean()) / xdata.std() y_norm = (ydata - ydata.mean()) / ydata.std()
When you run this updated code, you'll see the fitted tanh curve closely matches your data points!
内容的提问来源于stack exchange,提问作者Gopala

