shap.DeepExplainer调用子类化Keras模型报NoneType无len()错误求助
问题描述
基于TensorFlow Recommenders与Keras构建多任务推荐模型,计划使用SHAP库开展模型可解释性分析,因采用子类化Keras模型实现,运行SHAP相关代码时触发类型错误。
核心调用代码如下:
import shap background = train_np[np.random.choice(train_np.shape[0], 100, replace=False)] explainer = shap.DeepExplainer(model, background)
运行时触发报错如下:
/usr/local/lib/python3.7/dist-packages/shap/explainers/tf_utils.py in _get_model_output(model) 83 isinstance(model, tf.keras.Model): 84 if len(model.layers[-1]._inbound_nodes) == 0: ---> 85 if len(model.outputs) > 1: 86 warnings.warn("Only one model output supported.") 87 return model.outputs[0] TypeError: object of type 'NoneType' has no len()
诱发原因
- 直接触发报错的原因是SHAP的
DeepExplainer在解析Keras模型时,会直接读取model.outputs属性获取模型输出张量,但子类化Keras模型在未经过显式前向传播、未调用build()方法指定输入形状的前提下,model.outputs属性默认返回None,对None调用len()就会触发上述类型错误。 - 深层兼容问题是TensorFlow Recommenders框架下的多任务推荐模型属于典型的动态式子类化模型,和Sequential、Functional API实现的静态图模型不同,不会在模型定义阶段就静态生成计算图、绑定输入输出节点,SHAP旧版本内置的模型解析逻辑没有适配这类动态模型的特性,直接读取静态属性就会出现兼容问题。
- 额外注意:
DeepExplainer原生只支持单输出模型解释,即使解决了outputs为None的问题,多输出模型如果不做输出裁剪也会触发警告或报错。
可行解决方案
- 方案1:显式触发模型构建,绑定输入输出节点
在初始化SHAP解释器之前,先传入符合输入维度的样例数据执行一次前向传播,让模型实例自动构建计算图、生成静态的输入输出属性:import tensorflow as tf # 传入1条符合输入维度的样例数据,触发模型构建 sample_batch = tf.convert_to_tensor(background[:1]) _ = model(sample_batch) # 如果是多输出模型,可自定义包装层裁剪到需要解释的单路输出后再传入SHAP explainer = shap.DeepExplainer(model, background) - 方案2:包装为Functional API静态模型
把需要解释的推理逻辑单独封装,基于明确的输入张量用Functional API格式构建静态图模型,加载训练好的权重后再传入SHAP,完全规避子类化模型的静态属性读取问题:# 定义和输入特征维度匹配的输入层 input_layer = tf.keras.Input(shape=(feature_dim,)) # 调用训练好的模型对应模块计算目标输出,比如排序任务的输出 target_output = model.ranking_task_tower(input_layer) # 构建静态Functional模型 wrapped_model = tf.keras.Model(inputs=input_layer, outputs=target_output) # 用包装后的静态模型初始化解释器 explainer = shap.DeepExplainer(wrapped_model, background) - 方案3:更换适配性更强的SHAP解释器
如果上述静态图适配方案仍然存在兼容问题,可以放弃DeepExplainer,改用不需要解析模型内部结构的KernelExplainer或GradientExplainer,仅把模型作为可调用预测函数传入即可,对自定义子类模型的适配性更好:# 封装预测函数,返回需要解释的目标输出值 def pred_func(input_data): return model(tf.convert_to_tensor(input_data)).ranking_score.numpy() # 初始化解释器 explainer = shap.KernelExplainer(pred_func, background)
内容的提问来源于stack exchange,提问作者rakesh
相关产品推荐
相关产品推荐

