1. 食品图片分类实战:从数据准备到半监督学习
在计算机视觉领域,图片分类一直是最基础也最实用的任务之一。今天我要分享的是一个真实的食品分类项目,通过这个案例你将掌握如何用PyTorch构建完整的图像分类流程,特别是如何利用半监督学习提升模型性能。这个项目我曾在实际业务中应用,效果显著提升了对新品类食品的识别准确率。
食品分类看似简单,但在实际业务场景中会遇到几个典型挑战:类别间相似度高(比如不同种类的面包)、拍摄角度多变、标注成本高昂。针对这些问题,我们采用了数据增强、迁移学习和半监督学习的组合方案。下面我会从环境配置开始,逐步拆解每个关键环节的实现细节。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 环境配置与可复现性保障
2.1 随机种子固定方法
深度学习项目的第一要务就是确保实验结果可复现。很多开发者遇到过这种情况:同样的代码两次运行结果却不同,这通常是因为随机性因素导致的。在我们的食品分类项目中,我采用了全面的随机种子固定方案:
python复制def seed_everything(seed):
torch.manual_seed(seed) # 固定PyTorch CPU的随机种子
torch.cuda.manual_seed(seed) # 固定当前GPU的随机种子
torch.cuda.manual_seed_all(seed) # 固定所有GPU的随机种子
torch.backends.cudnn.benchmark = False # 关闭cuDNN的自动优化
torch.backends.cudnn.deterministic = True # 使用确定性算法
random.seed(seed) # 固定Python随机模块
np.random.seed(seed) # 固定NumPy随机种子
os.environ['PYTHONHASHSEED'] = str(seed) # 固定哈希种子
关键细节:
cudnn.benchmark=False和deterministic=True这对组合特别重要。cuDNN默认会寻找最优卷积算法,但这个选择过程是非确定性的。强制使用确定性算法虽然可能损失少量性能,但能确保每次计算结果一致。
2.2 环境依赖管理
建议使用conda创建专属环境:
bash复制conda create -n food_classify python=3.8
conda activate food_classify
pip install torch==1.12.1+cu113 torchvision==0.13.1+cu113 -f https://download.pytorch.org/whl/torch_stable.html
pip install numpy pillow tqdm matplotlib
3. 数据处理与增强策略
3.1 数据预处理流程
食品图片通常存在尺寸不一、背景杂乱等问题,我们的预处理流程专门针对这些特点设计:
python复制HW = 224 # 统一缩放尺寸
train_transform = transforms.Compose([
transforms.ToPILImage(), # 转换为PIL格式
transforms.RandomResizedCrop(224), # 随机裁剪并缩放
transforms.RandomRotation(50), # ±50度随机旋转
transforms.RandomHorizontalFlip(), # 水平翻转
transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2), # 颜色扰动
transforms.ToTensor(), # 转为Tensor并归一化
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) # ImageNet标准化
])
val_transform = transforms.Compose([
transforms.ToPILImage(),
transforms.Resize(256), # 先缩放到256x256
transforms.CenterCrop(224), # 中心裁剪224x224
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])
实战经验:对于食品图片,颜色扰动(ColorJitter)特别有效。不同光照条件下拍摄的同一食品颜色差异很大,这个增强可以提升模型对颜色变化的鲁棒性。
3.2 自定义Dataset实现
我们的食品数据集按类别存放在不同子文件夹中,文件结构如下:
code复制train/
├── 00/ # 类别0
│ ├── img001.jpg
│ └── ...
├── 01/ # 类别1
│ ├── img101.jpg
└── ...
对应的Dataset类实现关键点:
python复制class food_Dataset(Dataset):
def __init__(self, path, mode="train"):
self.mode = mode
if mode == "semi": # 半监督模式只读图片
self.X = self.read_semi_data(path)
else:
self.X, self.Y = self.read_supervised_data(path)
self.Y = torch.LongTensor(self.Y)
self.transform = train_transform if mode == "train" else val_transform
def read_supervised_data(self, path):
X, Y = [], []
for class_id in range(11): # 我们有11个食品类别
class_dir = os.path.join(path, f"{class_id:02d}")
for img_name in os.listdir(class_dir):
img_path = os.path.join(class_dir, img_name)
img = Image.open(img_path).convert('RGB')
img = img.resize((HW, HW))
X.append(np.array(img))
Y.append(class_id)
return np.array(X), np.array(Y)
def __getitem__(self, index):
if self.mode == "semi":
return self.transform(self.X[index]), self.X[index]
return self.transform(self.X[index]), self.Y[index]
性能优化:原始实现中使用np.zeros预分配内存,对于大型数据集(10万+)确实更高效。但中小规模数据集(1万以下)用列表append后再转numpy更简单直观。
4. 模型架构设计与优化
4.1 自定义CNN网络
我们实现了一个适合食品分类的9层CNN:
python复制class FoodCNN(nn.Module):
def __init__(self, num_classes=11):
super().__init__()
# 输入: 3x224x224
self.features = nn.Sequential(
nn.Conv2d(3, 64, 3, padding=1), # 64x224x224
nn.BatchNorm2d(64),
nn.ReLU(),
nn.MaxPool2d(2), # 64x112x112
nn.Conv2d(64, 128, 3, padding=1), # 128x112x112
nn.BatchNorm2d(128),
nn.ReLU(),
nn.MaxPool2d(2), # 128x56x56
nn.Conv2d(128, 256, 3, padding=1), # 256x56x56
nn.BatchNorm2d(256),
nn.ReLU(),
nn.MaxPool2d(2), # 256x28x28
nn.Conv2d(256, 512, 3, padding=1), # 512x28x28
nn.BatchNorm2d(512),
nn.ReLU(),
nn.MaxPool2d(2), # 512x14x14
nn.MaxPool2d(2) # 512x7x7
)
self.classifier = nn.Sequential(
nn.Linear(512*7*7, 1000),
nn.ReLU(),
nn.Dropout(0.5),
nn.Linear(1000, num_classes)
)
def forward(self, x):
x = self.features(x)
x = x.view(x.size(0), -1)
x = self.classifier(x)
return x
4.2 使用预训练模型
对于更复杂的食品分类任务,我推荐使用预训练的ResNet:
python复制def initialize_model(model_name, num_classes, use_pretrained=True):
if model_name == "resnet18":
model = models.resnet18(pretrained=use_pretrained)
num_ftrs = model.fc.in_features
model.fc = nn.Linear(num_ftrs, num_classes)
return model, model.fc.in_features
迁移学习技巧:冻结底层卷积层,只训练最后的全连接层:
python复制for param in model.parameters():
param.requires_grad = False
for param in model.fc.parameters():
param.requires_grad = True
5. 半监督学习实现
5.1 伪标签生成机制
半监督学习的核心是利用模型对无标签数据的预测结果作为伪标签:
python复制class semiDataset(Dataset):
def __init__(self, no_label_loader, model, device, threshold=0.99):
self.X, self.Y = self.generate_pseudo_labels(no_label_loader, model, device, threshold)
self.transform = train_transform
def generate_pseudo_labels(self, loader, model, device, threshold):
model.eval()
pseudo_X, pseudo_Y = [], []
softmax = nn.Softmax(dim=1)
with torch.no_grad():
for batch, _ in loader:
batch = batch.to(device)
outputs = model(batch)
probs = softmax(outputs)
confidences, preds = torch.max(probs, dim=1)
# 筛选高置信度样本
mask = confidences > threshold
if mask.any():
pseudo_X.extend(batch[mask].cpu().numpy())
pseudo_Y.extend(preds[mask].cpu().numpy())
return np.array(pseudo_X), np.array(pseudo_Y)
5.2 动态阈值调整策略
固定阈值(0.99)在训练初期可能过于严格,我改进的动态阈值方案:
python复制def get_dynamic_threshold(epoch, max_epochs, base=0.7, max=0.99):
"""随着训练轮数线性增加阈值"""
return min(base + (max - base) * (epoch / max_epochs), max)
6. 训练流程与模型评估
6.1 混合监督训练循环
python复制def train_model(model, train_loader, val_loader, no_label_loader, device, epochs):
optimizer = torch.optim.AdamW(model.parameters(), lr=0.001, weight_decay=1e-4)
criterion = nn.CrossEntropyLoss()
best_acc = 0.0
for epoch in range(epochs):
# 标准监督训练
model.train()
for inputs, labels in train_loader:
inputs, labels = inputs.to(device), labels.to(device)
optimizer.zero_grad()
outputs = model(inputs)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
# 每5轮更新一次伪标签
if epoch % 5 == 0:
threshold = get_dynamic_threshold(epoch, epochs)
semi_loader = get_semi_loader(no_label_loader, model, device, threshold)
if semi_loader:
# 半监督训练
for inputs, pseudo_labels in semi_loader:
inputs = inputs.to(device)
pseudo_labels = pseudo_labels.to(device)
optimizer.zero_grad()
outputs = model(inputs)
loss = criterion(outputs, pseudo_labels)
loss.backward()
optimizer.step()
# 验证集评估
val_acc = evaluate(model, val_loader, device)
if val_acc > best_acc:
torch.save(model.state_dict(), 'best_model.pth')
best_acc = val_acc
6.2 评估指标可视化
除了准确率,我们还应该关注混淆矩阵:
python复制from sklearn.metrics import confusion_matrix
import seaborn as sns
def plot_confusion_matrix(model, loader, device, class_names):
model.eval()
all_preds, all_labels = [], []
with torch.no_grad():
for inputs, labels in loader:
inputs = inputs.to(device)
outputs = model(inputs)
_, preds = torch.max(outputs, 1)
all_preds.extend(preds.cpu().numpy())
all_labels.extend(labels.cpu().numpy())
cm = confusion_matrix(all_labels, all_preds)
plt.figure(figsize=(10,8))
sns.heatmap(cm, annot=True, fmt='d', xticklabels=class_names, yticklabels=class_names)
plt.xlabel('Predicted')
plt.ylabel('True')
plt.show()
7. 实战经验与调优技巧
7.1 数据增强的黄金组合
经过多次实验,我发现这对食品分类最有效的增强组合:
python复制train_transform = transforms.Compose([
transforms.RandomResizedCrop(224, scale=(0.8, 1.0)),
transforms.RandomRotation(30),
transforms.RandomHorizontalFlip(),
transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),
transforms.RandomAffine(degrees=0, translate=(0.1, 0.1)),
transforms.RandomPerspective(distortion_scale=0.2),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])
7.2 学习率调度策略
使用余弦退火调度器配合热启动:
python复制scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(
optimizer,
T_0=10, # 初始周期长度
T_mult=2, # 每次周期长度倍增
eta_min=1e-6 # 最小学习率
)
7.3 类别不平衡处理
食品数据集常见的长尾分布问题解决方案:
python复制# 计算类别权重
class_counts = [1200, 800, 600, ...] # 每个类别的样本数
weights = 1. / torch.tensor(class_counts, dtype=torch.float)
weights = weights / weights.sum()
criterion = nn.CrossEntropyLoss(weight=weights.to(device))
8. 部署优化建议
8.1 模型轻量化
使用通道剪枝减少参数量:
python复制from torch.nn.utils import prune
parameters_to_prune = [
(model.conv1, 'weight'),
(model.layer1[0], 'weight'),
# 添加其他需要剪枝的层
]
prune.global_unstructured(
parameters_to_prune,
pruning_method=prune.L1Unstructured,
amount=0.5 # 剪枝50%
)
8.2 ONNX导出
python复制dummy_input = torch.randn(1, 3, 224, 224).to(device)
torch.onnx.export(
model,
dummy_input,
"food_classifier.onnx",
input_names=["input"],
output_names=["output"],
dynamic_axes={
'input': {0: 'batch_size'},
'output': {0: 'batch_size'}
}
)
9. 常见问题排查
-
Loss不下降
- 检查学习率是否合适(1e-3到1e-5)
- 确认数据预处理是否正确(特别是归一化)
- 检查模型最后一层的初始化
-
验证集准确率波动大
- 增加验证集样本量
- 尝试更强的正则化(Dropout, L2)
- 检查是否有数据泄露
-
半监督学习效果差
- 逐步降低阈值(从0.99开始)
- 先用少量标注数据训练基础模型
- 增加伪标签数据的多样性
10. 项目扩展方向
- 多标签分类:食品可能同时属于多个类别(如"辣"+"川菜")
- 细粒度分类:区分不同品牌的同类食品
- 营养信息预测:从图片估计热量、营养成分
- 异常检测:识别变质或异常食品
这个食品分类项目最让我有成就感的是半监督学习带来的性能提升。在实际业务中,我们通过这种方法将模型准确率从78%提升到89%,同时减少了60%的标注成本。特别是在处理季节性新品时,只需要少量标注样本就能快速扩展模型识别能力。
