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

TensorFlow图像分类模型训练至第2轮后精度停滞问题求助

交通手势4分类CNN模型训练停滞问题排查与解决思路

我搭建了一个用于交通手势(停止、通行、左转、右转)的4分类CNN模型,训练到第2个epoch后,后续所有epoch的精度都不再变化。我怀疑过数据问题,更换了不同数据集,也微调了学习率、dropout率、网络层数和滤波器数量,但问题依然存在。

问题代码

def fit_cnn_model(X_train, Y_train, X_test, Y_test, savedir, num_layers=2,
                  num_filters=64, kernel_size=3, pool_size=2, dropout_rate=0.3, learning_rate=0.0001,
                  n_epochs=10, batch_size=64, verbose=2):

    model = Sequential()
    date_time = datetime.now().strftime("%Y-%m-%d %H:%M:%S")

    for _ in range(num_layers):
        model.add(Conv2D(filters=num_filters, kernel_size=kernel_size, activation='relu', padding='same', kernel_initializer='he_uniform'))
        model.add(BatchNormalization())
        model.add(MaxPooling2D(pool_size=pool_size))
        model.add(Dropout(dropout_rate))

    model.add(Flatten())
    model.add(Dense(128, activation='relu'))
    model.add(Dense(4, activation='softmax'))

    opt = Adam(learning_rate=learning_rate)
    model.compile(optimizer=opt, loss='categorical_crossentropy', metrics=['accuracy'])

    history = model.fit(X_train, Y_train,
                        epochs=n_epochs,
                        batch_size=batch_size,
                        validation_data=(X_test, Y_test),
                        verbose=verbose)

    model.summary()

    scores = model.evaluate(X_test, Y_test, verbose=0)
    print("Loss: {:.2f}".format(scores[0]))
    print("Accuracy: {:.2f}%".format(scores[1] * 100))

    model_path = os.path.join(savedir, 'saved_model_' + date_time + '.h5')
    model.save(model_path)
    print(f"Model saved to {model_path}")

    return model

def load_data(image_folder, label_folder, target_size=(224, 224)):
    # Load and sort dataset
    image_files = os.listdir(image_folder)
    label_files = os.listdir(label_folder)
    image_files.sort()
    label_files.sort()

    images = []
    labels = []

    # Load images and labels
    for img_file in image_files:
        img_path = os.path.join(image_folder, img_file)

        label_file = img_file.replace(".jpg", ".txt").replace(".jpeg", ".txt")
        label_path = os.path.join(label_folder, label_file)

        img = Image.open(img_path).resize(target_size)
        img_array = np.array(img) / 255.0
        images.append(img_array)

        with open(label_path, 'r') as f:
            label_str = f.readline().strip()
            label_dict = {'stop': 0, 'continue': 1, 'left': 2, 'right': 3}
            label = label_dict.get(label_str, -1)
            if label == -1:
                raise ValueError("Unknown label: {}".format(label_str))

        labels.append(label)

    images = np.array(images)
    labels = np.array(labels)
    labels = to_categorical(labels, 4)


    X_train, X_test, Y_train, Y_test = train_test_split(images, labels, test_size=0.2)

    return X_train, Y_train, X_test, Y_test

排查方向与解决办法

1. 数据环节

  • 标签校验:逐一检查标签文件内容,确认没有拼写错误(比如continue写成contiue、大小写不一致),确保所有标签都能匹配到label_dict中的键。
  • 固定数据划分:在train_test_split中添加random_state=42,固定训练/测试集划分,避免每次训练数据分布差异导致结果不可复现。
  • 数据增强:在训练时添加数据增强,比如随机水平翻转、旋转±15度、缩放0.8-1.2倍,提升数据多样性,避免模型过早过拟合。示例代码:
    from tensorflow.keras.preprocessing.image import ImageDataGenerator
    
    datagen = ImageDataGenerator(
        horizontal_flip=True,
        rotation_range=15,
        zoom_range=0.2
    )
    datagen.fit(X_train)
    # 训练时使用datagen.flow
    history = model.fit(datagen.flow(X_train, Y_train, batch_size=batch_size),
                        epochs=n_epochs,
                        validation_data=(X_test, Y_test),
                        verbose=verbose)
    
  • 通道一致性:确认所有输入图像都是3通道RGB格式,若存在灰度图,需转换为3通道(img_array = np.repeat(img_array[..., np.newaxis], 3, axis=-1))。

2. 模型结构优化

  • 调整BN与激活顺序:当前代码中Conv2D直接带activation='relu',之后加BN,正确顺序应为卷积→BN→激活,这样BN能对线性输出做归一化,提升训练稳定性:
    model.add(Conv2D(filters=num_filters, kernel_size=kernel_size, padding='same', kernel_initializer='he_uniform'))
    model.add(BatchNormalization())
    model.add(Activation('relu'))
    
  • 优化特征提取流程:连续的池化会快速压缩特征图,丢失细节。可以改为每2层卷积加1次池化,同时递增滤波器数量(比如第一层32,第二层64):
    # 示例:2组卷积+池化,滤波器数量递增
    model.add(Conv2D(32, 3, padding='same', kernel_initializer='he_uniform'))
    model.add(BatchNormalization())
    model.add(Activation('relu'))
    model.add(Conv2D(32, 3, padding='same', kernel_initializer='he_uniform'))
    model.add(BatchNormalization())
    model.add(Activation('relu'))
    model.add(MaxPooling2D(2))
    model.add(Dropout(0.1))
    
    model.add(Conv2D(64, 3, padding='same', kernel_initializer='he_uniform'))
    model.add(BatchNormalization())
    model.add(Activation('relu'))
    model.add(Conv2D(64, 3, padding='same', kernel_initializer='he_uniform'))
    model.add(BatchNormalization())
    model.add(Activation('relu'))
    model.add(MaxPooling2D(2))
    model.add(Dropout(0.1))
    
  • 替换Flatten为全局平均池化:全局平均池化(GlobalAveragePooling2D)能减少参数数量,避免过拟合,同时保留空间特征:
    model.add(GlobalAveragePooling2D())
    model.add(Dense(128, activation='relu'))
    model.add(Dense(4, activation='softmax'))
    

3. 训练策略调整

  • 学习率优化:固定0.0001的学习率可能太小,导致模型快速收敛到局部最优。可以先调高初始学习率到0.001,再添加学习率衰减回调:
    from tensorflow.keras.callbacks import ReduceLROnPlateau
    
    lr_scheduler = ReduceLROnPlateau(monitor='val_accuracy', factor=0.5, patience=2, min_lr=1e-6)
    history = model.fit(..., callbacks=[lr_scheduler])
    
  • 添加训练监控回调:用ModelCheckpoint保存最优模型,同时绘制loss和accuracy曲线,明确是训练精度停滞还是验证精度停滞:
    from tensorflow.keras.callbacks import ModelCheckpoint
    import matplotlib.pyplot as plt
    
    checkpoint = ModelCheckpoint(os.path.join(savedir, 'best_model.h5'), monitor='val_accuracy', save_best_only=True)
    history = model.fit(..., callbacks=[checkpoint, lr_scheduler])
    
    # 绘制曲线
    plt.plot(history.history['accuracy'], label='Train Accuracy')
    plt.plot(history.history['val_accuracy'], label='Val Accuracy')
    plt.legend()
    plt.show()
    
  • 类别权重平衡:如果数据集类别不平衡,在model.fit中添加class_weight参数,比如某类样本只有其他类的1/2,权重设为2:
    class_weight = {0:1, 1:1, 2:2, 3:1} # 根据实际样本数量调整
    history = model.fit(..., class_weight=class_weight)
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 00:39:50