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

训练CNN时测试准确率远高于训练准确率的原因排查

训练MNIST时测试准确率高于训练准确率的问题分析与解决

环境配置

  • Python 3.9.5
  • torch 1.13.0+cu117
  • torchvision 0.14.0+cu117

问题现象

基于MNIST训练CNN图像分类模型,训练过程中测试准确率始终高于训练准确率,与预期相反。训练结果如下:

epoch=1, train loss=0.8197974562644958, train acc=0.7494, test loss=0.1455492526292801, test acc=0.9616
epoch=2, train loss=0.7107925415039062, train acc=0.7788333333333334, test loss=0.1208220049738884, test acc=0.9689
epoch=3, train loss=0.6579669713973999, train acc=0.7906666666666666, test loss=0.11497163027524948, test acc=0.9676
epoch=4, train loss=0.6305248141288757, train acc=0.7994333333333333, test loss=0.10593992471694946, test acc=0.97
epoch=5, train loss=0.5982099771499634, train acc=0.80585, test loss=0.09132635593414307, test acc=0.9714
epoch=6, train loss=0.5825754404067993, train acc=0.8125333333333333, test loss=0.09170813113451004, test acc=0.9723
epoch=7, train loss=0.5688086748123169, train acc=0.8155166666666667, test loss=0.08628570288419724, test acc=0.9737
epoch=8, train loss=0.5556393265724182, train acc=0.8193166666666667, test loss=0.08203426003456116, test acc=0.9762
epoch=9, train loss=0.546567976474762, train acc=0.8213833333333334, test loss=0.08405696600675583, test acc=0.9754
epoch=10, train loss=0.5374698638916016, train acc=0.8239333333333333, test loss=0.07133891433477402, test acc=0.9788
epoch=11, train loss=0.5179286599159241, train acc=0.82975, test loss=0.0744888037443161, test acc=0.9792
epoch=12, train loss=0.5131004452705383, train acc=0.8329, test loss=0.07630482316017151, test acc=0.9778
epoch=14, train loss=0.49787914752960205, train acc=0.8366666666666667, test loss=0.07209591567516327, test acc=0.9779
epoch=15, train loss=0.4968840777873993, train acc=0.83475, test loss=0.07035819441080093, test acc=0.9801
epoch=16, train loss=0.4877821207046509, train acc=0.83925, test loss=0.07009950280189514, test acc=0.9777
epoch=17, train loss=0.48330068588256836, train acc=0.84045, test loss=0.06527410447597504, test acc=0.9809
epoch=18, train loss=0.48005640506744385, train acc=0.8404166666666667, test loss=0.06624794006347656, test acc=0.9781
epoch=19, train loss=0.47614845633506775, train acc=0.8418833333333333, test loss=0.07185563445091248, test acc=0.9788

训练代码

from torch.utils.data import DataLoader
from torchvision import datasets, transforms
from pathlib import Path
import torch

from CNN import CNNmodel

SEED = 5
device = "cuda" if torch.cuda.is_available() else "cpu"
BATCH_SIZE = 16
data_root = Path("data/")

torch.manual_seed(SEED)
torch.cuda.manual_seed(SEED)

train_transform = transforms.Compose([
transforms.TrivialAugmentWide(num_magnitude_bins=8),
transforms.ToTensor()
])

test_transform = transforms.ToTensor()

train_data = datasets.MNIST(
root=data_root / "train",
train=True,
download=True,
transform=train_transform
)

test_data = datasets.MNIST(
root=data_root / "test",
train=False,
download=True,
transform=test_transform
)

train_dataloader = DataLoader(
train_data,
batch_size=BATCH_SIZE,
shuffle=True
)

test_dataloader = DataLoader(
test_data,
batch_size=BATCH_SIZE,
shuffle=False
)

channel_num = train_data[0][0].shape[0]
model = CNNmodel(in_shape=channel_num, hidden_shape=8, out_shape=len(train_data.classes)).to(device)
optimizer = torch.optim.SGD(params=model.parameters(), lr=0.01)
loss_fn = torch.nn.CrossEntropyLoss()
epochs = 20

def train_step(dataloader, loss_fn, optimizer, model, device):
    train_loss = 0
    train_acc = 0

    for batch, (X, y) in enumerate(dataloader):
        X, y = X.to(device), y.to(device)
    
        y_pred = model(X)
    
        loss = loss_fn(y_pred, y)
        train_loss += loss
        
        optimizer.zero_grad()
    
        loss.backward()
    
        optimizer.step()
    
        y_pred_class = torch.argmax(torch.softmax(y_pred, dim=1), dim=1)
        train_acc += (y_pred_class == y).sum().item()/len(y_pred)
    
    train_loss /= len(dataloader)
    train_acc /= len(dataloader)
    
    return (train_loss, train_acc)

def test_step(dataloader, loss_fn, model, device):
    test_loss = 0
    test_acc = 0

    with torch.inference_mode():
        for batch, (X, y) in enumerate(dataloader):
            X, y = X.to(device), y.to(device)
    
            y_pred = model(X)
    
            loss = loss_fn(y_pred, y)
            test_loss += loss
    
            y_pred_class = torch.argmax(torch.softmax(y_pred, dim=1), dim=1)
            test_acc += (y_pred_class == y).sum().item()/len(y_pred)
        
        test_loss /= len(dataloader)
        test_acc /= len(dataloader)
    
    return (test_loss, test_acc)

for epoch in range(epochs):
    train_loss, train_acc = train_step(
        dataloader=train_dataloader,
        loss_fn=loss_fn,
        optimizer=optimizer,
        model=model,
        device=device
    )

    test_loss, test_acc = test_step(
        dataloader=test_dataloader,
        loss_fn=loss_fn,
        model=model,
        device=device
    )
    
    torch.cuda.empty_cache()
    print(f"epoch={epoch}, train loss={train_loss}, train acc={train_acc}, test loss={test_loss}, test acc={test_acc}\n")

模型结构

import torch.nn as nn

class CNNmodel(nn.Module):
    def __init__(self, in_shape, hidden_shape, out_shape) -> None:
        super().__init__()
        self.conv_block_1 = nn.Sequential(
            nn.Conv2d(
                in_channels=in_shape,
                out_channels=hidden_shape,
                kernel_size=3,
                stride=1,
                padding=1
            ),
            nn.ReLU(),
            nn.Conv2d(
                in_channels=hidden_shape,
                out_channels=hidden_shape,
                kernel_size=3,
                stride=1,
                padding=1
            ),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2, stride=2)
        )
        self.conv_block_2 = nn.Sequential(
            nn.Conv2d(
                in_channels=hidden_shape,
                out_channels=hidden_shape,
                kernel_size=3,
                stride=1,
                padding=1
            ),
            nn.ReLU(),
            nn.Conv2d(
                in_channels=hidden_shape,
                out_channels=hidden_shape,
                kernel_size=3,
                stride=1,
                padding=1
            ),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2)
        )
        self.classifier = nn.Sequential(
            nn.Flatten(),
            nn.Linear(in_features=hidden_shape*7*7,
                      out_features=out_shape)
        )
    
    def forward(self, x):
        return self.classifier(self.conv_block_2(self.conv_block_1(x)))

原因分析

  1. 训练集数据增强提升了任务难度:训练集使用TrivialAugmentWide做随机图像变换,生成的带噪声样本比原始MNIST测试集更难识别,导致训练准确率偏低。
  2. 学习率过低导致收敛缓慢:SGD优化器搭配0.01的学习率,对于MNIST任务来说偏保守,20个epoch内模型还未充分拟合训练数据,但足够应对简单的测试集。
  3. 模型容量不足:卷积层仅用8个通道,模型拟合能力有限,无法很好地学习增强后训练集的复杂特征,但能轻松适配干净的测试集。
  4. 训练准确率计算时机问题:训练时每步更新参数后直接计算准确率,此时模型参数处于动态变化状态,结果会受参数更新噪声影响,而测试时参数固定,结果更稳定。

解决方法

  • 调整数据增强策略:降低增强强度(如减少num_magnitude_bins),或单独创建无增强的训练DataLoader,用于计算训练准确率,避免增强样本干扰评估结果。
  • 优化学习率设置:将SGD学习率提升至0.05或0.1,同时可搭配StepLR等学习率调度器,在训练后期逐步降低学习率,加快模型收敛。
  • 增大模型容量:把hidden_shape从8调整为16或32,或添加额外的卷积层、全连接层,提升模型拟合复杂训练数据的能力。
  • 规范训练准确率计算流程:每个epoch训练完成后,切换模型到评估模式(model.eval()),用无增强的训练集计算准确率,消除参数更新过程中的噪声影响。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 01:14:57