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

为何SHAP交互式摘要图在首个标签页中重复绘制?

解决SHAP绘图在IPython标签页重复显示的问题

问题原因

  1. Matplotlib绘图状态未重置:每次调用SHAP绘图函数后,未清理Matplotlib的Figure对象,导致新绘图叠加在旧图上。
  2. 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 09:54:56