You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在TensorFlow中冻结自编码器卷积层并保存编码器部分

具体实现方案

冻结前三个卷积层参数

你代码里的前三个卷积层分别是编码器段的conv1、conv3、conv6三个Conv2D层,冻结参数的核心逻辑是将目标层的trainable属性设为False,被设置的层在训练过程中不会更新权重。
注意冻结操作必须在模型编译前完成,如果是训练完成后再做冻结微调,冻结后需要重新编译模型才能生效。如果不需要三个卷积层附属的BatchNormalization层更新滑动统计量,可以把相邻的BN层也加入冻结列表。

提取并保存训练后的编码器

训练完成后不需要重新搭建编码器结构,直接基于训练好的自编码器模型,以原模型输入为输入,定位到编码器段最后一层的输出作为编码器输出,单独实例化一个Model对象即可,之后可以直接调用保存接口存储模型。


修改后的可直接运行代码

# 补全原代码缺失的输入层定义,输入维度根据你的实际数据调整
input_img = Input(shape=(28, 28, 1))

def encoder(input_img):
    #encoder
    #input = 28 x 28 x 1 (wide and thin)
    conv1 = Conv2D(64, (2,2), activation='relu', padding='same', name='conv1')(input_img) #28 x 28 x 64
    conv2 = BatchNormalization(name='bn1')(conv1)
    conv3 = Conv2D(32, (2,2), activation='relu', padding='same', name='conv2')(conv2)
    conv4 = BatchNormalization(name='bn2')(conv3)
    pool5 = MaxPooling2D(pool_size=(2,2), name='pool1')(conv4) #14 x 14 x 32
    conv6 = Conv2D(16, (2,2), activation='relu', padding='same', name='conv3')(pool5) #14 x 14 x 16
    conv7 = BatchNormalization(name='bn3')(conv6)
    conv8 = Conv2D(8, (2,2), activation='relu', padding='same', name='conv4')(conv7)
    conv9 = BatchNormalization(name='bn4')(conv8)
    conv10 = Conv2D(4, (2,2), activation='relu', padding='same', name='encoder_output')(conv9)
    return conv10

def decoder(conv11):    
    #decoder
    conv12 = Conv2D(4, (2,2), activation='relu', padding='same')(conv11)
    conv13 = Conv2D(8, (2,2), activation='relu', padding='same')(conv12)
    conv14 = BatchNormalization()(conv13)
    conv15 = Conv2D(16, (2,2), activation='relu', padding='same')(conv14)
    conv16 = BatchNormalization()(conv15)
    conv17 = Conv2D(32, (2,2), activation='relu', padding='same')(conv16)
    conv18 = BatchNormalization()(conv17)
    conv19 = Conv2D(64, (2,2), activation='relu', padding='same')(conv18)
    conv20 = BatchNormalization()(conv19)
    up21 = UpSampling2D((2,2))(conv20)
    decoded = Conv2D(3, (2,2), activation='sigmoid', padding='same')(up21)
    return decoded

autoencoder = Model(input_img, decoder(encoder(input_img)))

# 冻结前三个卷积层,这里用层名定位比索引更稳妥,不会因为层顺序变化出错
freeze_layer_names = ['conv1', 'conv2', 'conv3']
for layer in autoencoder.layers:
    if layer.name in freeze_layer_names:
        layer.trainable = False
    # 如果需要连带冻结对应BN层,把bn1、bn2、bn3也加进freeze_layer_names即可

# 冻结完成后再编译模型
autoencoder.compile(loss='mae', optimizer = 'SGD')
# 打印模型摘要,通过Non-trainable params数值确认冻结生效
autoencoder.summary()

# 原有训练逻辑不变
train = np.concatenate((normal[0:1900,:,:,:],un_informative[0:1900,:,:,:]),axis=0) 
valid = np.concatenate((normal[1900:,:,:,:],un_informative[1900:,:,:,:]),axis=0)
history = autoencoder.fit(train ,train , batch_size=batch_size,epochs=200,verbose=1, validation_data=(valid, valid))

# 提取训练好的编码器
trained_encoder = Model(
    inputs=autoencoder.input,
    outputs=autoencoder.get_layer('encoder_output').output
)

# 保存编码器,二选一即可
trained_encoder.save('trained_encoder.h5') # 保存为单H5文件
# trained_encoder.save('trained_encoder_savedmodel') # 保存为TensorFlow SavedModel格式

提示:如果不想给层加name参数,也可以通过autoencoder.summary()打印的层顺序,用get_layer(index=层序号)的方式定位目标层,注意序号从0开始计数,核对无误后再执行冻结和提取操作。


内容的提问来源于stack exchange,提问作者aliashraf afrah

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.30 08:31:06