train/test 函数 · 主循环
一、train:训练一个 epoch 的标准五步
食物分类.py ------ 训练函数
def train(dataloader, model, loss_fn, optimizer):
model.train()
batch_size_num = 1
for X, y in dataloader:
X, y = X.to(device), y.to(device)
pred = model.forward(X)
loss = loss_fn(pred, y)
optimizer.zero_grad()
loss.backward()
optimizer.step()
loss_value = loss.item()
print(f"batch {batch_size_num} loss: {loss_value:.7f}")
batch_size_num += 1
函数作用:遍历整个训练集一次(一个 epoch),每个 batch 做五件事:
• X, y = X.to(device), y.to(device):数据和标签搬到 GPU/CPU;
• model.forward(X):前向传播,得到每张图 20 个类别分数;
• loss_fn(pred, y):交叉熵损失,衡量预测与真实标签的差距;
• optimizer.zero_grad():清空上一轮梯度(不清空会累加,新手最常见错误);
• loss.backward() 反向传播求梯度,optimizer.step() 更新参数。
两个细节:model.train() 切到训练模式(BatchNorm、Dropout 等层在训练/测试下行为不同,必须切换);loss.item() 把张量转成普通 Python 数字用于打印。
二、test:评估准确率
食物分类.py ------ 测试函数
def test(dataloader, model, loss_fn):
size = len(dataloader.dataset)
num_batches = len(dataloader)
model.eval()
test_loss,correct = 0,0
with torch.no_grad():
for X, y in dataloader:
X, y = X.to(device), y.to(device)
pred = model.forward(X)
test_loss += loss_fn(pred, y).item()
correct += (pred.argmax(1) == y).type(torch.float).sum().item()
test_loss /= num_batches
correct /= size
print(f"Test result: \n Accuracy: {(100*correct)}%, Avg loss: {test_loss}")
model.eval() 切到评估模式;torch.no_grad() 关闭梯度计算------评估不需要反向传播,关掉省显存/内存并提速。重点代码是这一行:
correct += (pred.argmax(1) == y).type(torch.float).sum().item()
pred.argmax(1) 取每行分数最大的类别下标作预测,与真实标签 y 逐位比较,相等转 1 求和就是预测正确张数;除以总数 size 乘 100 即准确率百分比。
三、主循环与小结
loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
epochs = 20
for epoch in range(epochs):
print(f"Epoch {epoch+1}")
train(train_loader, model, loss_fn, optimizer)
print("Finished Training")
test(test_loader, model, loss_fn)
CrossEntropyLoss 交叉熵是分类任务的标准损失;Adam 是自适应学习率优化器,lr=0.001 是常用初始值。每轮打印轮次,训练结束后在测试集上评估。
至此完整链路已经跑通:清单生成 → Dataset 读数据 → DataLoader 组批 → 网络前向 → 损失反向传播 → 测试评估。
数据增强
一、为什么要做数据增强
前几篇的数据集训练集只有 274 张图,平均每类不到 14 张。样本这么少,网络很容易"背答案"------训练图滚瓜烂熟,见到新图就露馅,这就是过拟合。
增强思路:训练时对每张图随机施加变换(旋转、翻转、调色......),模型每次看到的都是"同一张图的随机变体"。数据量名义上没变,见过的形态却成倍增加,泛化能力明显提升,而且不需要采集新数据。数据增强.py 就是给训练集加了这串随机变换。
二、训练集流水线逐项拆解
数据增强.py ------ 训练集增强流水线
data_transforms = {
'train':
transforms.Compose([
transforms.Resize((256, 256)),
transforms.RandomRotation(45),
transforms.CenterCrop(256),
transforms.RandomHorizontalFlip(p=0.5),
transforms.RandomVerticalFlip(p=0.5),
transforms.ColorJitter(brightness=0.2, contrast=0.1,
saturation=0.1, hue=0.1),
transforms.RandomGrayscale(p=0.1),
transforms.ToTensor(),
transforms.Normalize(mean=0.485, 0.456, 0.406,
std=0.229, 0.224, 0.225)
]),
...
}
Compose 里的变换按顺序执行,逐项看:
• Resize((256,256)):统一尺寸,流水线第一道工序;
• RandomRotation(45):随机旋转 -45°~+45°。注意:旋转后四角露黑边、对角线方向超画布,必须和下一个变换配套;
• CenterCrop(256):从中心裁回 256×256,裁掉黑边和越界部分------先旋转、再中心裁剪,尺寸才兜得住;
• RandomHorizontalFlip(p=0.5) / RandomVerticalFlip(p=0.5):各以 50% 概率左右/上下翻转,食物大多左右对称,翻转不改变类别;
• ColorJitter(brightness=0.2, contrast=0.1, saturation=0.1, hue=0.1):随机扰动亮度 ±20%、对比度 ±10%、饱和度 ±10%、色相 ±0.1,模拟不同光照和拍摄条件;
• RandomGrayscale(p=0.1):10% 概率转灰度图,逼模型不完全依赖颜色(不同品种的哈密瓜颜色差异很大);
• ToTensor():转张量,像素 0~255 → 0~1,通道 H×W×C → C×H×W;
• Normalize(mean=0.485,0.456,0.406, std=0.229,0.224,0.225):按 (x-mean)/std 逐通道标准化。这组均值方差是 ImageNet 上百万张图统计出来的社区事实标准,把输入分布拉回零均值单位方差附近,梯度更稳、收敛更快;用 ImageNet 预训练模型时输入分布必须和它训练时一致,这组值就必须用。
三、训练集增强,验证集不增强
数据增强.py ------ 验证集流水线(只有确定性变换)
'valid':
transforms.Compose([
transforms.Resize((256, 256)),
transforms.ToTensor(),
transforms.Normalize(mean=0.485, 0.456, 0.406,
std=0.229, 0.224, 0.225)
])
验证集只有 Resize + ToTensor + Normalize,没有任何随机变换 。原因:增强是为了让训练"更难、更多样",逼模型学本质特征;验证集要回答"模型真实水平如何",必须固定可复现------同一张测试图每次喂进去必须一模一样,评估才有意义。规矩:随机增强只加训练集,验证集只做确定性预处理。