如何为每个Serial number批量拟合Logistic曲线并输出参数至DataFrame?
Hey there! I see you're trying to fit logistic curves to each group in your large DataFrame and collect the parameters neatly. Let's tweak your code to get that done properly. Here's a step-by-step solution:
Step 1: Fix Data Loading & Prep
First, your original code converts the Serial number to float, which isn't ideal—serial numbers are usually identifiers, not numeric values. Let's adjust the loading step to keep it as an integer (or string if needed):
import pandas as pd import numpy as np from scipy.optimize import curve_fit # Load data correctly, keep Serial number as integer ohlalala = pd.read_csv('1851_2019.csv', sep=';') # Select relevant columns, don't cast Serial number to float ohlala = ohlalala[['Serial number', 'mrwSmpVWi', 'mrwSmpP']].copy() # Ensure the numeric columns are float (in case they loaded as strings) ohlala[['mrwSmpVWi', 'mrwSmpP']] = ohlala[['mrwSmpVWi', 'mrwSmpP']].astype(float)
Step 2: Define the Logistic Function Outside the Loop
No need to redefine the function every time you loop through a group—define it once at the top for efficiency:
def logifunc(x, c, a, b): return c / (1 + a * np.exp(-b * x))
Step 3: Collect Parameters for Each Group
Create an empty list to store results, then loop through each group, fit the curve, and append the parameters to the list. We'll add a try-except block to handle any fitting failures (since some groups might have data that can't be fitted):
# Empty list to hold results param_results = [] grouped_data = ohlala.groupby('Serial number') for group_name, grouped_df in grouped_data: # Extract x and y values clearly (no need for transpose!) x = grouped_df['mrwSmpVWi'].values y = grouped_df['mrwSmpP'].values try: # Fit the logistic curve with initial guesses result, pcov = curve_fit(logifunc, x, y, p0=[110, 400, -2]) # Append results as a dictionary (easy to convert to DataFrame later) param_results.append({ 'Serial number': group_name, 'c': result[0], 'a': result[1], 'b': result[2] }) except RuntimeError as e: # Handle fitting failures (log the error if needed) print(f"Failed to fit curve for Serial number {group_name}: {e}") # You can choose to append NaN values for failed fits param_results.append({ 'Serial number': group_name, 'c': np.nan, 'a': np.nan, 'b': np.nan })
Step 4: Convert Results to DataFrame
Finally, turn the list of dictionaries into a DataFrame where each row represents one Serial number and its fitted parameters:
params_df = pd.DataFrame(param_results) print(params_df.head())
Key Notes:
- Avoiding Transpose: Using
grouped_df['column_name'].valuesis much more readable than transposing the DataFrame to extract rows. - Error Handling: The try-except block prevents your code from crashing if a group's data can't be fitted (e.g., not enough data points, poor initial guesses).
- Efficiency: Defining the function outside the loop saves unnecessary overhead, which is helpful for your large 250k-row dataset.
内容的提问来源于stack exchange,提问作者1lk4

