模型训练 不同批次大小与学习率下的训练损失对比 (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
