TensorFlow Estimator无法训练第二组全连接层问题咨询
嘿,我完全懂你碰到的这个坑——用TensorFlow Estimator做分阶段迁移学习,先训带自定义损失的CNN卷积层,再冻结卷积层换个损失训新的全连接层,结果第二步直接报错,原以为不用提前配置第二个损失对吧?
其实问题出在Estimator的核心机制上:Estimator实例一旦创建,它绑定的模型函数(model_fn)就固定了,包括损失函数、网络结构这些都没法动态切换。所以你不能直接复用第一个训练好的Estimator,得换个思路,用可配置的模型函数+分阶段初始化的方式来实现。
下面给你一步步的解决方案,附代码示例:
1. 先写一个支持双模式的模型函数
我们把模型函数改成可通过参数切换的,让它能在「训练卷积层」和「训练新全连接层」两种模式下工作:
def flexible_model_fn(features, labels, mode, params): # 定义要复用的卷积层基础结构 conv_base = tf.keras.Sequential([ tf.keras.layers.Conv2D(32, (3,3), activation='relu', input_shape=(28,28,1)), tf.keras.layers.MaxPooling2D((2,2)), tf.keras.layers.Conv2D(64, (3,3), activation='relu'), tf.keras.layers.MaxPooling2D((2,2)), tf.keras.layers.Flatten() ]) # 根据参数控制卷积层是否可训练:训练卷积层时开启,训全连接层时冻结 conv_base.trainable = (params['train_mode'] == 'conv') # 提取卷积层输出的特征 conv_features = conv_base(features['image']) # 根据模式选择全连接层和损失函数 if params['train_mode'] == 'conv': # 原全连接层 + 原损失函数 logits = tf.keras.layers.Dense(10)(conv_features) loss = tf.losses.sparse_categorical_crossentropy(labels, logits) else: # 新全连接层 + 新损失函数(这里用MSE举例,你换成自己的即可) logits = tf.keras.layers.Dense(5)(conv_features) loss = tf.losses.mean_squared_error(labels, logits) # 训练操作(两种模式共用,优化器会自动忽略不可训练变量) if mode == tf.estimator.ModeKeys.TRAIN: optimizer = tf.optimizers.Adam() train_op = optimizer.minimize( loss, global_step=tf.train.get_or_create_global_step() ) return tf.estimator.EstimatorSpec(mode=mode, loss=loss, train_op=train_op) # 省略评估和预测逻辑,你可以根据自己的需求补充(比如计算准确率、返回预测结果等)
2. 第一阶段:训练卷积层
创建第一个Estimator实例,指定模式为训练卷积层,用原损失函数训练:
# 初始化Estimator,指定训练模式为卷积层 conv_estimator = tf.estimator.Estimator( model_fn=flexible_model_fn, model_dir='./pre_trained_conv', # 保存卷积层权重的目录 params={'train_mode': 'conv'} ) # 开始训练(train_input_fn是你自己定义的训练数据输入函数) conv_estimator.train(input_fn=train_input_fn, steps=1000)
3. 第二阶段:冻结卷积层,训练新全连接层
创建新的Estimator实例,加载之前训练好的卷积层权重,切换模式为训练全连接层,用新损失函数:
# 可选:精准匹配卷积层变量,避免加载不存在的全连接层变量报错 warm_start_settings = tf.estimator.WarmStartSettings( ckpt_to_initialize_from='./pre_trained_conv', # 匹配卷积层相关的变量名(根据你实际的层名调整) vars_to_warm_start='conv2d.*|max_pooling2d.*|flatten.*' ) # 初始化新的Estimator,指定训练模式为全连接层 fc_estimator = tf.estimator.Estimator( model_fn=flexible_model_fn, model_dir='./new_fc_model', # 新的模型目录,避免覆盖原卷积层模型 params={'train_mode': 'fc'}, warm_start_from=warm_start_settings # 加载预训练的卷积层权重 ) # 用新的训练数据(如果需要)开始训练新全连接层 fc_estimator.train(input_fn=new_train_input_fn, steps=500)
关键注意点
- 为什么不能复用原来的Estimator?因为Estimator在创建时就把model_fn里的所有逻辑(包括损失函数、网络结构)固定了,没法中途修改。必须重新创建实例,通过params切换模式。
- 冻结卷积层的核心是把
conv_base.trainable设为False,这样优化器在训练时不会更新卷积层的权重。 - 用
warm_start_from加载预训练权重时,如果不指定vars_to_warm_start,它会尝试加载所有变量,但第二阶段新增的全连接层变量在原模型里不存在,会报错,所以精准匹配卷积层变量更稳妥。
这样就能实现你要的分阶段训练:先训卷积层,再冻结它训新全连接层换损失的需求啦。
内容的提问来源于stack exchange,提问作者Sam Ragusa
相关产品推荐
相关产品推荐

