VGG16嵌入SENet模块后在Imagenette数据集训练准确率为0求助
我尝试在VGG16网络的多个中间位置嵌入SENet模块,再使用Imagenette数据集训练模型,但训练得到的准确率为0。我是深度学习入门使用者,希望能得到帮助排查解决该问题。
已完成操作
- 定义SE模块并嵌入到VGG16的多处位置
- 加载Imagenette数据集,为缩短训练时长仅取用部分数据
- 划分训练集与测试集后完成模型编译与训练,相关代码如下:
import tensorflow as tf from tensorflow.keras.applications.vgg16 import VGG16 as Model from tensorflow.keras.applications.vgg16 import preprocess_input import numpy as np from matplotlib import pyplot as plt # %matplotlib inline from tensorflow.keras.preprocessing.image import load_img, img_to_array from tensorflow.keras.preprocessing import image from tensorflow.keras import backend as K from vis.utils import utils from tensorflow.keras.applications.vgg16 import decode_predictions import json import io from tensorflow.keras import layers import os from keras.models import Sequential from keras.layers import Dense, Conv2D, MaxPool2D , Flatten from keras.preprocessing.image import ImageDataGenerator from tensorflow.keras.layers import GlobalAveragePooling2D, Reshape, Dense, Permute, Multiply def reshape(img,label): img = tf.cast(img, tf.float32) img = tf.image.resize(img, (224,224)) img = img/255.0 return img, label def reshape(img,label): img = tf.cast(img, tf.float32) img = tf.image.resize(img, (224,224)) img = img/255.0 return img, label import tensorflow_datasets as tfds imgn_dataset = tfds.builder('imagenette') imgn_dataset.info.features['label'].num_classes == 10 imgn_dataset.download_and_prepare() datasets = imgn_dataset.as_dataset(as_supervised = True) train_data, test_data = datasets['train'], datasets['validation'] train_data = train_data.map(reshape) train_data = train_data.batch(128) train_data = train_data.prefetch(tf.data.experimental.AUTOTUNE) test_data = test_data.map(reshape) test_data = test_data.batch(128) test_data = test_data.prefetch(tf.data.experimental.AUTOTUNE) # Load model model = Model(weights='imagenet', include_top=True) model.summary() # Squeeze and Excitation def se_block(input, channels, r=8): # Squeeze x = GlobalAveragePooling2D()(input) # Excitation x = Dense(channels//r, activation="relu")(x) x = Dense(channels, activation="sigmoid")(x) return Multiply()([input, x]) input_dim = (224,224,3) img_input = layers.Input(shape = input_dim) sen = se_block(img_input, 3) vgg_model = model.get_layer("block1_conv1")(sen) vgg_model = model.get_layer("block1_conv2")(vgg_model) vgg_model = model.get_layer("block1_pool")(vgg_model) sen = se_block(vgg_model, 64) vgg_model = model.get_layer("block2_conv1")(sen) vgg_model = model.get_layer("block2_conv2")(vgg_model) vgg_model = model.get_layer("block2_pool")(vgg_model) sen = se_block(vgg_model, 128) vgg_model = model.get_layer("block3_conv1")(sen) vgg_model = model.get_layer("block3_conv2")(vgg_model) vgg_model = model.get_layer("block3_conv3")(vgg_model) vgg_model = model.get_layer("block3_pool")(vgg_model) sen = se_block(vgg_model, 256) vgg_model = model.get_layer("block4_conv1")(sen) vgg_model = model.get_layer("block4_conv2")(vgg_model) vgg_model = model.get_layer("block4_conv3")(vgg_model) vgg_model = model.get_layer("block4_pool")(vgg_model) sen = se_block(vgg_model, 512) vgg_model = model.get_layer("block5_conv1")(sen) vgg_model = model.get_layer("block5_conv2")(vgg_model) vgg_model = model.get_layer("block5_conv3")(vgg_model) vgg_model = model.get_layer("block5_pool")(vgg_model) sen = se_block(vgg_model, 512) vgg_model = model.get_layer("flatten")(sen) vgg_model = model.get_layer("fc1")(vgg_model) vgg_model = model.get_layer("fc2")(vgg_model) vgg_model = tf.keras.models.Model(img_input, vgg_model) vgg_model.summary() vgg_model.get_layer("block1_conv1").trainable = False vgg_model.get_layer("block1_conv2").trainable = False vgg_model.get_layer("block1_pool").trainable = False vgg_model.get_layer("block2_conv1").trainable = False vgg_model.get_layer("block2_conv1").trainable = False vgg_model.get_layer("block2_pool").trainable = False vgg_model.get_layer("block3_conv1").trainable = False vgg_model.get_layer("block3_conv1").trainable = False vgg_model.get_layer("block3_conv1").trainable = False vgg_model.get_layer("block3_pool").trainable = False vgg_model.get_layer("block4_conv1").trainable = False vgg_model.get_layer("block4_conv1").trainable = False vgg_model.get_layer("block4_conv1").trainable = False vgg_model.get_layer("block4_pool").trainable = False vgg_model.get_layer("block5_conv1").trainable = False vgg_model.get_layer("block5_conv1").trainable = False vgg_model.get_layer("block5_conv1").trainable = False vgg_model.get_layer("block5_pool").trainable = False vgg_model.get_layer("flatten").trainable = False vgg_model.get_layer("fc1").trainable = False vgg_model.get_layer("fc2").trainable = False vgg_model.compile(optimizer = 'adam', loss ='mse', metrics=['accuracy'] ) vgg_model.summary() train_data = train_data.take(1) test_data = test_data.take(1) vgg_model.fit(train_data, epochs=10, validation_data=test_data)
问题排查与修复方案
1. 输出维度不匹配
你加载的VGG16开启了include_top=True,原始fc2层输出是1000类(适配ImageNet数据集),而Imagenette只有10个类别,输出维度完全不匹配,无法得到正确分类结果。
修复方法:加载VGG16时设置include_top=False,去掉原始顶层,新增适配10分类的输出层,激活函数使用softmax。
2. 损失函数选择错误
你使用的MSE(均方误差)是回归任务的损失函数,10分类属于多分类任务,且你的标签是整数编码,需要使用SparseCategoricalCrossentropy作为损失函数。
3. 模型几乎没有可训练参数
你把所有VGG内置层的trainable都设为False,只有新增SE模块的少量全连接层可训练,同时仅取1个batch(128个样本)训练,样本量不足以覆盖10个分类的特征,模型完全无法学习有效信息。
修复方法:
- 放开VGG高层卷积层和全连接层的训练权限,或者先训练新增的SE模块和自定义分类头,收敛后再微调VGG底层参数
- 去掉
train_data.take(1)和test_data.take(1)的限制,至少使用1/10以上的训练数据,保证样本覆盖所有10个类别
4. 预处理逻辑与预训练权重不兼容
VGG16的预训练权重适配的预处理是自带的preprocess_input(按ImageNet数据集均值做中心化),你直接除以255的操作会破坏输入分布与预训练权重的适配性。
修复方法:把预处理逻辑替换为VGG16自带的preprocess_input。
5. 代码冗余错误
代码中重复定义两次reshape函数,设置层可训练状态时重复书写多次同一层名(比如block3_conv1写了三次,仅第一层生效,后两层未被设置),虽不直接导致准确率为0,但会埋下逻辑隐患,建议清理冗余代码。
内容的提问来源于stack exchange,提问作者Ram

