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

VGG16嵌入SENet模块后在Imagenette数据集训练准确率为0求助

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 09:57:03