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

Python手写逻辑回归实现中系数的绘制与排序方法求助

Hey there! I get that you're trying to visualize and sort your logistic regression coefficients to spot the least impactful features, and it’s a bummer that sns.coefplot() got deprecated. No worries though—we can easily replicate this functionality using seaborn.barplot() (or plain matplotlib) with a few straightforward steps. Let’s walk through how to integrate this into your existing code.

Step-by-Step Solution

1. Map Coefficients to Feature Names

First, we need to link your numerical theta values back to the actual feature names from the dataset. This makes the plot meaningful instead of just showing arbitrary indices.

2. Organize Data for Sorting

We’ll use a pandas DataFrame to pair each feature name with its coefficient. This simplifies sorting and plotting.

3. Sort by Impact Magnitude

Sorting coefficients by their absolute value is the most useful here—negative coefficients just mean the feature has an inverse relationship with the target, not a smaller impact.

4. Plot the Sorted Coefficients

A horizontal bar plot will make it easy to read feature names and compare coefficient sizes at a glance.

Modified Code for Your Logistic Regression Function

Replace the plotting section in your Logistic_Regression function with this code:

def Logistic_Regression(X,Y,alpha,theta,num_iters):
    m = len(Y)
    for x in range(num_iters):
        new_theta = Gradient_Descent(X,Y,theta,m,alpha)
        theta = new_theta
        if x % 100 == 0:
            Accuracy(theta)
    
    # ---------------------- Coefficient Visualization Code ----------------------
    # 1. Extract feature names (skip first column which is the intercept)
    feature_names = df.columns[2:-1].tolist()
    # Add intercept name to match theta's first element
    feature_names.insert(0, 'Intercept')
    
    # 2. Create DataFrame for coefficients
    coef_df = pd.DataFrame({
        'Feature': feature_names,
        'Coefficient': theta.flatten()  # Convert 2D theta array to 1D
    })
    
    # 3. Sort by absolute coefficient value (descending order)
    coef_df['Abs_Coefficient'] = coef_df['Coefficient'].abs()
    coef_df_sorted = coef_df.sort_values(by='Abs_Coefficient', ascending=False)
    
    # 4. Plot the sorted coefficients
    plt.figure(figsize=(12, 8))
    sns.barplot(x='Coefficient', y='Feature', data=coef_df_sorted, palette='viridis')
    plt.title('Sorted Logistic Regression Coefficients (By Impact Magnitude)')
    plt.xlabel('Coefficient Value')
    plt.ylabel('Feature')
    plt.grid(axis='x', linestyle='--', alpha=0.7)
    plt.show()
    
    # Optional: To view smallest impact features first, use this instead
    # coef_df_sorted_asc = coef_df.sort_values(by='Abs_Coefficient', ascending=True)
    # sns.barplot(x='Coefficient', y='Feature', data=coef_df_sorted_asc, palette='viridis')
    # plt.title('Coefficients (Smallest to Largest Impact)')
    # plt.show()

Key Notes

  • Intercept Handling: The first element of theta is the intercept term, so we add 'Intercept' to our feature list to keep everything aligned with your data processing.
  • Interpretability: Sorting by absolute value highlights which features drive predictions the most, regardless of whether they increase or decrease the likelihood of a malignant diagnosis.
  • Feature Removal: Once you identify low-impact features (smallest absolute coefficients), you can drop them from your training data and re-train the model to test if performance stays consistent. Just make sure to do this on the training set first (and use cross-validation!) to avoid overfitting.

Quick Side Tip

You’re currently applying both standardization ((X - mean)/std) and MinMax scaling. For logistic regression, standardization alone is usually preferred—it makes coefficients more interpretable (each unit change corresponds to a standard deviation shift in the feature).

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 08:20:44