机器学习 LightGBM模型的SHAP和LIME --python实战(1)
·
一、 导入库与环境设置
首先是软件的安装
conda install numpy pandas matplotlib seaborn scikit-learn lightgbm
# 使用pip安装conda没有的库
pip install shap lime
首先是导入python的库
# 0. 导入库与环境设置
# ===================================================================
import numpy as np # 导入NumPy库,用于高效的数值计算,特别是数组操作。通常简写为np。
import pandas as pd # 导入Pandas库,用于数据处理和分析,核心数据结构是DataFrame。通常简写为pd。
import matplotlib.pyplot as plt # 导入Matplotlib的pyplot模块,用于数据可视化和绘图。通常简写为plt。
import seaborn as sns # 导入Seaborn库,它基于Matplotlib,提供了更高级、更美观的统计图形。通常简写为sns。
import shap # 导入SHAP库,用于模型解释。
import lime # 导入LIME库,是另一种用于模型解释的库。
import lime.lime_tabular # 从LIME库中导入专门用于处理表格数据的模块。
import warnings # 导入warnings模块,用于控制警告信息的输出。
import lightgbm as lgb # 导入LightGBM库,一个高性能的梯度提升框架。通常简写为lgb。
from sklearn.model_selection import train_test_split, GridSearchCV, KFold # 从scikit-learn导入模型选择相关的工具:数据分割、网格搜索和K折交叉验证。
from sklearn.linear_model import LinearRegression # 从scikit-learn导入线性回归模型。
from sklearn.tree import DecisionTreeRegressor # 从scikit-learn导入决策树回归模型。
from sklearn.ensemble import RandomForestRegressor, BaggingRegressor # 从scikit-learn导入集成学习模型:随机森林和Bagging。
from sklearn.metrics import mean_absolute_error, mean_squared_error, r2_score # 从scikit-learn导入用于评估回归模型性能的指标。
from sklearn.inspection import PartialDependenceDisplay # 从scikit-learn导入用于绘制部分依赖图的工具。
from sklearn.preprocessing import StandardScaler # 从scikit-learn导入用于数据标准化的工具。
# --- Matplotlib 和 Seaborn 美化设置 ---
plt.rcParams['font.family'] = 'sans-serif' # 设置Matplotlib的默认字体族为无衬线字体,以避免某些环境下中文或特殊符号显示为方框。
plt.style.use('seaborn-v0_8-whitegrid') # 应用Seaborn的'whitegrid'样式,使图表背景为白色网格,更加美观。
warnings.filterwarnings("ignore") # 忽略所有警告信息。在最终报告中可以这样做,但在开发阶段建议查看警告。
二、 数据加载与初步探查
# ===================================================================
# 1. 数据加载与初步探查
# ===================================================================
print("--- 1. 数据加载与初步探查 ---") # 打印阶段标题。
data = pd.read_csv('random_ml_dataset.csv') # 使用pandas的read_csv函数加载名为'数据15.3.csv'的CSV文件到DataFrame `data` 中。
pd.set_option('display.max_columns', None) # 设置pandas的显示选项,确保在打印DataFrame时能显示所有列,而不是用省略号代替。
print("数据信息 (data.info()):") # 打印提示信息。
data.info() # 调用.info()方法,打印DataFrame的简要信息,包括索引类型、列名、非空值数量和数据类型。
print("\n描述性统计 (data.describe()):") # 打印提示信息。
print(data.describe()) # 调用.describe()方法,为数值类型的列生成描述性统计数据,如计数、均值、标准差、最小值、四分位数和最大值。
print("\n数据前5行 (data.head()):") # 打印提示信息。
print(data.head()) # 调用.head()方法,显示DataFrame的前5行,以便直观地查看数据内容和格式。
导入 数据以及查看,然后下面主要就是看一下数据的分布
# ===================================================================
# 2. 探索性数据分析 (EDA)
# ===================================================================
print("\n--- 2. 探索性数据分析 (EDA) ---") # 打印阶段标题。
# --- 2.1 目标变量分布 ---
plt.figure(figsize=(12, 5)) # 创建一个尺寸为12x5英寸的图形窗口。
plt.subplot(1, 2, 1) # 在1行2列的网格中,选择第1个位置创建子图。
sns.histplot(data.iloc[:, 0], kde=True, bins=30) # 使用Seaborn绘制目标变量(第一列)的直方图和核密度估计曲线。
plt.title(f'Distribution of Target Variable "{data.columns[0]}"') # 设置子图的标题,动态显示目标变量的名称。
plt.xlabel("Value") # 设置X轴标签。
plt.ylabel("Frequency") # 设置Y轴标签。
plt.subplot(1, 2, 2) # 在1行2列的网格中,选择第2个位置创建子图。
sns.boxplot(y=data.iloc[:, 0]) # 使用Seaborn绘制目标变量(第一列)的箱线图。
plt.title(f'Boxplot of Target Variable "{data.columns[0]}"') # 设置子图的标题。
plt.tight_layout() # 自动调整子图参数,使其填充整个图像区域,避免重叠。
plt.savefig('EDA_目标变量分布.png') # 将生成的图像保存为文件。
plt.show() # 显示图像。
# --- 2.2 特征与目标变量的相关性热力图 ---
plt.figure(figsize=(12, 10)) # 创建一个尺寸为12x10英寸的图形窗口。
correlation_matrix = data.corr() # 计算DataFrame中所有数值列之间的皮尔逊相关系数矩阵。
sns.heatmap(correlation_matrix, annot=True, cmap='coolwarm', fmt='.2f', linewidths=.5) # 使用Seaborn绘制热力图来可视化相关系数矩阵。
# annot=True: 在格子上显示数值。cmap='coolwarm': 设置颜色映射。fmt='.2f': 数值格式化为两位小数。
plt.title('Correlation Heatmap of Features and Target') # 设置图像标题。
plt.savefig('EDA_相关性热力图.png') # 保存图像。
plt.show() # 显示图像。
# --- 2.3 查看与目标变量最相关的几个特征的散点图 ---
target_corr = correlation_matrix.iloc[1:, 0].abs().sort_values(ascending=False) # 提取除目标变量自身外的所有特征与目标变量的相关系数,取绝对值并降序排列。
top_features = target_corr.head(3).index # 获取相关性最高的3个特征的名称。
print(f"\n与目标变量最相关的3个特征是: {list(top_features)}") # 打印这3个特征的名称。
plt.figure(figsize=(15, 5)) # 创建一个尺寸为15x5英寸的图形窗口。
for i, feature inenumerate(top_features): # 遍历这3个特征。
plt.subplot(1, 3, i + 1) # 在1行3列的网格中,为每个特征创建一个子图。
sns.scatterplot(x=data[feature], y=data.iloc[:, 0]) # 绘制当前特征与目标变量的散点图。
plt.title(f'"{feature}" vs "{data.columns[0]}"') # 设置子图标题。
plt.xlabel(feature) # 设置X轴标签为当前特征名。
plt.ylabel(data.columns[0]) # 设置Y轴标签为目标变量名。
plt.tight_layout() # 自动调整布局。
plt.savefig('EDA_高相关特征散点图.png') # 保存图像。
plt.show() # 显示图像。
三、数据划分
这边是将将30%的数据划为测试集,70%为训练集。大家可以自己调整的,一般训练集是0.6-0.8,具体的可以查看数据以及相关的文章来确定。
# ===================================================================
# 3. 数据准备
# ===================================================================
print("\n--- 3. 数据准备 ---") # 打印阶段标题。
X = data.iloc[:, :-1] # 将除最后一列外的所有列作为特征变量X。
y = data.iloc[:, -1] # 将最后一列作为目标变量y。
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3, random_state=42) # 使用train_test_split函数将数据划分为训练集和测试集。
# test_size=0.3: 将30%的数据划为测试集,70%为训练集。
# random_state=42: 设置随机种子,确保每次划分结果都一样,便于复现实验。
print(f"训练集大小: {X_train.shape}, 测试集大小: {X_test.shape}") # 打印训练集和测试集的形状,确认划分成功。
这边是将将30%的数据划为测试集,70%为训练集。大家可以自己调整的,一般训练集是0.6-0.8,具体的可以查看数据以及相关的文章来确定。
四、基线模型比较
在开始机器学习之前,要有一个性能基准(Baseline),任何更复杂的模型都应该要超越这个基准才有价值。大家可以用比较自己需要的几个机器学习方法即可
# ===================================================================
# 4. 基线模型比较(包含LightGBM)
# ===================================================================
print("\n--- 4. 基线模型比较 ---") # 打印阶段标题。
models = { # 创建一个字典,键是模型名称,值是模型对象实例。
"Linear Regression": LinearRegression(), # 线性回归模型。
"Decision Tree": DecisionTreeRegressor(random_state=42), # 决策树回归模型。
"Random Forest": RandomForestRegressor(random_state=42, n_estimators=100), # 随机森林回归模型。
"LightGBM (Baseline)": lgb.LGBMRegressor(random_state=42, verbose=-1) # 使用默认参数的LightGBM回归模型。
}
results = {} # 创建一个空字典,用于存储每个模型的评估结果。
for name, model in models.items(): # 遍历models字典中的每一个模型。
model.fit(X_train, y_train) # 使用训练数据(X_train, y_train)来训练当前模型。
y_pred = model.predict(X_test) # 使用训练好的模型对测试集X_test进行预测。
r2 = r2_score(y_test, y_pred) # 计算真实值y_test和预测值y_pred之间的R²分数。
mse = mean_squared_error(y_test, y_pred) # 计算均方误差(MSE)。
results[name] = {'R²': r2, 'MSE': mse} # 将模型的R²和MSE结果存入results字典。
print(f"{name} -> R²: {r2:.4f}, MSE: {mse:.4f}") # 打印当前模型的性能指标。
results_df = pd.DataFrame(results).T # 将results字典转换为DataFrame,并转置(.T)使模型名称成为行索引。
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(15, 6)) # 创建一个1行2列的子图网格。
results_df['R²'].plot(kind='bar', ax=ax1, title='R² Score Comparison of Different Models') # 在第一个子图上绘制R²分数的条形图。
ax1.set_ylabel('R² Score') # 设置y轴标签。
ax1.tick_params(axis='x', rotation=45) # 将x轴的标签旋转45度,防止重叠。
results_df['MSE'].plot(kind='bar', ax=ax2, title='MSE Comparison of Different Models', color='orange') # 在第二个子图上绘制MSE的条形图。
ax2.set_ylabel('MSE') # 设置y轴标签。
ax2.tick_params(axis='x', rotation=45) # 旋转x轴标签。
plt.tight_layout() # 自动调整布局。
plt.savefig('模型比较_性能对比.png') # 保存图像。
plt.show() # 显示图像。
五、LightGBM超参数调优
一般机器学习的默认参数不一定最适合我们的特定数据集。因此,这一阶段的目标是通过网格搜索交叉验证(Grid Search CV) 来找到 LightGBM 的最佳超参数组合。
# ===================================================================
# 5. LightGBM超参数调优
# ===================================================================
print("\n--- 5. LightGBM超参数调优 ---") # 打印阶段标题。
# LightGBM参数网格
param_grid = { # 定义一个字典,作为超参数的搜索空间。
'n_estimators': [100, 200], # 树的数量(迭代次数)。
'max_depth': [-1, 10], # 每棵树的最大深度,-1表示不限制。
'num_leaves': [31, 50], # 每棵树的叶子节点数。
'learning_rate': [0.05, 0.1], # 学习率。
'feature_fraction': [0.8, 0.9], # 每次迭代时随机选择特征的比例。
'bagging_fraction': [0.8, 0.9] # 每次迭代时随机选择数据的比例(不进行重采样)。
}
# 初始化LightGBM模型
lgb_model = lgb.LGBMRegressor(random_state=42, verbose=-1) # 创建一个LightGBM回归器实例。
# 网格搜索
grid_search = GridSearchCV( # 创建一个GridSearchCV对象,用于执行网格搜索和交叉验证。
estimator=lgb_model, # 要调优的模型。
param_grid=param_grid, # 要搜索的参数网格。
cv=5, # 交叉验证的折数,这里是5折。
n_jobs=-1, # 使用所有可用的CPU核心进行并行计算,-1表示使用所有核心。
verbose=1, # 打印详细的运行日志,1表示打印简要日志。
scoring='r2'# 用于评估模型性能的指标,这里选择R²。
)
print("开始LightGBM超参数调优...") # 打印提示信息。
grid_search.fit(X_train, y_train) # 在训练集上执行网格搜索。这会自动进行交叉验证。
# 获取最佳模型
final_model = grid_search.best_estimator_ # 从grid_search结果中获取性能最好的模型实例。
print(f"\n最佳参数组合: {grid_search.best_params_}") # 打印找到的最佳参数组合。
print(f"交叉验证最佳R²得分: {grid_search.best_score_:.4f}") # 打印在交叉验证中得到的最高R²分数。
更多推荐
所有评论(0)