预训练VGG19中仅训练SENET模块与全连接层是否可行?
问题描述
我正在使用预训练VGG19模型,并在其层间加入《Squeeze and Excitation Network》论文中提及的Squeeze-and-Excitation注意力(SENET)模块,仅训练最后全连接层与SENET模块。
这种方式是否可行?我的操作是否正确,还是需要训练整个网络?(我选择仅训练SENET的思路是,它可通过在训练过程中学习注意力权重来抑制或增强特定特征通道。)
参考代码
import tensorflow as tf import numpy as np from tensorflow.keras.layers import Activation, Conv2D, Input, BatchNormalization, Reshape, GlobalAveragePooling2D, AveragePooling2D, GlobalMaxPooling2D, Flatten , ReLU, Layer, Dense from tensorflow.keras.activations import sigmoid, softmax, relu, tanh from tensorflow.keras import Sequential from tensorflow.keras.applications import VGG19 from tensorflow.keras.models import Model class SENET_Attn(Layer): """ Channel Attention Block as reported in SENET """ def __init__(self,out_dim, ratio, layer_name="SENET"): super(SENET_Attn, self).__init__() self.out_dim = out_dim self.ratio = ratio self.layer_name = layer_name def build(self, ratio, layer_name="SENET"): self.Global_Average_Pooling = GlobalAveragePooling2D(keepdims= True) self.Fully_connected_1_1 = Dense(units= self.out_dim/self.ratio, name=self.layer_name+'_fully_connected1', kernel_initializer="glorot_uniform") self.Relu = ReLU() self.Fully_connected_2 = Dense(units=self.out_dim, name=layer_name+'_fully_connected2', activation = "tanh") self.Sigmoid = Activation("sigmoid") def call(self, inputs): inputs = tf.cast(inputs, dtype = "float32") squeeze = self.Global_Average_Pooling(inputs) excitation = self.Fully_connected_1_1(squeeze) excitation = self.Relu(excitation) excitation = self.Fully_connected_2(excitation) excitation = self.Sigmoid(excitation) excitation = tf.reshape(excitation, [-1,1,1,self.out_dim]) scale = inputs * excitation return scale Vgg = VGG19(include_top = False,input_shape = (50,50,3)) # SENET-Attn VGG19 ratio =16 input_layer = Input(shape=(50,50,3)) out = Vgg.layers[1](input_layer) # Block one of VGG19 out = Vgg.layers[2](out) out = Vgg.layers[3](out) out = Vgg.layers[4](out) # Block two of VGG19 out = Vgg.layers[5](out) out = Vgg.layers[6](out) #out = SENET_Attn(out.shape[-1], ratio, )(out) # SENET Attention out = Vgg.layers[7](out) # Block three of VGG19 out = Vgg.layers[8](out) out = Vgg.layers[9](out) out = Vgg.layers[10](out) out = Vgg.layers[11](out) out = SENET_Attn(out.shape[-1], ratio, )(out) # SENET Attention out = Vgg.layers[12](out) # Block four of VGG19 out = Vgg.layers[13](out) out = SENET_Attn(out.shape[-1], ratio, )(out) # SENET Attention out = Vgg.layers[14](out) out = Vgg.layers[15](out) out = SENET_Attn(out.shape[-1], ratio, )(out) # SENET Attention out = Vgg.layers[16](out) out = SENET_Attn(out.shape[-1], ratio, )(out) # SENET Attention out = Vgg.layers[17](out) #Block five of VGG19 out = Vgg.layers[18](out) out = SENET_Attn(out.shape[-1], ratio, )(out) # SENET Attention out = Vgg.layers[19](out) out = Vgg.layers[20](out) out = SENET_Attn(out.shape[-1], ratio, )(out) # SENET Attention out = Vgg.layers[21](out) out = SENET_Attn(out.shape[-1], ratio, )(out) # SENET Attention flatten = Flatten()(out) out = Dense(100)(flatten) out = tf.keras.layers.ReLU()(out) out = Dense(100)(out) out = tf.keras.layers.ReLU()(out) out = Dense(2)(out) out = tf.keras.layers.Softmax()(out) model = Model(input_layer, out) model.compile(optimizer= tf.keras.optimizers.SGD(), loss= tf.keras.losses.CategoricalCrossentropy() , metrics=['accuracy']) model.summary()
回答
可行性与操作修正
这种方式完全可行,你的核心思路没问题,但代码存在几个关键细节需要修正:
必须显式冻结预训练VGG层
当前代码没有冻结VGG19的卷积层,训练时会导致所有参数被更新。要实现仅训练SENET和全连接层,需添加以下代码:for layer in Vgg.layers: layer.trainable = FalseSENET模块实现修正
你的SENET_Attn类有两处不符合原论文的设计:build方法的参数错误,标准Keras Layer的build方法应接收input_shape作为参数,而非自定义参数- 第二个全连接层不应使用
tanh激活,原论文中该层无激活,仅在最后用Sigmoid生成注意力权重
修正后的SENET模块:
class SENET_Attn(Layer): """通道注意力模块,参考SENET论文实现""" def __init__(self, out_dim, ratio, layer_name="SENET"): super(SENET_Attn, self).__init__() self.out_dim = out_dim self.ratio = ratio self.layer_name = layer_name def build(self, input_shape): self.Global_Average_Pooling = GlobalAveragePooling2D(keepdims=True) self.Fully_connected_1 = Dense(units=self.out_dim // self.ratio, name=f"{self.layer_name}_fc1", kernel_initializer="glorot_uniform") self.Relu = ReLU() self.Fully_connected_2 = Dense(units=self.out_dim, name=f"{self.layer_name}_fc2") self.Sigmoid = Activation("sigmoid") super().build(input_shape) # 必须调用父类build方法 def call(self, inputs): inputs = tf.cast(inputs, dtype="float32") squeeze = self.Global_Average_Pooling(inputs) excitation = self.Fully_connected_1(squeeze) excitation = self.Relu(excitation) excitation = self.Fully_connected_2(excitation) excitation = self.Sigmoid(excitation) excitation = tf.reshape(excitation, [-1, 1, 1, self.out_dim]) scale = inputs * excitation return scale
是否需要训练整个网络?
- 小数据集场景:不需要训练整个网络,冻结VGG仅训练SENET和全连接层即可,既能利用预训练模型的通用特征提取能力,又能避免过拟合。
- 大数据集场景:可以分阶段训练:先冻结VGG训练SENET和全连接层,再解冻VGG顶部2-3个卷积块(如block4、block5),用更小的学习率(原学习率的1/10)进行微调,让预训练特征更好适配你的任务。
总结
你的思路方向正确,通过SENET的通道注意力机制调整特征权重,同时冻结预训练层减少计算量和过拟合风险。修正代码中的冻结操作和SENET实现细节后,即可正常训练目标层;是否微调整个网络取决于你的数据集规模。
内容的提问来源于stack exchange,提问作者harshit tiwari
相关产品推荐
相关产品推荐

