如何微调TensorFlow Hub模型?无法访问base_model.trainable及layers属性
解决TensorFlow Hub模型微调时无法访问trainable/layers属性的问题
问题根源
从TensorFlow Hub加载的预训练模型多为封装后的SavedModel,通过hub.KerasLayer加载时,外层的base_model是一个包装层,而非原生Keras模型结构,因此无法直接访问trainable或layers属性。
解决步骤
提取原生预训练模型
通过KerasLayer的内部层级获取真正的预训练模型主体:import tensorflow_hub as hub import tensorflow as tf # 原模型加载代码 base_model = hub.KerasLayer("你的Hub模型地址") # 提取内部原生模型 actual_base_model = base_model.layers[0]若上述方式无效,可打印
base_model.summary()查看结构,定位到包含预训练权重的内部层节点。开启并控制微调范围
获取原生模型后,即可正常设置训练属性,还可选择性冻结部分层:# 开启模型可训练 actual_base_model.trainable = True # 冻结前N层,仅训练后续层 fine_tune_at = 100 for layer in actual_base_model.layers[:fine_tune_at]: layer.trainable = False重新编译模型
微调阶段需使用极小的学习率,避免破坏预训练权重:model.compile(optimizer=tf.keras.optimizers.Adam(learning_rate=1e-5), loss=tf.keras.losses.BinaryCrossentropy(from_logits=True), metrics=['accuracy'])
内容的提问来源于stack exchange,提问作者Osama Mohammed
相关产品推荐
相关产品推荐

