模型训练 不同批次大小与学习率下的训练损失对比(1)

模型训练 不同批次大小与学习率下的训练损失对比 (1)

batchsize 和 learning rate存在什么关系,可以怎么做

flyfish

一、线性缩放规则(Linear Scaling Rule)

公式

ηnew=ηold×BnewBold\eta_{new} = \eta_{old} \times \frac{B_{new}}{B_{old}}ηnew=ηold×BoldBnew

批次扩大多少倍,学习率就同步扩大多少倍。

经验规则:批次越大,梯度估计的噪声越小、方向越稳定,可以承受更大的更新步长,因此按比例放大学习率,就能保持收敛速度和训练动力学基本一致。

适用条件

优化器:原生SGD / 带动量的SGD(对梯度方差最敏感)

二、平方根缩放规则(Square Root Scaling Rule)

公式

ηnew=ηold×BnewBold\eta_{new} = \eta_{old} \times \sqrt{\frac{B_{new}}{B_{old}}}ηnew=ηold×BoldBnew

学习率与批次大小的平方根成正比。

梯度估计的方差 与批次大小成反比(Var(g)∝1/BVar(g) \propto 1/BVar(g)∝1/B),梯度的标准差 与 B\sqrt{B}B 成反比。要保持梯度更新的信噪比不变,学习率只需和批次的平方根同步增长即可。

适用条件

优化器:SGD系列

三、常数规则(自适应优化器默认)

做法

批次大小变化时,学习率基本保持不变,或仅做小幅微调,不需要严格按比例缩放。

Adam、AdamW、RMSprop等自适应优化器,会根据梯度的一阶矩、二阶矩自动归一化更新步长,对批次大小带来的梯度方差变化鲁棒性极强,批次的波动会被优化器自身抵消。

适用条件

优化器:Adam / AdamW / RMSprop 等自适应优化器

验证下

Ubuntu22.04

安装中文字体

bash 复制代码
sudo apt install fonts-noto-cjk
sudo apt install fonts-wqy-microhei
sudo apt install fonts-noto-cjk

SGD的情况下 不同批次大小与学习率下的训练损失对比

bash 复制代码
批次大小  32  →  学习率 0.0100
批次大小  64  →  学习率 0.0200
批次大小 128  →  学习率 0.0400
批次大小 256  →  学习率 0.0800
python 复制代码
# 导入PyTorch核心库、神经网络模块与优化器模块
import torch
import torch.nn as nn
import torch.optim as optim
# 导入torchvision提供的数据集和数据预处理工具
from torchvision import datasets, transforms
# 导入数据加载器,用于按批次批量加载训练数据
from torch.utils.data import DataLoader
# 导入matplotlib绘图库,用于可视化训练损失曲线
import matplotlib.pyplot as plt


# 配置matplotlib中文显示
plt.rcParams['font.sans-serif'] = ['AR PL UMing CN']
#plt.rcParams['font.sans-serif'] = ['SimHei']  # 设置中文字体为黑体,解决中文乱码
plt.rcParams['axes.unicode_minus'] = False    # 解决坐标轴负号显示为方块的问题

# 定义一个简单的全连接神经网络类,继承PyTorch的nn.Module基类
class SimpleNN(nn.Module):
    def __init__(self):
        # 调用父类nn.Module的初始化方法
        super(SimpleNN, self).__init__()
        # 第一个全连接层:输入维度是28*28(MNIST单张图片的像素总数),输出维度128(隐藏层神经元数量)
        self.fc1 = nn.Linear(28*28, 128)
        # 第二个全连接层:输入维度128,输出维度10(对应MNIST的10个数字类别:0~9)
        self.fc2 = nn.Linear(128, 10)
    
    # 定义前向传播的计算逻辑
    def forward(self, x):
        # 将二维图片张量展平为一维向量:-1表示自动匹配批次大小,28*28是单张图片的特征维度
        x = x.view(-1, 28*28)
        # 第一个全连接层的输出经过ReLU激活函数,引入非线性变换
        x = torch.relu(self.fc1(x))
        # 第二个全连接层输出最终的分类得分(也叫logits)
        x = self.fc2(x)
        return x

# 定义数据预处理的流水线
transform = transforms.Compose([
    transforms.ToTensor(),  # 将图片转为PyTorch张量,同时把像素值从[0,255]缩放到[0,1]区间
    transforms.Normalize((0.5,), (0.5,))  # 单通道归一化:均值0.5、标准差0.5,将数值映射到[-1, 1]区间
])

# 加载MNIST手写数字训练数据集
# root:数据本地存储路径;train=True表示加载训练集;download=True表示本地无数据时自动下载
# transform:对每张图片应用上面定义的预处理操作
train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform)

# 定义模型训练函数(新增日志打印功能)
# 输入:批次大小batch_size、学习率learning_rate、训练轮数epochs(默认1轮)
# 输出:每一轮训练的平均损失列表
def train_model(batch_size, learning_rate, epochs=1):
    # 创建数据加载器:按指定批次大小加载数据,shuffle=True表示每轮训练打乱数据顺序
    train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)
    # 实例化我们定义的简单神经网络
    model = SimpleNN()
    # 定义损失函数:交叉熵损失,是多分类任务的标准损失函数
    criterion = nn.CrossEntropyLoss()
    # 定义优化器:随机梯度下降(SGD),负责更新网络参数,学习率为传入的learning_rate
    optimizer = optim.SGD(model.parameters(), lr=learning_rate)
    
    losses = []  # 用于存储每一个epoch的平均训练损失
    total_batch = len(train_loader)  # 总批次数量,用于显示进度
    
    for epoch in range(epochs):
        epoch_loss = 0  # 记录当前epoch的总损失
        # 逐批次遍历训练数据,用enumerate获取批次索引
        for batch_idx, (data, target) in enumerate(train_loader):
            optimizer.zero_grad()  # 清空上一步计算的梯度,避免梯度累积
            output = model(data)  # 前向传播:输入图片数据,得到模型预测结果
            loss = criterion(output, target)  # 计算预测结果与真实标签的损失值
            loss.backward()  # 反向传播:自动计算损失对每个网络参数的梯度
            optimizer.step()  # 参数更新:根据梯度和优化器规则调整网络参数
            epoch_loss += loss.item()  # 累加当前批次的损失数值(item()提取张量的纯数值)
            
            # ========== 批次级日志:每100个批次打印一次进度 ==========
            if (batch_idx + 1) % 100 == 0:
                print(f"  批次进度 [{batch_idx+1}/{total_batch}]  当前批次损失: {loss.item():.4f}")
        
        # 计算当前epoch的平均损失
        avg_loss = epoch_loss / len(train_loader)
        losses.append(avg_loss)
        
        # ========== 轮次级日志:每轮结束打印平均损失 ==========
        print(f"  第 {epoch+1}/{epochs} 轮训练完成 | 平均损失: {avg_loss:.4f}")
    
    # 返回所有epoch的平均损失
    return losses

# 设置4组不同的批次大小
batch_sizes = [32, 64, 128, 256]
# 基准初始学习率(以batch_size=32为基准)
initial_lr = 0.01
# 按照「线性缩放规则」计算对应批次大小的学习率:批次大小扩大k倍,学习率也扩大k倍
learning_rates = [initial_lr * (batch_size / 32) for batch_size in batch_sizes]

# 打印输出所有学习率数值
print("="*60)
print("各批次大小对应的学习率:")
for bs, lr in zip(batch_sizes, learning_rates):
    print(f"批次大小 {bs:3d}  →  学习率 {lr:.4f}")
print("="*60)

# 用不同批次大小+对应学习率分别训练模型,保存所有结果
losses_dict = {}  # 字典存储:键是配置标签,值是对应的损失列表
epochs = 5  # 统一训练5轮

# ========== 实验组级日志:每组训练开始/结束都打印提示 ==========
for batch_size, lr in zip(batch_sizes, learning_rates):
    print(f"\n>>> 开始训练:批次大小 = {batch_size}, 学习率 = {lr:.4f}")
    print("-" * 50)
    losses = train_model(batch_size, lr, epochs)
    losses_dict[f'批次大小: {batch_size}, 学习率: {lr:.4f}'] = losses
    print("-" * 50)
    print(f"<<< 本组训练完成:批次大小 = {batch_size}, 学习率 = {lr:.4f}")

# 绘制训练损失对比曲线
plt.figure(figsize=(12, 8))  # 设置画布尺寸
# 遍历损失字典,逐条绘制损失曲线并添加标签
for label, losses in losses_dict.items():
    plt.plot(losses, label=label)
plt.xlabel('训练轮数(Epochs)', fontsize=12)  # x轴中文标签
plt.ylabel('训练损失值(Loss)', fontsize=12)   # y轴中文标签
plt.title('不同批次大小与学习率下的训练损失对比', fontsize=14)  # 中文标题
plt.legend(fontsize=11)  # 显示图例
plt.grid(True, alpha=0.3)  # 显示网格线,调低透明度更美观
plt.show()  # 弹出并展示图表
plt.savefig('loss_curve.png', dpi=150, bbox_inches='tight')
print("损失曲线已保存为 loss_curve.png")
bash 复制代码
============================================================
各批次大小对应的学习率:
批次大小  32  →  学习率 0.0100
批次大小  64  →  学习率 0.0200
批次大小 128  →  学习率 0.0400
批次大小 256  →  学习率 0.0800
============================================================

>>> 开始训练:批次大小 = 32, 学习率 = 0.0100
--------------------------------------------------
  批次进度 [100/1875]  当前批次损失: 1.6352
  批次进度 [200/1875]  当前批次损失: 1.0548
  批次进度 [300/1875]  当前批次损失: 0.8466
  批次进度 [400/1875]  当前批次损失: 0.7757
  批次进度 [500/1875]  当前批次损失: 0.7975
  批次进度 [600/1875]  当前批次损失: 0.3258
  批次进度 [700/1875]  当前批次损失: 0.3457
  批次进度 [800/1875]  当前批次损失: 0.2457
  批次进度 [900/1875]  当前批次损失: 0.3036
  批次进度 [1000/1875]  当前批次损失: 0.6772
  批次进度 [1100/1875]  当前批次损失: 0.3987
  批次进度 [1200/1875]  当前批次损失: 0.3310
  批次进度 [1300/1875]  当前批次损失: 0.3145
  批次进度 [1400/1875]  当前批次损失: 0.3355
  批次进度 [1500/1875]  当前批次损失: 0.7144
  批次进度 [1600/1875]  当前批次损失: 0.3756
  批次进度 [1700/1875]  当前批次损失: 0.3367
  批次进度 [1800/1875]  当前批次损失: 0.2308
  第 1/5 轮训练完成 | 平均损失: 0.5718
  批次进度 [100/1875]  当前批次损失: 0.3595
  批次进度 [200/1875]  当前批次损失: 0.1613
  批次进度 [300/1875]  当前批次损失: 0.3172
  批次进度 [400/1875]  当前批次损失: 0.2900
  批次进度 [500/1875]  当前批次损失: 0.1812
  批次进度 [600/1875]  当前批次损失: 0.3495
  批次进度 [700/1875]  当前批次损失: 0.3604
  批次进度 [800/1875]  当前批次损失: 0.5645
  批次进度 [900/1875]  当前批次损失: 0.0936
  批次进度 [1000/1875]  当前批次损失: 0.4675
  批次进度 [1100/1875]  当前批次损失: 0.2053
  批次进度 [1200/1875]  当前批次损失: 0.1359
  批次进度 [1300/1875]  当前批次损失: 0.2548
  批次进度 [1400/1875]  当前批次损失: 0.3273
  批次进度 [1500/1875]  当前批次损失: 0.3306
  批次进度 [1600/1875]  当前批次损失: 0.1869
  批次进度 [1700/1875]  当前批次损失: 0.1657
  批次进度 [1800/1875]  当前批次损失: 0.2237
  第 2/5 轮训练完成 | 平均损失: 0.3114
  批次进度 [100/1875]  当前批次损失: 0.2830
  批次进度 [200/1875]  当前批次损失: 0.3480
  批次进度 [300/1875]  当前批次损失: 0.2826
  批次进度 [400/1875]  当前批次损失: 0.1801
  批次进度 [500/1875]  当前批次损失: 0.2994
  批次进度 [600/1875]  当前批次损失: 0.1535
  批次进度 [700/1875]  当前批次损失: 0.3579
  批次进度 [800/1875]  当前批次损失: 0.1798
  批次进度 [900/1875]  当前批次损失: 0.1556
  批次进度 [1000/1875]  当前批次损失: 0.2362
  批次进度 [1100/1875]  当前批次损失: 0.1987
  批次进度 [1200/1875]  当前批次损失: 0.0641
  批次进度 [1300/1875]  当前批次损失: 0.2014
  批次进度 [1400/1875]  当前批次损失: 0.5216
  批次进度 [1500/1875]  当前批次损失: 0.3150
  批次进度 [1600/1875]  当前批次损失: 0.3834
  批次进度 [1700/1875]  当前批次损失: 0.1556
  批次进度 [1800/1875]  当前批次损失: 0.2070
  第 3/5 轮训练完成 | 平均损失: 0.2651
  批次进度 [100/1875]  当前批次损失: 0.2428
  批次进度 [200/1875]  当前批次损失: 0.3619
  批次进度 [300/1875]  当前批次损失: 0.1764
  批次进度 [400/1875]  当前批次损失: 0.1087
  批次进度 [500/1875]  当前批次损失: 0.3758
  批次进度 [600/1875]  当前批次损失: 0.1345
  批次进度 [700/1875]  当前批次损失: 0.0807
  批次进度 [800/1875]  当前批次损失: 0.1328
  批次进度 [900/1875]  当前批次损失: 0.1897
  批次进度 [1000/1875]  当前批次损失: 0.3104
  批次进度 [1100/1875]  当前批次损失: 0.4025
  批次进度 [1200/1875]  当前批次损失: 0.1703
  批次进度 [1300/1875]  当前批次损失: 0.2997
  批次进度 [1400/1875]  当前批次损失: 0.1920
  批次进度 [1500/1875]  当前批次损失: 0.3362
  批次进度 [1600/1875]  当前批次损失: 0.1902
  批次进度 [1700/1875]  当前批次损失: 0.0784
  批次进度 [1800/1875]  当前批次损失: 0.0549
  第 4/5 轮训练完成 | 平均损失: 0.2301
  批次进度 [100/1875]  当前批次损失: 0.1082
  批次进度 [200/1875]  当前批次损失: 0.0583
  批次进度 [300/1875]  当前批次损失: 0.1302
  批次进度 [400/1875]  当前批次损失: 0.0337
  批次进度 [500/1875]  当前批次损失: 0.0824
  批次进度 [600/1875]  当前批次损失: 0.4375
  批次进度 [700/1875]  当前批次损失: 0.1050
  批次进度 [800/1875]  当前批次损失: 0.0712
  批次进度 [900/1875]  当前批次损失: 0.1036
  批次进度 [1000/1875]  当前批次损失: 0.1185
  批次进度 [1100/1875]  当前批次损失: 0.1109
  批次进度 [1200/1875]  当前批次损失: 0.2373
  批次进度 [1300/1875]  当前批次损失: 0.1645
  批次进度 [1400/1875]  当前批次损失: 0.0766
  批次进度 [1500/1875]  当前批次损失: 0.0966
  批次进度 [1600/1875]  当前批次损失: 0.2200
  批次进度 [1700/1875]  当前批次损失: 0.0352
  批次进度 [1800/1875]  当前批次损失: 0.0739
  第 5/5 轮训练完成 | 平均损失: 0.2021
--------------------------------------------------
<<< 本组训练完成:批次大小 = 32, 学习率 = 0.0100

>>> 开始训练:批次大小 = 64, 学习率 = 0.0200
--------------------------------------------------
  批次进度 [100/938]  当前批次损失: 0.9697
  批次进度 [200/938]  当前批次损失: 0.6736
  批次进度 [300/938]  当前批次损失: 0.4151
  批次进度 [400/938]  当前批次损失: 0.4282
  批次进度 [500/938]  当前批次损失: 0.2553
  批次进度 [600/938]  当前批次损失: 0.1728
  批次进度 [700/938]  当前批次损失: 0.5420
  批次进度 [800/938]  当前批次损失: 0.4763
  批次进度 [900/938]  当前批次损失: 0.4012
  第 1/5 轮训练完成 | 平均损失: 0.5583
  批次进度 [100/938]  当前批次损失: 0.3009
  批次进度 [200/938]  当前批次损失: 0.3271
  批次进度 [300/938]  当前批次损失: 0.2340
  批次进度 [400/938]  当前批次损失: 0.2521
  批次进度 [500/938]  当前批次损失: 0.2633
  批次进度 [600/938]  当前批次损失: 0.2147
  批次进度 [700/938]  当前批次损失: 0.2542
  批次进度 [800/938]  当前批次损失: 0.3538
  批次进度 [900/938]  当前批次损失: 0.1956
  第 2/5 轮训练完成 | 平均损失: 0.3165
  批次进度 [100/938]  当前批次损失: 0.1604
  批次进度 [200/938]  当前批次损失: 0.2603
  批次进度 [300/938]  当前批次损失: 0.1824
  批次进度 [400/938]  当前批次损失: 0.3237
  批次进度 [500/938]  当前批次损失: 0.2117
  批次进度 [600/938]  当前批次损失: 0.4976
  批次进度 [700/938]  当前批次损失: 0.2220
  批次进度 [800/938]  当前批次损失: 0.1152
  批次进度 [900/938]  当前批次损失: 0.2298
  第 3/5 轮训练完成 | 平均损失: 0.2715
  批次进度 [100/938]  当前批次损失: 0.2721
  批次进度 [200/938]  当前批次损失: 0.2150
  批次进度 [300/938]  当前批次损失: 0.3241
  批次进度 [400/938]  当前批次损失: 0.3123
  批次进度 [500/938]  当前批次损失: 0.1980
  批次进度 [600/938]  当前批次损失: 0.3012
  批次进度 [700/938]  当前批次损失: 0.3659
  批次进度 [800/938]  当前批次损失: 0.2001
  批次进度 [900/938]  当前批次损失: 0.1783
  第 4/5 轮训练完成 | 平均损失: 0.2366
  批次进度 [100/938]  当前批次损失: 0.2924
  批次进度 [200/938]  当前批次损失: 0.1481
  批次进度 [300/938]  当前批次损失: 0.3940
  批次进度 [400/938]  当前批次损失: 0.3254
  批次进度 [500/938]  当前批次损失: 0.0584
  批次进度 [600/938]  当前批次损失: 0.2413
  批次进度 [700/938]  当前批次损失: 0.1744
  批次进度 [800/938]  当前批次损失: 0.1823
  批次进度 [900/938]  当前批次损失: 0.1940
  第 5/5 轮训练完成 | 平均损失: 0.2074
--------------------------------------------------
<<< 本组训练完成:批次大小 = 64, 学习率 = 0.0200

>>> 开始训练:批次大小 = 128, 学习率 = 0.0400
--------------------------------------------------
  批次进度 [100/469]  当前批次损失: 0.6556
  批次进度 [200/469]  当前批次损失: 0.5254
  批次进度 [300/469]  当前批次损失: 0.4676
  批次进度 [400/469]  当前批次损失: 0.3694
  第 1/5 轮训练完成 | 平均损失: 0.5786
  批次进度 [100/469]  当前批次损失: 0.4052
  批次进度 [200/469]  当前批次损失: 0.3785
  批次进度 [300/469]  当前批次损失: 0.2894
  批次进度 [400/469]  当前批次损失: 0.3029
  第 2/5 轮训练完成 | 平均损失: 0.3180
  批次进度 [100/469]  当前批次损失: 0.2925
  批次进度 [200/469]  当前批次损失: 0.3477
  批次进度 [300/469]  当前批次损失: 0.2321
  批次进度 [400/469]  当前批次损失: 0.2878
  第 3/5 轮训练完成 | 平均损失: 0.2698
  批次进度 [100/469]  当前批次损失: 0.2638
  批次进度 [200/469]  当前批次损失: 0.3178
  批次进度 [300/469]  当前批次损失: 0.2517
  批次进度 [400/469]  当前批次损失: 0.1637
  第 4/5 轮训练完成 | 平均损失: 0.2343
  批次进度 [100/469]  当前批次损失: 0.1985
  批次进度 [200/469]  当前批次损失: 0.2089
  批次进度 [300/469]  当前批次损失: 0.1935
  批次进度 [400/469]  当前批次损失: 0.1959
  第 5/5 轮训练完成 | 平均损失: 0.2060
--------------------------------------------------
<<< 本组训练完成:批次大小 = 128, 学习率 = 0.0400

>>> 开始训练:批次大小 = 256, 学习率 = 0.0800
--------------------------------------------------
  批次进度 [100/235]  当前批次损失: 0.3574
  批次进度 [200/235]  当前批次损失: 0.3908
  第 1/5 轮训练完成 | 平均损失: 0.6198
  批次进度 [100/235]  当前批次损失: 0.2594
  批次进度 [200/235]  当前批次损失: 0.2785
  第 2/5 轮训练完成 | 平均损失: 0.3185
  批次进度 [100/235]  当前批次损失: 0.3030
  批次进度 [200/235]  当前批次损失: 0.2215
  第 3/5 轮训练完成 | 平均损失: 0.2635
  批次进度 [100/235]  当前批次损失: 0.2253
  批次进度 [200/235]  当前批次损失: 0.1440
  第 4/5 轮训练完成 | 平均损失: 0.2254
  批次进度 [100/235]  当前批次损失: 0.2538
  批次进度 [200/235]  当前批次损失: 0.2317
  第 5/5 轮训练完成 | 平均损失: 0.1971
--------------------------------------------------
<<< 本组训练完成:批次大小 = 256, 学习率 = 0.0800

AdamW的情况下 不同批次大小与学习率下的训练损失对比

bash 复制代码
批次大小  32  →  学习率 0.0100
批次大小  64  →  学习率 0.0200
批次大小 128  →  学习率 0.0400
批次大小 256  →  学习率 0.0800
python 复制代码
# 导入PyTorch核心库、神经网络模块与优化器模块
import torch
import torch.nn as nn
import torch.optim as optim
# 导入torchvision提供的数据集和数据预处理工具
from torchvision import datasets, transforms
# 导入数据加载器,用于按批次批量加载训练数据
from torch.utils.data import DataLoader
# 导入matplotlib绘图库,用于可视化训练损失曲线
import matplotlib.pyplot as plt
# ========== 导入时间模块,用于生成不重复的文件名 ==========
import datetime

# 配置matplotlib中文显示
plt.rcParams['font.sans-serif'] = ['AR PL UMing CN']
plt.rcParams['axes.unicode_minus'] = False    # 解决坐标轴负号显示为方块的问题

# 定义一个简单的全连接神经网络类,继承PyTorch的nn.Module基类
class SimpleNN(nn.Module):
    def __init__(self):
        # 调用父类nn.Module的初始化方法
        super(SimpleNN, self).__init__()
        # 第一个全连接层:输入维度是28*28(MNIST单张图片的像素总数),输出维度128(隐藏层神经元数量)
        self.fc1 = nn.Linear(28*28, 128)
        # 第二个全连接层:输入维度128,输出维度10(对应MNIST的10个数字类别:0~9)
        self.fc2 = nn.Linear(128, 10)
    
    # 定义前向传播的计算逻辑
    def forward(self, x):
        # 将二维图片张量展平为一维向量:-1表示自动匹配批次大小,28*28是单张图片的特征维度
        x = x.view(-1, 28*28)
        # 第一个全连接层的输出经过ReLU激活函数,引入非线性变换
        x = torch.relu(self.fc1(x))
        # 第二个全连接层输出最终的分类得分(也叫logits)
        x = self.fc2(x)
        return x

# 定义数据预处理的流水线
transform = transforms.Compose([
    transforms.ToTensor(),  # 将图片转为PyTorch张量,同时把像素值从[0,255]缩放到[0,1]区间
    transforms.Normalize((0.5,), (0.5,))  # 单通道归一化:均值0.5、标准差0.5,将数值映射到[-1, 1]区间
])

# 加载MNIST手写数字训练数据集
# root:数据本地存储路径;train=True表示加载训练集;download=True表示本地无数据时自动下载
# transform:对每张图片应用上面定义的预处理操作
train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform)

# 定义模型训练函数
# 输入:批次大小batch_size、学习率learning_rate、训练轮数epochs(默认1轮)
# 输出:每一轮训练的平均损失列表
def train_model(batch_size, learning_rate, epochs=1):
    # 创建数据加载器:按指定批次大小加载数据,shuffle=True表示每轮训练打乱数据顺序
    train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)
    # 实例化我们定义的简单神经网络
    model = SimpleNN()
    # 定义损失函数:交叉熵损失,是多分类任务的标准损失函数
    criterion = nn.CrossEntropyLoss()
    # ========== 优化器从SGD更换为AdamW ==========
    optimizer = optim.AdamW(model.parameters(), lr=learning_rate)
    
    losses = []  # 用于存储每一个epoch的平均训练损失
    total_batch = len(train_loader)  # 总批次数量,用于显示进度
    
    for epoch in range(epochs):
        epoch_loss = 0  # 记录当前epoch的总损失
        # 逐批次遍历训练数据,用enumerate获取批次索引
        for batch_idx, (data, target) in enumerate(train_loader):
            optimizer.zero_grad()  # 清空上一步计算的梯度,避免梯度累积
            output = model(data)  # 前向传播:输入图片数据,得到模型预测结果
            loss = criterion(output, target)  # 计算预测结果与真实标签的损失值
            loss.backward()  # 反向传播:自动计算损失对每个网络参数的梯度
            optimizer.step()  # 参数更新:根据梯度和优化器规则调整网络参数
            epoch_loss += loss.item()  # 累加当前批次的损失数值(item()提取张量的纯数值)
            
            # 批次级日志:每100个批次打印一次进度
            if (batch_idx + 1) % 100 == 0:
                print(f"  批次进度 [{batch_idx+1}/{total_batch}]  当前批次损失: {loss.item():.4f}")
        
        # 计算当前epoch的平均损失
        avg_loss = epoch_loss / len(train_loader)
        losses.append(avg_loss)
        
        # 轮次级日志:每轮结束打印平均损失
        print(f"  第 {epoch+1}/{epochs} 轮训练完成 | 平均损失: {avg_loss:.4f}")
    
    # 返回所有epoch的平均损失
    return losses

# 设置4组不同的批次大小
batch_sizes = [32, 64, 128, 256]
# 基准初始学习率(以batch_size=32为基准)
initial_lr = 0.01
# 按照「线性缩放规则」计算对应批次大小的学习率:批次大小扩大k倍,学习率也扩大k倍
learning_rates = [initial_lr * (batch_size / 32) for batch_size in batch_sizes]

# 打印输出所有学习率数值
print("="*60)
print("各批次大小对应的学习率:")
for bs, lr in zip(batch_sizes, learning_rates):
    print(f"批次大小 {bs:3d}  →  学习率 {lr:.4f}")
print("="*60)

# 用不同批次大小+对应学习率分别训练模型,保存所有结果
losses_dict = {}  # 字典存储:键是配置标签,值是对应的损失列表
epochs = 5  # 统一训练5轮

# 实验组级日志:每组训练开始/结束都打印提示
for batch_size, lr in zip(batch_sizes, learning_rates):
    print(f"\n>>> 开始训练:批次大小 = {batch_size}, 学习率 = {lr:.4f}")
    print("-" * 50)
    losses = train_model(batch_size, lr, epochs)
    losses_dict[f'批次大小: {batch_size}, 学习率: {lr:.4f}'] = losses
    print("-" * 50)
    print(f"<<< 本组训练完成:批次大小 = {batch_size}, 学习率 = {lr:.4f}")

# 绘制训练损失对比曲线
plt.figure(figsize=(12, 8))  # 设置画布尺寸
# 遍历损失字典,逐条绘制损失曲线并添加标签
for label, losses in losses_dict.items():
    plt.plot(losses, label=label)
plt.xlabel('训练轮数(Epochs)', fontsize=12)  # x轴中文标签
plt.ylabel('训练损失值(Loss)', fontsize=12)   # y轴中文标签
plt.title('不同批次大小与学习率下的训练损失对比', fontsize=14)  # 中文标题
plt.legend(fontsize=11)  # 显示图例
plt.grid(True, alpha=0.3)  # 显示网格线,调低透明度更美观

# ========== 先保存再显示 + 时间戳命名,避免覆盖历史图片 ==========
# 生成带年月日时分秒的时间戳,确保每次运行文件名唯一
timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
save_filename = f'loss_curve_{timestamp}.png'
plt.savefig(save_filename, dpi=150, bbox_inches='tight')
print(f"\n损失曲线已保存为:{save_filename}")

plt.show()  # 最后再弹出展示图片
bash 复制代码
============================================================
各批次大小对应的学习率:
批次大小  32  →  学习率 0.0100
批次大小  64  →  学习率 0.0200
批次大小 128  →  学习率 0.0400
批次大小 256  →  学习率 0.0800
============================================================

>>> 开始训练:批次大小 = 32, 学习率 = 0.0100
--------------------------------------------------
  批次进度 [100/1875]  当前批次损失: 0.7074
  批次进度 [200/1875]  当前批次损失: 0.4518
  批次进度 [300/1875]  当前批次损失: 0.3376
  批次进度 [400/1875]  当前批次损失: 0.4451
  批次进度 [500/1875]  当前批次损失: 0.7398
  批次进度 [600/1875]  当前批次损失: 0.4991
  批次进度 [700/1875]  当前批次损失: 0.2943
  批次进度 [800/1875]  当前批次损失: 0.3946
  批次进度 [900/1875]  当前批次损失: 0.2536
  批次进度 [1000/1875]  当前批次损失: 0.1256
  批次进度 [1100/1875]  当前批次损失: 0.3620
  批次进度 [1200/1875]  当前批次损失: 0.0940
  批次进度 [1300/1875]  当前批次损失: 0.3555
  批次进度 [1400/1875]  当前批次损失: 0.5256
  批次进度 [1500/1875]  当前批次损失: 0.4057
  批次进度 [1600/1875]  当前批次损失: 0.1910
  批次进度 [1700/1875]  当前批次损失: 0.1254
  批次进度 [1800/1875]  当前批次损失: 0.3020
  第 1/5 轮训练完成 | 平均损失: 0.4332
  批次进度 [100/1875]  当前批次损失: 0.7709
  批次进度 [200/1875]  当前批次损失: 0.2387
  批次进度 [300/1875]  当前批次损失: 0.1073
  批次进度 [400/1875]  当前批次损失: 0.1168
  批次进度 [500/1875]  当前批次损失: 0.1148
  批次进度 [600/1875]  当前批次损失: 0.0997
  批次进度 [700/1875]  当前批次损失: 0.1523
  批次进度 [800/1875]  当前批次损失: 0.3041
  批次进度 [900/1875]  当前批次损失: 0.1091
  批次进度 [1000/1875]  当前批次损失: 0.5886
  批次进度 [1100/1875]  当前批次损失: 0.2839
  批次进度 [1200/1875]  当前批次损失: 0.1292
  批次进度 [1300/1875]  当前批次损失: 0.4519
  批次进度 [1400/1875]  当前批次损失: 0.3524
  批次进度 [1500/1875]  当前批次损失: 0.5323
  批次进度 [1600/1875]  当前批次损失: 0.3373
  批次进度 [1700/1875]  当前批次损失: 0.3269
  批次进度 [1800/1875]  当前批次损失: 0.0967
  第 2/5 轮训练完成 | 平均损失: 0.3343
  批次进度 [100/1875]  当前批次损失: 0.0901
  批次进度 [200/1875]  当前批次损失: 1.0776
  批次进度 [300/1875]  当前批次损失: 0.7595
  批次进度 [400/1875]  当前批次损失: 0.3049
  批次进度 [500/1875]  当前批次损失: 0.1472
  批次进度 [600/1875]  当前批次损失: 0.2124
  批次进度 [700/1875]  当前批次损失: 0.4502
  批次进度 [800/1875]  当前批次损失: 0.2131
  批次进度 [900/1875]  当前批次损失: 0.2419
  批次进度 [1000/1875]  当前批次损失: 0.4421
  批次进度 [1100/1875]  当前批次损失: 0.0684
  批次进度 [1200/1875]  当前批次损失: 0.2464
  批次进度 [1300/1875]  当前批次损失: 0.2520
  批次进度 [1400/1875]  当前批次损失: 0.0985
  批次进度 [1500/1875]  当前批次损失: 0.2886
  批次进度 [1600/1875]  当前批次损失: 0.4231
  批次进度 [1700/1875]  当前批次损失: 0.6077
  批次进度 [1800/1875]  当前批次损失: 0.5246
  第 3/5 轮训练完成 | 平均损失: 0.3067
  批次进度 [100/1875]  当前批次损失: 0.3217
  批次进度 [200/1875]  当前批次损失: 0.2226
  批次进度 [300/1875]  当前批次损失: 0.2225
  批次进度 [400/1875]  当前批次损失: 0.0406
  批次进度 [500/1875]  当前批次损失: 0.3473
  批次进度 [600/1875]  当前批次损失: 0.2913
  批次进度 [700/1875]  当前批次损失: 0.5247
  批次进度 [800/1875]  当前批次损失: 0.3245
  批次进度 [900/1875]  当前批次损失: 0.2462
  批次进度 [1000/1875]  当前批次损失: 0.0996
  批次进度 [1100/1875]  当前批次损失: 0.2958
  批次进度 [1200/1875]  当前批次损失: 0.2020
  批次进度 [1300/1875]  当前批次损失: 0.3223
  批次进度 [1400/1875]  当前批次损失: 0.1879
  批次进度 [1500/1875]  当前批次损失: 0.5078
  批次进度 [1600/1875]  当前批次损失: 0.3093
  批次进度 [1700/1875]  当前批次损失: 0.3731
  批次进度 [1800/1875]  当前批次损失: 0.2966
  第 4/5 轮训练完成 | 平均损失: 0.2973
  批次进度 [100/1875]  当前批次损失: 0.0335
  批次进度 [200/1875]  当前批次损失: 0.7626
  批次进度 [300/1875]  当前批次损失: 0.0817
  批次进度 [400/1875]  当前批次损失: 0.0820
  批次进度 [500/1875]  当前批次损失: 0.2879
  批次进度 [600/1875]  当前批次损失: 0.0996
  批次进度 [700/1875]  当前批次损失: 0.2235
  批次进度 [800/1875]  当前批次损失: 0.1550
  批次进度 [900/1875]  当前批次损失: 0.2963
  批次进度 [1000/1875]  当前批次损失: 0.3390
  批次进度 [1100/1875]  当前批次损失: 0.3379
  批次进度 [1200/1875]  当前批次损失: 0.2295
  批次进度 [1300/1875]  当前批次损失: 0.1238
  批次进度 [1400/1875]  当前批次损失: 0.1198
  批次进度 [1500/1875]  当前批次损失: 0.2774
  批次进度 [1600/1875]  当前批次损失: 0.6945
  批次进度 [1700/1875]  当前批次损失: 0.3286
  批次进度 [1800/1875]  当前批次损失: 0.2767
  第 5/5 轮训练完成 | 平均损失: 0.3083
--------------------------------------------------
<<< 本组训练完成:批次大小 = 32, 学习率 = 0.0100

>>> 开始训练:批次大小 = 64, 学习率 = 0.0200
--------------------------------------------------
  批次进度 [100/938]  当前批次损失: 0.5482
  批次进度 [200/938]  当前批次损失: 0.3401
  批次进度 [300/938]  当前批次损失: 0.4672
  批次进度 [400/938]  当前批次损失: 0.4406
  批次进度 [500/938]  当前批次损失: 0.6053
  批次进度 [600/938]  当前批次损失: 0.7528
  批次进度 [700/938]  当前批次损失: 0.7292
  批次进度 [800/938]  当前批次损失: 0.5596
  批次进度 [900/938]  当前批次损失: 0.4549
  第 1/5 轮训练完成 | 平均损失: 0.5903
  批次进度 [100/938]  当前批次损失: 0.3569
  批次进度 [200/938]  当前批次损失: 0.4743
  批次进度 [300/938]  当前批次损失: 0.1535
  批次进度 [400/938]  当前批次损失: 0.4414
  批次进度 [500/938]  当前批次损失: 0.4942
  批次进度 [600/938]  当前批次损失: 0.5983
  批次进度 [700/938]  当前批次损失: 0.4599
  批次进度 [800/938]  当前批次损失: 0.2854
  批次进度 [900/938]  当前批次损失: 0.6797
  第 2/5 轮训练完成 | 平均损失: 0.4188
  批次进度 [100/938]  当前批次损失: 0.4070
  批次进度 [200/938]  当前批次损失: 0.4916
  批次进度 [300/938]  当前批次损失: 0.3676
  批次进度 [400/938]  当前批次损失: 0.2911
  批次进度 [500/938]  当前批次损失: 0.3157
  批次进度 [600/938]  当前批次损失: 0.3064
  批次进度 [700/938]  当前批次损失: 0.6298
  批次进度 [800/938]  当前批次损失: 0.5509
  批次进度 [900/938]  当前批次损失: 0.3873
  第 3/5 轮训练完成 | 平均损失: 0.4029
  批次进度 [100/938]  当前批次损失: 0.3332
  批次进度 [200/938]  当前批次损失: 0.2279
  批次进度 [300/938]  当前批次损失: 0.3602
  批次进度 [400/938]  当前批次损失: 0.3223
  批次进度 [500/938]  当前批次损失: 0.6749
  批次进度 [600/938]  当前批次损失: 0.3604
  批次进度 [700/938]  当前批次损失: 0.2739
  批次进度 [800/938]  当前批次损失: 0.6317
  批次进度 [900/938]  当前批次损失: 0.4104
  第 4/5 轮训练完成 | 平均损失: 0.4010
  批次进度 [100/938]  当前批次损失: 0.4399
  批次进度 [200/938]  当前批次损失: 0.6083
  批次进度 [300/938]  当前批次损失: 0.2271
  批次进度 [400/938]  当前批次损失: 0.2572
  批次进度 [500/938]  当前批次损失: 0.2920
  批次进度 [600/938]  当前批次损失: 0.5541
  批次进度 [700/938]  当前批次损失: 0.2858
  批次进度 [800/938]  当前批次损失: 0.4519
  批次进度 [900/938]  当前批次损失: 0.1556
  第 5/5 轮训练完成 | 平均损失: 0.4042
--------------------------------------------------
<<< 本组训练完成:批次大小 = 64, 学习率 = 0.0200

>>> 开始训练:批次大小 = 128, 学习率 = 0.0400
--------------------------------------------------
  批次进度 [100/469]  当前批次损失: 1.9822
  批次进度 [200/469]  当前批次损失: 1.8678
  批次进度 [300/469]  当前批次损失: 1.7571
  批次进度 [400/469]  当前批次损失: 1.7680
  第 1/5 轮训练完成 | 平均损失: 2.2750
  批次进度 [100/469]  当前批次损失: 1.6298
  批次进度 [200/469]  当前批次损失: 1.6884
  批次进度 [300/469]  当前批次损失: 1.6939
  批次进度 [400/469]  当前批次损失: 1.5797
  第 2/5 轮训练完成 | 平均损失: 1.6785
  批次进度 [100/469]  当前批次损失: 1.6625
  批次进度 [200/469]  当前批次损失: 1.5799
  批次进度 [300/469]  当前批次损失: 1.7810
  批次进度 [400/469]  当前批次损失: 1.6197
  第 3/5 轮训练完成 | 平均损失: 1.6477
  批次进度 [100/469]  当前批次损失: 1.6437
  批次进度 [200/469]  当前批次损失: 1.7153
  批次进度 [300/469]  当前批次损失: 1.6571
  批次进度 [400/469]  当前批次损失: 1.5806
  第 4/5 轮训练完成 | 平均损失: 1.6481
  批次进度 [100/469]  当前批次损失: 1.5170
  批次进度 [200/469]  当前批次损失: 1.7285
  批次进度 [300/469]  当前批次损失: 1.6151
  批次进度 [400/469]  当前批次损失: 1.6410
  第 5/5 轮训练完成 | 平均损失: 1.6465
--------------------------------------------------
<<< 本组训练完成:批次大小 = 128, 学习率 = 0.0400

>>> 开始训练:批次大小 = 256, 学习率 = 0.0800
--------------------------------------------------
  批次进度 [100/235]  当前批次损失: 1.9311
  批次进度 [200/235]  当前批次损失: 1.7620
  第 1/5 轮训练完成 | 平均损失: 5.0282
  批次进度 [100/235]  当前批次损失: 1.5785
  批次进度 [200/235]  当前批次损失: 1.3281
  第 2/5 轮训练完成 | 平均损失: 1.4718
  批次进度 [100/235]  当前批次损失: 1.3532
  批次进度 [200/235]  当前批次损失: 1.3954
  第 3/5 轮训练完成 | 平均损失: 1.3255
  批次进度 [100/235]  当前批次损失: 1.1862
  批次进度 [200/235]  当前批次损失: 1.2859
  第 4/5 轮训练完成 | 平均损失: 1.2770
  批次进度 [100/235]  当前批次损失: 1.1989
  批次进度 [200/235]  当前批次损失: 1.3160
  第 5/5 轮训练完成 | 平均损失: 1.2684
--------------------------------------------------
<<< 本组训练完成:批次大小 = 256, 学习率 = 0.0800

AdamW的情况下 不同批次大小与相同的学习率下的训练损失对比

各批次大小对应的学习率:

bash 复制代码
批次大小  32  →  学习率 0.0100
批次大小  64  →  学习率 0.0100
批次大小 128  →  学习率 0.0100
批次大小 256  →  学习率 0.0100
相关推荐
知识分享小能手14 分钟前
概理论与数理统计学习教程,从入门到精通,平稳随机过程(27)
学习·数据挖掘·概率论
2601_9672642826 分钟前
极客时间 AI Agent 全栈工程师训练营,全套学习资源齐全
人工智能·学习
薛定e的猫咪30 分钟前
(NeurIPS 2022)GraphGPS:MPNN 与全局注意力的融合之道
人工智能·深度学习·学习·算法
Z59981784133 分钟前
c#软件开发学习笔记--WPF(Canvas、数据绑定、MVVM模式、值转换器)
笔记·学习·c#
茯苓gao33 分钟前
从零开发 EtherCAT 主站(六):SOEM 初始化流程详解,主站是如何发现所有从站的?
笔记·嵌入式硬件·学习·信息与通信
老当益壮梁奶奶1 小时前
Linux软件编程学习笔记(九):消息队列、共享内存与信号灯详解
linux·c语言·笔记·学习
MartinYeung51 小时前
[论文学习]激活差异揭示后门:SAE架构对比研究
人工智能·学习·架构
小雪崩1 小时前
嵌入式学习 day35:进程间通信-消息队列、共享空间及有名信号量
linux·c语言·学习
一条破秋裤2 小时前
STM32 学习笔记:STM32 简介与学习路线
笔记·stm32·学习