为何SHAP交互式摘要图在首个标签页中重复绘制?
解决SHAP绘图在IPython标签页重复显示的问题
问题原因
- Matplotlib绘图状态未重置:每次调用SHAP绘图函数后,未清理Matplotlib的Figure对象,导致新绘图叠加在旧图上。
- Output上下文未完全隔离:IPywidgets的Output组件未在每次函数调用前清空,加上SHAP绘图的全局状态泄漏,导致图表跨标签页重复显示。
修复步骤
1. 管理Matplotlib绘图生命周期
修改SHAPInterpreter类的summary_plot方法,确保每次绘图创建独立Figure,并在显示后关闭,避免状态残留:
def summary_plot(self, max_display=10, feature_names=None, plot_type='dot', color_bar=False): """Generate a SHAP summary plot.""" # 新建独立Figure,避免与已有绘图冲突 plt.figure() if feature_names is not None: feature_indices = [self.X.columns.get_loc(name) for name in feature_names] shap_values = self.shap_values[:, feature_indices] X = self.X.iloc[:, feature_indices] else: shap_values = self.shap_values X = self.X shap.summary_plot( shap_values, X, max_display=max_display, plot_type=plot_type, color_bar=color_bar, show=False, plot_size=(10, 6)) plt.show() # 关闭当前Figure,释放资源 plt.close()
2. 强化Output上下文隔离
在interactive_summary_plot2的update_plot函数中,添加全局清理Matplotlib状态的代码,确保每次更新绘图前完全重置:
def update_plot(button): with out: clear_output(wait=True) # 关闭所有残留的Matplotlib Figure plt.close('all') num_features = slider.value plot_type = dropdown.value feature_names = [name for name, checkbox in checkboxes.items() if checkbox.value] self.summary_plot(max_display=num_features, feature_names=feature_names, plot_type=plot_type, color_bar=True)
3. 修复标签页装饰器的上下文清空
修改add_to_tab装饰器,确保每次调用函数前清空对应标签页的Output内容,避免累积:
def add_to_tab(tab, title): def decorator(func): def wrapper(*args, **kwargs): for i in range(len(tab.children)): if tab.get_title(i) == title: break else: tab.children += (widgets.Output(),) tab.set_title(len(tab.children) - 1, title) i = len(tab.children) - 1 with tab.children[i]: # 先清空标签页内容,再执行函数 clear_output(wait=True) func(*args, **kwargs) return wrapper return decorator
4. 处理Dependence Plot的状态泄漏
修改dependence_plot方法,同样添加绘图后的清理:
def dependence_plot(self, feature, interaction_index=None, show=True): """Generate a SHAP dependence plot.""" plt.figure() result = shap.dependence_plot(feature, self.shap_values, self.X, interaction_index=interaction_index, show=show) plt.close() return result
修改后的完整代码
import shap import numpy as np import matplotlib.pyplot as plt from sklearn.model_selection import train_test_split from IPython.display import clear_output import ipywidgets as widgets class SHAPInterpreter: """ A class that builds on top of the SHAP library to compute and plot SHAP values for a LightGBM model. """ def __init__(self, model, X, y, downsample=False, sample_frac=0.2, random_state=None): """ Initialize the SHAPInterpreter. Parameters: model (lightgbm.LGBMModel): A trained LightGBM model. X (pandas.DataFrame): The feature matrix. y (pandas.Series): The target vector. downsample (bool): Whether to downsample the data. sample_frac (float): The fraction of data to use for plotting. random_state (int): The random state to use for downsampling. """ self.model = model self.X = X self.y = y self.explainer = shap.TreeExplainer(model) if downsample: self.X, _, self.y, _ = train_test_split(self.X, self.y, test_size=sample_frac, stratify=self.y, random_state=random_state) self.shap_values = self.explainer.shap_values(self.X) self.feature_names = self.X.columns.tolist() def summary_plot(self, max_display=10, feature_names=None, plot_type='dot', color_bar=False): """Generate a SHAP summary plot.""" plt.figure() if feature_names is not None: feature_indices = [self.X.columns.get_loc(name) for name in feature_names] shap_values = self.shap_values[:, feature_indices] X = self.X.iloc[:, feature_indices] else: shap_values = self.shap_values X = self.X shap.summary_plot( shap_values, X, max_display=max_display, plot_type=plot_type, color_bar=color_bar, show=False, plot_size=(10, 6)) plt.show() plt.close() def dependence_plot(self, feature, interaction_index=None, show=True): """Generate a SHAP dependence plot.""" plt.figure() result = shap.dependence_plot(feature, self.shap_values, self.X, interaction_index=interaction_index, show=show) plt.close() return result def interactive_summary_plot2(self): """Create an interactive SHAP summary plot.""" slider = widgets.IntSlider( value=min(10, self.X.shape[1]), min=1, max=self.X.shape[1], step=1, description='Number of features:', ) dropdown = widgets.Dropdown( options=['dot', 'bar'], value='dot', description='Plot type:', ) checkboxes = {} checkboxes_box = widgets.VBox( layout=widgets.Layout(overflow_y='scroll')) button = widgets.Button(description='Update plot') out = widgets.Output() def update_checkboxes(change): num_features = change['new'] checkboxes.clear() mean_shap_values = np.abs(self.shap_values).mean(axis=0) sorted_feature_names = self.X.columns[np.argsort(mean_shap_values)[::-1]] checkboxes.update({col: widgets.Checkbox(value=(i < num_features), description=col) for i, col in enumerate(sorted_feature_names[:num_features])}) checkboxes_box.children = [ widgets.Label(value='Select features to display:'), widgets.VBox(list(checkboxes.values()), layout=widgets.Layout(overflow_y='scroll', height='150px', border='solid 1px'))] slider.observe(update_checkboxes, names='value') def update_plot(button): with out: clear_output(wait=True) plt.close('all') num_features = slider.value plot_type = dropdown.value feature_names = [name for name, checkbox in checkboxes.items() if checkbox.value] self.summary_plot(max_display=num_features, feature_names=feature_names, plot_type=plot_type, color_bar=True) button.on_click(update_plot) update_checkboxes({'new': slider.value}) display(slider, checkboxes_box, dropdown, button, out) # 假设modeler是已定义的模型对象 # shap_interpreter = SHAPInterpreter( # modeler.best_model, # modeler.test_set[0], # modeler.test_set[1], # downsample=True, # sample_frac=0.2, # random_state=139) def add_to_tab(tab, title): def decorator(func): def wrapper(*args, **kwargs): for i in range(len(tab.children)): if tab.get_title(i) == title: break else: tab.children += (widgets.Output(),) tab.set_title(len(tab.children) - 1, title) i = len(tab.children) - 1 with tab.children[i]: clear_output(wait=True) func(*args, **kwargs) return wrapper return decorator def run_functions_in_tabs(func_dict, tab=None): if tab is None: tab = widgets.Tab() display(tab) for title, func_info in func_dict.items(): func = func_info.get('func') args = func_info.get('args', []) kwargs = func_info.get('kwargs', {}) decorated_func = add_to_tab(tab, title)(func) decorated_func(*args, **kwargs) # func_dict = { # 'Summary Plot': {'func': shap_interpreter.interactive_summary_plot2}, # 'Dependence Plot': {'func': shap_interpreter.dependence_plot, 'args': ['DTB_cnt_8wk']}, # 'ROC Curve': {'func': modeler.plot_roc_curve} # } # run_functions_in_tabs(func_dict)
内容的提问来源于stack exchange,提问作者exan
相关产品推荐
相关产品推荐

