1. 引言
本文围绕一段完整的 PyTorch 食物图像分类代码展开,逐段讲解从数据集整理、自定义 Dataset、数据预处理、CNN 网络搭建到训练与测试循环的完整流程。代码以食物图片二分类(train/test 目录)为背景,适合初学者理解 PyTorch 图像分类任务的基本套路。
2. 环境与依赖
代码依赖以下核心库:
- os:文件路径遍历与拼接。
- torch / torch.nn:张量计算与神经网络模块。
- numpy:数值计算(本代码中主要用于间接支持张量操作)。
- PIL(Pillow):读取图片文件。
- torch.utils.data:Dataset 与 DataLoader 数据加载工具。
- torchvision.transforms:图像预处理与张量转换。
3. 数据集目录整理:train_test_file 函数
该函数的作用是把指定目录下的图片路径和类别标签写入 txt 文件,供后续 Dataset 读取。核心逻辑如下:
python
def train_test_file(root, dir):
# 使用with自动关闭文件,避免文件泄露
with open(dir + '.txt', 'w') as file_txt:
path = os.path.join(root, dir)
dirs_cache = [] # 缓存类别文件夹,修复NameError问题
for roots, directories, files in os.walk(path):
if len(directories) != 0:
dirs_cache = directories
else:
now_dir = roots.split(os.sep) # os.sep自动适配Windows\\和Linux/
for file in files:
path_1 = os.path.join(roots, file)
print(path_1)
# 防止类别不存在报错
if now_dir[-1] in dirs_cache:
file_txt.write(path_1 + ' ' + str(dirs_cache.index(now_dir[-1])) + '\n')
逐行解析:
- with open(dir + '.txt', 'w'):以写模式打开文件,with 语句确保文件自动关闭,避免资源泄露。
- os.walk(path):递归遍历目录,返回三元组(当前目录路径、子目录列表、文件列表)。
- dirs_cache = directories:当当前目录还有子目录时,缓存类别文件夹名。这里修复了原代码中可能出现的 NameError 问题------当进入叶子目录时,directories 为空,需要用之前缓存的类别列表。
- roots.split(os.sep):按系统分隔符拆分路径,os.sep 在 Windows 下是反斜杠,在 Linux 下是正斜杠,保证跨平台兼容。
- now_dir-1 in dirs_cache:判断当前叶子目录名是否属于已知类别,防止类别不存在时报错。
- dirs_cache.index(now_dir-1):取类别在列表中的下标作为数字标签,写入 txt 文件。
最终生成的 train.txt / test.txt 每行格式为:图片绝对路径 + 空格 + 类别编号。
4. 自定义 Dataset 与 getitem 机制
代码先通过一个简单的 USE_getitem 类演示了 Python 的索引协议:
python
class USE_getitem():
def __init__(self, text):
self.text = text
def __getitem__(self, index):
result = self.text[index].upper()
return result
def __len__(self):
return len(self.text)
这个类实现了两个魔法方法:getitem 让对象支持下标访问(p1),len 让对象支持 len() 调用。这是 PyTorch Dataset 的核心协议------DataLoader 正是通过这两个方法按索引取样本并统计长度。
接下来是真正的食物数据集类:
python
class food_dataset(Dataset):
def __init__(self, file_path, transform=None):
self.file_path = file_path
self.imgs = []
self.labels = []
self.transform = transform
with open(self.file_path, 'r', encoding='gbk') as f:
samples = [x.strip().split(' ') for x in f.readlines()]
for img_path, label in samples:
self.imgs.append(img_path)
self.labels.append(int(label)) # 直接转int,不再存字符串
def __len__(self):
return len(self.imgs)
def __getitem__(self, idx):
image = Image.open(self.imgs[idx]).convert("RGB") # 统一转RGB,解决灰度图通道报错
if self.transform:
image = self.transform(image)
label = self.labels[idx]
return image, label
关键点:
- init:读取 txt 文件,按空格拆分每行,分别存入图片路径列表和标签列表。标签直接转为 int,避免后续计算类型不匹配。
- getitem:用 PIL 打开图片并统一转为 RGB 三通道,解决灰度图导致的通道数报错问题;随后应用 transform 预处理,返回(图像张量,标签)元组。
- len:返回样本总数,供 DataLoader 计算迭代轮次。
5. 图像预处理与 DataLoader
python
data_transforms = {
# 机器学习的时候,数据进行归一化? 图片做归一化 ->0~1
'trainda':
transforms.Compose([
transforms.Resize([256, 256]), # 图像变换大小 opencv int8 0~255
transforms.RandomRotation(45), # 随机旋转,-45到45度之间随机选
transforms.CenterCrop(256), # 从中心开始裁剪[256,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), # 参数1为亮度,参数2为对比度,参数3为饱和度,参数4为色相
transforms.RandomGrayscale(p=0.1), # 概率转换成灰度,3通道就是R=G=B
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) # 标准化,均值,标准差
]),
'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])
]),
}
transforms.Compose 把多个预处理操作串联成一个管道:
- Resize(256, 256):把图片统一缩放到 256×256,保证网络输入尺寸一致。
- RandomRotation(45):在 -45 到 45 度之间随机旋转图片,增强模型对旋转变化的鲁棒性。
- CenterCrop(256):从中心裁剪出 256×256 的区域,配合 Resize 保证输入尺寸一致。
- RandomHorizontalFlip(p=0.5):以 0.5 的概率随机水平翻转图片。
- RandomVerticalFlip(p=0.5):以 0.5 的概率随机垂直翻转图片。
- ColorJitter(brightness=0.2, contrast=0.1, saturation=0.1, hue=0.1):随机调整亮度、对比度、饱和度和色相,增强模型对颜色变化的鲁棒性。
- RandomGrayscale(p=0.1):以 0.1 的概率将图片转为灰度图(三通道 R=G=B)。
- ToTensor():把 PIL 图像或 numpy 数组转为 PyTorch 张量,并将像素值从 0-255 归一化到 0-1,同时把通道维度从 HWC 调整为 CHW。
- Normalize(mean, std):用 ImageNet 的均值和标准差对张量做标准化,加速模型收敛。
训练集使用完整的数据增强管道,验证集只做 Resize、ToTensor 和 Normalize,不做随机增强,保证评估结果稳定。
随后创建数据集和加载器:
python
training_data = food_dataset('./train.txt', transform=data_transforms['trainda'])
test_data = food_dataset('./test.txt', transform=data_transforms['valid'])
train_dataloader = DataLoader(training_data, batch_size=4, shuffle=True)
test_dataloader = DataLoader(test_data, batch_size=4, shuffle=True)
DataLoader 负责把 Dataset 包装成可迭代对象:batch_size=4 表示每次取 4 个样本组成一个 batch;shuffle=True 表示每个 epoch 开始前打乱数据顺序,提升训练稳定性。
6. CNN 网络结构解析
python
class CNN(nn.Module):
def __init__(self):
super(CNN, self).__init__()
self.conv1 = nn.Sequential(
nn.Conv2d(
in_channels=3,
out_channels=16,
kernel_size=5,
stride=1,
padding=2,
),
nn.ReLU(),
nn.MaxPool2d(kernel_size=2),
)
self.conv2 = nn.Sequential(
nn.Conv2d(16, 32, 5, 1, 2),
nn.ReLU(),
nn.Conv2d(32, 32, 5, 1, 2),
nn.ReLU(),
nn.MaxPool2d(2),
)
self.conv3 = nn.Sequential(
nn.Conv2d(32, 128, 5, 1, 2),
nn.ReLU(),
)
self.out = nn.Linear(128 * 64 * 64, 20)
def forward(self, x):
x = self.conv1(x)
x = self.conv2(x)
x = self.conv3(x)
x = x.view(x.size(0), -1)
output = self.out(x)
return output
网络由三个卷积块和一个全连接输出层组成,逐层分析:
- conv1:输入 3 通道(RGB),输出 16 通道,卷积核 5×5,padding=2 保持尺寸不变,ReLU 激活后接 2×2 最大池化,特征图从 256×256 降为 128×128。
- conv2:16→32 通道卷积 + ReLU,再接 32→32 卷积 + ReLU,最后 2×2 池化,特征图降为 64×64。
- conv3:32→128 通道卷积 + ReLU,无池化,特征图保持 64×64。
- out:全连接层,输入维度为 128×64×64(展平后的特征总数),输出 20 个类别分数。
forward 方法中,x.view(x.size(0), -1) 把每个样本的 128×64×64 特征图展平为一维向量,再送入全连接层得到 20 类别的 logits。
7. 训练与测试循环
python
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = CNN().to(device)
loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)
这段代码完成三件事:
- device 选择:优先使用 GPU(cuda),否则回退到 CPU。
- 模型与损失函数:实例化 CNN 并移动到指定设备;CrossEntropyLoss 内部包含 Softmax 和交叉熵计算,适合多分类任务。
- 优化器:Adam 优化器,学习率 1e-4,更新模型全部参数。
python
def train_loop(dataloader, model, loss_fn, optimizer):
model.train()
for batch, (X, y) in enumerate(dataloader):
X, y = X.to(device), y.to(device)
pred = model(X)
loss = loss_fn(pred, y)
optimizer.zero_grad()
loss.backward()
optimizer.step()
if batch % 10 == 0:
print(f"Train Batch:{batch}, Loss: {loss.item():.4f}")
训练循环的标准五步:
- model.train():切换到训练模式,启用 Dropout 和 BatchNorm 的训练行为。
- 前向传播:把 batch 数据送入模型得到预测 pred。
- 计算损失:pred 与真实标签 y 计算交叉熵。
- 反向传播:optimizer.zero_grad() 清空旧梯度,loss.backward() 计算新梯度,optimizer.step() 更新参数。
- 日志输出:每 10 个 batch 打印一次当前损失,便于观察收敛情况。
python
def test_loop(dataloader, model, loss_fn):
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(X)
test_loss += loss_fn(pred, y).item()
correct += (pred.argmax(1) == y).type(torch.float).sum().item()
test_loss /= len(dataloader)
correct /= len(dataloader.dataset)
print(f"==== Test ==== Loss: {test_loss:.4f}, Acc: {correct:.4f}")
测试循环与训练的关键区别:
- model.eval():切换到评估模式,关闭 Dropout 等训练专用行为。
- torch.no_grad():关闭梯度计算,节省内存并加速推理。
- pred.argmax(1):取每个样本预测分数最高的类别下标,与真实标签比较,统计正确数。
- 准确率计算:正确数除以总样本数,得到测试集准确率。
8. 完整训练流程与注意事项
把上述函数串联起来即可开始训练:
python
epochs = 10
for t in range(epochs):
print(f"Epoch {t+1}\n-------------------------------")
train_loop(train_dataloader, model, loss_fn, optimizer)
test_loop(test_dataloader, model, loss_fn)
print("Done!")
实际使用中需要注意以下几点:
- 类别数匹配:全连接层输出 20 对应 20 个类别,需与数据集实际类别数一致。
- 内存占用:128×64×64 的展平特征较大,若显存不足可减小 batch_size 或增加池化层。
- 数据增强:当前仅做了 Resize 和 ToTensor,可加入 RandomHorizontalFlip、RandomRotation 等增强手段提升泛化能力。
- 学习率调整:1e-4 是较保守的初始值,训练后期可配合学习率衰减策略。
9. 总结
本文完整解析了 PyTorch 食物图像分类的代码链路:先用 os.walk 整理数据集生成标签文件,再通过自定义 Dataset 读取图片和标签,配合 transforms 预处理和 DataLoader 批量加载,最后搭建三层 CNN 并完成训练与测试循环。理解这套流程后,可以轻松迁移到其他图像分类任务,只需调整类别数、网络结构和数据增强策略即可。