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

LSTM模型获取特征重要性时出现TypeError报错求助

问题:LSTM模型无法获取特征重要性,报TypeError错误

我已构建好LSTM模型,但无法获取目标变量y的特征重要性,尝试更换环境后仍出现如下TypeError报错:

import os
os.environ["SM_FRAMEWORK"] = "tf.keras"

lstm_classifier = LSTMClassifier(Cluster_0)

lstm_classifier.preprocess_data()

X_test = lstm_classifier.get_X_test()
y_test = lstm_classifier.get_y_test()

feature_importances = lstm_classifier.get_feature_importance(X_test, y_test)

报错信息:

TypeError                                 Traceback (most recent call last)
<ipython-input-61-046c9c5bbede> in <cell line: 12>()
     10 
     11 
---> 12 feature_importances = lstm_classifier.get_feature_importance(X_test, y_test)

2 frames
<ipython-input-57-8bc9dce06edb> in get_feature_importance(self, X, y, n_repeats, random_state)
    150 
    151         # Compute the permutation importance
---> 152         result = permutation_importance(predict_proba_wrapped, X, y, n_repeats=n_repeats,
    153                                         random_state=random_state, n_jobs=-1)
    154 

/usr/local/lib/python3.10/dist-packages/sklearn/inspection/_permutation_importance.py in permutation_importance(estimator, X, y, scoring, n_repeats, n_jobs, random_state, sample_weight, max_samples)
    249         scorer = scoring
    250     elif scoring is None or isinstance(scoring, str):
---> 251         scorer = check_scoring(estimator, scoring=scoring)
    252     else:
    253         scorers_dict = _check_multimetric_scoring(estimator, scoring)

/usr/local/lib/python3.10/dist-packages/sklearn/metrics/_scorer.py in check_scoring(estimator, scoring, allow_none)
    472     """
    473     if not hasattr(estimator, "fit"):
---> 474         raise TypeError(
    475             "estimator should be an estimator implementing 'fit' method, %r was passed"
    476             % estimator

TypeError: estimator should be an estimator implementing 'fit' method, <function LSTMClassifier.get_feature_importance.<locals>.predict_proba_wrapped at 0x79fb6f5d7eb0> was passed

报错原因

permutation_importance函数的第一个参数要求是实现了fit方法的sklearn兼容估算器,但你传入的是自定义的predict_proba_wrapped函数,该函数没有fit方法,因此触发TypeError。

解决方法

方法1:让模型类兼容sklearn接口

修改LSTMClassifier类,确保它实现sklearn风格的fit和predict_proba(或predict)方法,然后在调用permutation_importance时传入模型实例本身:

# 修改get_feature_importance方法中的调用
result = permutation_importance(self, X, y, n_repeats=n_repeats,
                                random_state=random_state, n_jobs=-1, scoring='accuracy')

这里的self就是你的lstm_classifier实例,只要它符合sklearn估算器的接口规范,就能被permutation_importance正确识别。

方法2:手动传入自定义评分器

如果不想修改模型类,可以手动创建一个评分器,跳过check_scoring对估算器fit方法的检查:

from sklearn.metrics import accuracy_score, make_scorer

# 自定义评分函数,适配你的模型输出格式
def custom_scorer(y_true, y_pred_proba):
    # 假设y_pred_proba是概率矩阵,取最大值对应的类别作为预测结果
    y_pred = y_pred_proba.argmax(axis=1)
    return accuracy_score(y_true, y_pred)

# 创建评分器
scorer = make_scorer(custom_scorer)

# 调用permutation_importance时传入该评分器
result = permutation_importance(predict_proba_wrapped, X, y, n_repeats=n_repeats,
                                random_state=random_state, n_jobs=-1, scoring=scorer)

方法3:用KerasWrapper包装Keras模型

如果你的LSTM是用Keras构建的,直接使用sklearn.wrappers.KerasClassifier(分类任务)或KerasRegressor(回归任务)包装模型,使其兼容sklearn接口:

from sklearn.wrappers import KerasClassifier
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import LSTM, Dense

# 定义模型构建函数
def build_lstm_model(input_shape):
    model = Sequential()
    model.add(LSTM(64, input_shape=input_shape))
    model.add(Dense(2, activation='softmax'))
    model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy'])
    return model

# 包装模型
lstm_classifier = KerasClassifier(build_fn=lambda: build_lstm_model((X_train.shape[1], X_train.shape[2])), 
                                  epochs=10, batch_size=32)
lstm_classifier.fit(X_train, y_train)

# 计算特征重要性
result = permutation_importance(lstm_classifier, X_test, y_test, n_repeats=5, random_state=42)

内容的提问来源于stack exchange,提问作者lily li

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 14:55:56