使用simpletransformers训练Seq2Seq模型时遇tensorflow无io属性错误
问题:实例化Seq2SeqModel时触发AttributeError
我尝试用simpletransformers库训练Seq2Seq模型,实例化Seq2SeqModel时一直报以下错误:
复现代码
import tensorflow as tf from simpletransformers.seq2seq import Seq2SeqModel, Seq2SeqArgs import logging model_args = { "reprocess_input_data": True, "overwrite_output_dir": True, "max_seq_length": 50, "train_batch_size": 16, "num_train_epochs": 3, "save_eval_checkpoints": False, "save_model_every_epoch": False, "evaluate_generated_text": True, "evaluate_during_training_verbose": True, "use_multiprocessing": False, "max_length": 50, "manual_seed": 42, } logging.basicConfig(level=logging.INFO) transformers_logger = logging.getLogger("transformers") transformers_logger.setLevel(logging.WARNING) model = Seq2SeqModel( "roberta", "roberta-base", "bert-base-uncased", from_tf=True, args=model_args, use_cuda=False )
错误信息
AttributeError: module 'tensorflow' has no attribute 'io'
当前环境
- tensorflow版本:2.11.0
- Python版本:3.8.3
解决办法
- 排查模块覆盖问题:在报错代码前添加
print(dir(tf)),查看输出中是否包含io。如果没有,说明你的tf模块被同名文件或其他库篡改,检查项目目录下是否有tensorflow.py这类文件,删除或重命名它。 - 降级TensorFlow版本:tf 2.11.0与simpletransformers/transformers的兼容性不佳,建议降级到2.10.x版本,执行命令:
pip install tensorflow==2.10.0 - 更新依赖库:simpletransformers依赖的transformers版本过旧可能导致该问题,执行命令更新:
pip install --upgrade transformers simpletransformers - 移除
from_tf=True参数:你使用的roberta-base是PyTorch格式的预训练权重,强制指定from_tf=True会触发格式转换冲突。去掉该参数,让模型自动加载PyTorch权重即可。
内容的提问来源于stack exchange,提问作者Mehul Gupta
相关产品推荐
相关产品推荐

