使用BertModelLayer做三分类时报InputSpec类型错误如何解决?
错误根因
该报错是bert-for-tf2库(BertModelLayer所属的依赖包)与当前环境的TensorFlow/Keras版本不兼容导致的:高版本TensorFlow(2.4+)的Keras模块调整了InputSpec的类型校验逻辑,旧版bert-for-tf2生成的InputSpec实例无法通过校验。
排查步骤
- 执行命令查看当前环境依赖版本:
pip show tensorflow bert-for-tf2
如果TensorFlow版本≥2.4,且bert-for-tf2版本≤0.14.6,即可确定是版本兼容问题。
解决方案
按优先级从高到低选择即可:
- 升级bert-for-tf2到最新兼容版本
执行升级命令,新版本已经适配了高版本TensorFlow的InputSpec校验规则:
pip install --upgrade bert-for-tf2
升级完成后重新运行代码即可。
2. 降级TensorFlow到兼容版本
如果无法升级bert-for-tf2,可以将TensorFlow降级到2.3.x版本,该版本的Keras InputSpec逻辑与旧版bert-for-tf2完全兼容:
pip install tensorflow==2.3.0
- 临时手动修复(无需修改依赖版本)
找到报错的文件/usr/local/lib/python3.7/dist-packages/bert/transformer.py,修改第209行代码:- 首先在文件头部添加导入语句:
from tensorflow.keras.layers import InputSpec - 把原有行
self.input_spec = keras.layers.InputSpec(shape=input_shape)替换为self.input_spec = InputSpec(shape=input_shape)
- 首先在文件头部添加导入语句:
额外注意
你提供的代码存在变量未定义问题:create_model函数内部直接使用了bert_config_file,但该变量没有作为参数传入,解决版本问题后还需要把函数定义修改为def create_model(max_seq_len, bert_ckpt_file, bert_config_file):,调用时传入BERT配置文件的路径即可。
内容的提问来源于stack exchange,提问作者Hrisav Bhowmick
相关产品推荐
相关产品推荐

