1. 项目概述
这个基于Python和CNN的深度学习项目,旨在通过卷积神经网络实现对不同颜色鞋子的识别分类。作为一名长期从事计算机视觉开发的工程师,我发现颜色识别在实际应用中有着广泛的需求,比如电商平台的商品分类、智能仓储管理等领域。
项目采用经典的卷积神经网络架构,通过PyTorch框架实现。相比传统的图像处理方法,CNN能够自动提取图像的多层次特征,避免了手工设计特征的繁琐过程。在颜色识别任务上,CNN可以同时捕捉颜色信息和形状纹理特征,实现更准确的分类。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 核心设计思路
2.1 数据准备与预处理
数据是深度学习项目的基石。我们需要收集包含不同颜色鞋子的图片数据集,建议每类颜色至少准备500-1000张图片。数据来源可以是:
- 公开数据集如ImageNet的子集
- 网络爬虫获取的电商平台图片
- 自行拍摄的真实场景照片
预处理步骤包括:
- 统一图片尺寸为224x224像素
- 数据增强:随机旋转、翻转、亮度调整
- 归一化处理:将像素值缩放到[0,1]范围
python复制transform = transforms.Compose([
transforms.Resize((224,224)),
transforms.RandomHorizontalFlip(),
transforms.RandomRotation(15),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225])
])
2.2 网络架构设计
我们采用改进的ResNet18作为基础模型,针对颜色识别任务进行优化:
- 保留原始ResNet的特征提取部分
- 修改最后的全连接层,输出节点数等于颜色类别数
- 添加注意力机制模块,增强对颜色特征的关注
python复制class ColorNet(nn.Module):
def __init__(self, num_classes):
super(ColorNet, self).__init__()
self.base = models.resnet18(pretrained=True)
self.attention = nn.Sequential(
nn.Conv2d(512, 512, kernel_size=3, padding=1),
nn.Sigmoid()
)
self.base.fc = nn.Linear(512, num_classes)
def forward(self, x):
x = self.base.conv1(x)
x = self.base.bn1(x)
x = self.base.relu(x)
x = self.base.maxpool(x)
x = self.base.layer1(x)
x = self.base.layer2(x)
x = self.base.layer3(x)
x = self.base.layer4(x)
att = self.attention(x)
x = x * att
x = self.base.avgpool(x)
x = torch.flatten(x, 1)
x = self.base.fc(x)
return x
3. 模型训练与优化
3.1 训练参数设置
训练过程中需要精心调整以下超参数:
- 学习率:初始设为0.001,使用余弦退火策略
- 批量大小:根据GPU内存设为32或64
- 训练轮数:通常50-100个epoch
- 损失函数:交叉熵损失
- 优化器:AdamW
python复制model = ColorNet(num_classes=len(classes)).to(device)
criterion = nn.CrossEntropyLoss()
optimizer = optim.AdamW(model.parameters(), lr=0.001)
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=50)
3.2 训练过程监控
使用TensorBoard记录训练过程中的关键指标:
- 训练/验证集的损失和准确率
- 学习率变化曲线
- 混淆矩阵可视化
python复制from torch.utils.tensorboard import SummaryWriter
writer = SummaryWriter()
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()
# 记录训练指标
writer.add_scalar('Loss/train', loss.item(), epoch)
scheduler.step()
# 验证集评估
model.eval()
with torch.no_grad():
val_loss = 0
correct = 0
total = 0
for inputs, labels in val_loader:
inputs, labels = inputs.to(device), labels.to(device)
outputs = model(inputs)
loss = criterion(outputs, labels)
val_loss += loss.item()
_, predicted = outputs.max(1)
total += labels.size(0)
correct += predicted.eq(labels).sum().item()
val_acc = correct / total
writer.add_scalar('Accuracy/val', val_acc, epoch)
4. 模型评估与部署
4.1 性能评估指标
除了准确率外,还需要关注:
- 各类别的精确率、召回率和F1分数
- 混淆矩阵分析
- 推理速度(FPS)
python复制from sklearn.metrics import classification_report, confusion_matrix
def evaluate(model, dataloader):
model.eval()
all_preds = []
all_labels = []
with torch.no_grad():
for inputs, labels in dataloader:
inputs = inputs.to(device)
outputs = model(inputs)
_, preds = torch.max(outputs, 1)
all_preds.extend(preds.cpu().numpy())
all_labels.extend(labels.numpy())
print(classification_report(all_labels, all_preds, target_names=classes))
cm = confusion_matrix(all_labels, all_preds)
plt.figure(figsize=(10,8))
sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=classes, yticklabels=classes)
plt.xlabel('Predicted')
plt.ylabel('Actual')
plt.show()
4.2 模型部署方案
训练好的模型可以通过以下方式部署:
- Flask/Django构建Web API
- 转换为ONNX格式优化推理速度
- 移动端部署(使用PyTorch Mobile)
python复制# Flask API示例
from flask import Flask, request, jsonify
import torch
from PIL import Image
import io
app = Flask(__name__)
model = load_model() # 加载训练好的模型
@app.route('/predict', methods=['POST'])
def predict():
if 'file' not in request.files:
return jsonify({'error': 'no file uploaded'})
file = request.files['file']
img_bytes = file.read()
img = Image.open(io.BytesIO(img_bytes))
# 预处理
img_tensor = transform(img).unsqueeze(0)
# 预测
with torch.no_grad():
output = model(img_tensor)
_, pred = torch.max(output, 1)
return jsonify({'class': classes[pred.item()]})
if __name__ == '__main__':
app.run(host='0.0.0.0', port=5000)
5. 常见问题与解决方案
5.1 数据不平衡问题
当某些颜色的样本数量远多于其他颜色时,可以:
- 对少数类进行过采样
- 对多数类进行欠采样
- 在损失函数中添加类别权重
python复制# 计算类别权重
from sklearn.utils.class_weight import compute_class_weight
class_weights = compute_class_weight('balanced', classes=np.unique(train_labels), y=train_labels)
class_weights = torch.FloatTensor(class_weights).to(device)
criterion = nn.CrossEntropyLoss(weight=class_weights)
5.2 过拟合问题
解决方法包括:
- 增加数据增强方式
- 添加Dropout层
- 使用L2正则化
- 早停策略
python复制# 在模型中添加Dropout
self.base.fc = nn.Sequential(
nn.Dropout(0.5),
nn.Linear(512, num_classes)
)
5.3 颜色识别中的光照变化
针对不同光照条件下的颜色变化:
- 在数据集中包含多种光照条件的样本
- 使用色彩空间转换(如HSV)代替RGB
- 添加色彩归一化层
python复制# HSV色彩空间转换
def rgb_to_hsv(image):
image = image.convert('HSV')
return image
# 在数据增强中添加
transforms.Lambda(lambda x: rgb_to_hsv(x)),
6. 项目扩展与优化方向
6.1 多任务学习
可以同时预测鞋子的颜色和款式:
python复制class MultiTaskNet(nn.Module):
def __init__(self, num_colors, num_styles):
super().__init__()
self.base = models.resnet18(pretrained=True)
self.color_head = nn.Linear(512, num_colors)
self.style_head = nn.Linear(512, num_styles)
def forward(self, x):
x = self.base(x)
color = self.color_head(x)
style = self.style_head(x)
return color, style
6.2 模型轻量化
使用MobileNetV3等轻量级网络:
python复制from torchvision.models import mobilenet_v3_small
model = mobilenet_v3_small(pretrained=True)
model.classifier[3] = nn.Linear(1024, num_classes)
6.3 自监督预训练
利用无标注数据进行预训练:
python复制# 使用SimCLR等自监督方法
from lightly.models import SimCLR
encoder = models.resnet18()
model = SimCLR(encoder, num_ftrs=512)
在实际部署这个项目时,我发现有几个关键点需要特别注意:首先,确保训练数据中包含了各种光照条件下的鞋子图片,这对颜色识别的鲁棒性至关重要;其次,对于相似颜色(如深红和暗红)的区分,可以尝试在HSV色彩空间中进行训练;最后,模型部署时要考虑推理速度,必要时可以对模型进行量化处理。
