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

预训练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()
回答

可行性与操作修正

这种方式完全可行,你的核心思路没问题,但代码存在几个关键细节需要修正:

  1. 必须显式冻结预训练VGG层
    当前代码没有冻结VGG19的卷积层,训练时会导致所有参数被更新。要实现仅训练SENET和全连接层,需添加以下代码:

    for layer in Vgg.layers:
        layer.trainable = False
    
  2. SENET模块实现修正
    你的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 13:11:37