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

调用TFBertModel时KerasTensor类型不兼容错误求助

解决TFBertModel与KerasTensor输入不兼容问题

问题核心

使用TFBertModel构建Keras分类模型时,bert_model([input_ids, attention_mask])报错,提示不接受KerasTensor类型,仅支持Tensor、numpy.ndarray等。

解决方案

1. 改用字典形式传递输入参数

TFBertModel官方推荐通过字典传递命名输入,而非列表。修改报错行代码:

bert_output = bert_model({'input_ids': input_ids, 'attention_mask': attention_mask})

2. 封装BERT模型为自定义Keras层

通过自定义层包装TFBertModel,让Keras正确识别输入输出的张量类型:

from transformers import TFBertModel
import tensorflow as tf
from tensorflow.keras import Model
from tensorflow.keras import layers

# 自定义BERT包装层
class BertWrapper(tf.keras.layers.Layer):
    def __init__(self, bert_model, **kwargs):
        super().__init__(**kwargs)
        self.bert_model = bert_model
    
    def call(self, inputs):
        input_ids, attention_mask = inputs
        return self.bert_model({'input_ids': input_ids, 'attention_mask': attention_mask})

# 加载BERT模型并启用返回字典配置
bert_model = TFBertModel.from_pretrained('bert-base-uncased', return_dict=True)

# 定义输入层
input_ids = tf.keras.layers.Input(shape=(128,), dtype='int32', name='input_ids')
attention_mask = tf.keras.layers.Input(shape=(128,), dtype='int32', name='attention_mask')

# 使用包装层调用BERT
bert_wrapper = BertWrapper(bert_model)
bert_output = bert_wrapper([input_ids, attention_mask])

pooled_output = bert_output.pooler_output

# 后续分类层逻辑保持不变
x = layers.Dense(128, activation='relu')(pooled_output)
x = layers.Dropout(0.3)(x)
output = layers.Dense(2, activation='softmax')(x)

model = Model(inputs=[input_ids, attention_mask], outputs=output)
model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=5e-5),
              loss='sparse_categorical_crossentropy',
              metrics=['accuracy'])
model.fit(train_dataset, epochs=3, validation_data=test_dataset)

3. 调整版本组合适配Python 3.12

Python 3.12属于较新版本,部分旧版Transformers对其支持不足:

  • 升级Transformers至4.48.0及以上,配合当前TensorFlow 2.18.0使用
  • 若仍有问题,可临时降级Python至3.11,该版本对TensorFlow和Transformers的兼容性更成熟

4. 加载模型时显式启用返回字典配置

加载TFBertModel时设置return_dict=True,确保输出结构符合Keras预期:

bert_model = TFBertModel.from_pretrained('bert-base-uncased', return_dict=True)

内容的提问来源于stack exchange,提问作者Faiz khan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 08:52:08