TensorFlow 2.1中keras.utils.layer_utils与generic_utils的tf.keras等价导入问题
独立Keras迁移到tf.keras(TensorFlow 2.1.0)工具模块导入解决方案
两个缺失模块的正确导入路径
- layer_utils模块
TensorFlow 2.1.0版本中对应导入路径为:
from tensorflow.python.keras.utils import layer_utils
如果你仅用到该模块下的特定工具函数,比如参数计数、层信息获取等,也可以直接调用对应层的内置方法,无需导入整个模块:
# 替代layer_utils.count_params()的写法 param_count = your_layer.count_params()
- generic_utils模块
对应导入路径为:
from tensorflow.python.keras.utils import generic_utils
该模块最常用的进度条工具Progbar也可以直接从公开接口导入:
from tensorflow.keras.utils import Progbar
其他可选替代方案
如果不想依赖tensorflow.python的内部路径(内部模块路径可能在后续TensorFlow版本中发生变动),可以根据你实际用到的功能替换为公开接口或第三方工具:
- 若用到generic_utils的进度条功能:可以直接用第三方库
tqdm替代,使用更灵活,适配性更强 - 若用到layer_utils的权重处理相关功能:直接调用tf.keras.layers.Layer的内置方法,或者使用tf.keras.initializers、tf.keras.regularizers下的公开接口实现对应逻辑
原有导入语句完整适配参考(TensorFlow 2.1.0)
import tensorflow as tf from tensorflow.keras import backend as K from tensorflow.keras.optimizers import Adam, SGD, RMSprop from tensorflow.keras.layers import Flatten, Dense, Input, Conv2D, MaxPooling2D, Dropout from tensorflow.keras.layers import GlobalAveragePooling2D, GlobalMaxPooling2D, TimeDistributed from tensorflow.keras.utils import get_source_inputs from tensorflow.python.keras.utils import layer_utils from tensorflow.keras.utils import get_file from tensorflow.keras.losses import categorical_crossentropy from tensorflow.keras.models import Model from tensorflow.python.keras.utils import generic_utils from tensorflow.keras.layers import Layer, InputSpec from tensorflow.keras import initializers, regularizers
内容的提问来源于stack exchange,提问作者Farrel Ferdian
相关产品推荐
相关产品推荐

