TensorFlow 2运行nmt-chatbot时tf.contrib.training.HParams报错替代方案咨询
问题解决方法
你遇到的报错是TensorFlow 2.x版本正式移除了旧版1.x中的contrib模块导致的,原有tf.contrib.training.HParams可通过以下三个方案替代:
- 方案1:最小改动替换内置类
直接将原有代码中的tf.contrib.training.HParams替换为TF2内置的tf.keras.utils.HParams,其余代码不需要做任何调整,原有HParams的所有调用方法都可以正常兼容。
修改后代码片段如下:
def create_hparams(flags): """Create training hparams.""" return tf.keras.utils.HParams( # Data src=flags.src, tgt=flags.tgt, train_prefix=flags.train_prefix, dev_prefix=flags.dev_prefix, test_prefix=flags.test_prefix, vocab_prefix=flags.vocab_prefix, embed_prefix=flags.embed_prefix, out_dir=flags.out_dir, # 原有代码后面的其余参数直接原封不动保留即可
- 方案2:开启TF1兼容模式
如果后续还会遇到大量TF1/2不兼容的报错,可以直接在项目入口文件的最开头添加以下两行代码,直接启用TensorFlow的1.x兼容模式,不需要修改包括HParams在内的任何原有代码:
import tensorflow.compat.v1 as tf tf.disable_v2_behavior()
- 方案3:使用独立hparams库
如果不想依赖TensorFlow内置工具,可以单独安装官方拆分出的hparams工具包:pip install hparams
之后修改导入逻辑即可,用法和原有HParams基本一致。
内容的提问来源于stack exchange,提问作者Icarus
相关产品推荐
相关产品推荐

