如何用Python绘制Linear Discriminant Analysis的变量分布图?
Hey there! Sounds like you're aiming to replicate that useful XLSTAT-style variable loading plot for your Linear Discriminant Analysis (LDA) in Python—those plots are perfect for understanding which features drive your discriminant axes, so let's break down how to build one step by step.
First, Understand What You're Targeting
XLSTAT's LDA variable plot is typically a loading plot: it shows the projection of your original features onto the LDA discriminant axes (LD1, LD2, etc.). Each point/arrow represents a feature, and its position tells you how strongly it contributes to separating your classes along that axis.
Step 1: Extract LDA Loadings (Variable Coefficients)
Assuming you're using scikit-learn (the most common Python library for LDA), the LinearDiscriminantAnalysis class has a scalings_ attribute that gives you exactly these loadings. This matches the variable coefficients XLSTAT outputs.
Here's how to extract them into a readable dataframe:
from sklearn.discriminant_analysis import LinearDiscriminantAnalysis import pandas as pd # Replace with your own feature matrix (X) and target labels (y) lda = LinearDiscriminantAnalysis() lda.fit(X, y) # Convert loadings to a dataframe for easier handling # Each column is a discriminant axis (LD1, LD2...) loadings = pd.DataFrame( lda.scalings_, index=X.columns, # Use your feature names here columns=[f'LD{i+1}' for i in range(lda.scalings_.shape[1])] )
Note: If XLSTAT was using standardized variables (a common setting), make sure to scale your features first with StandardScaler before fitting LDA to match the results.
Step 2: Plot the Loadings (XLSTAT-style)
We'll use matplotlib to build a plot that mirrors XLSTAT's look. You can choose between scatter points with labels (like XLSTAT's default) or arrows (another common XLSTAT style) to represent variables.
Option 1: Scatter Point Style (Matching XLSTAT's Basic Plot)
import matplotlib.pyplot as plt plt.figure(figsize=(8, 6)) # Plot each variable as a scatter point with its name for feature in loadings.index: ld1_val = loadings.loc[feature, 'LD1'] ld2_val = loadings.loc[feature, 'LD2'] # Plot the point plt.scatter(ld1_val, ld2_val, s=120, c='#007acc', edgecolor='white') # Add the feature name next to the point plt.text(ld1_val + 0.02, ld2_val + 0.02, feature, fontsize=10) # Add reference lines (XLSTAT includes these to center the plot) plt.axvline(0, color='gray', linestyle='--', alpha=0.7) plt.axhline(0, color='gray', linestyle='--', alpha=0.7) # Add labels and title (match XLSTAT's clean style) plt.title('LDA Variable Loading Plot', fontsize=14, pad=15) plt.xlabel('Linear Discriminant 1 (LD1)', fontsize=12) plt.ylabel('Linear Discriminant 2 (LD2)', fontsize=12) # Clean up the plot plt.grid(alpha=0.3) plt.tight_layout() plt.show()
Option 2: Arrow Style (XLSTAT's Directional Plot)
If you prefer arrows to show the direction and magnitude of each variable's contribution, use this instead:
plt.figure(figsize=(8, 6)) # Draw arrows for each feature for feature in loadings.index: ld1_val = loadings.loc[feature, 'LD1'] ld2_val = loadings.loc[feature, 'LD2'] # Draw arrow from origin to the loading point plt.arrow( 0, 0, ld1_val, ld2_val, head_width=0.03, head_length=0.03, color='#007acc', length_includes_head=True ) # Place feature name outside the arrow plt.text(ld1_val * 1.1, ld2_val * 1.1, feature, fontsize=10) # Add reference lines plt.axvline(0, color='gray', linestyle='--', alpha=0.7) plt.axhline(0, color='gray', linestyle='--', alpha=0.7) # Set axis limits to fit all arrows plt.xlim(-1.3, 1.3) plt.ylim(-1.3, 1.3) # Add labels and title plt.title('LDA Variable Direction Plot', fontsize=14, pad=15) plt.xlabel('LD1', fontsize=12) plt.ylabel('LD2', fontsize=12) plt.grid(alpha=0.3) plt.tight_layout() plt.show()
Quick Tips for Matching XLSTAT Exactly
- Variable Standardization: XLSTAT often standardizes variables by default. If your results don't match, add
StandardScaler()to your pipeline before LDA. - Color Schemes: Adjust the
cparameter inscatter()orarrow()to match XLSTAT's default colors (usually a blue palette). - Axis Scaling: Use
plt.xlim()/plt.ylim()to match the axis ranges XLSTAT uses for your data.
内容的提问来源于stack exchange,提问作者Ammar Sabir Cheema

