1. 项目概述
天气识别是计算机视觉领域的一个经典应用场景,通过深度学习模型对天气图像进行分类识别。使用PyTorch框架实现天气识别系统,能够帮助我们掌握图像分类任务的全流程实现,包括数据准备、模型构建、训练优化和部署应用。
这个项目适合有一定Python基础,想要入门计算机视觉和PyTorch框架的开发者。通过本实践,你将学习到:
- 如何使用PyTorch构建卷积神经网络
- 如何处理天气图像数据集
- 如何训练和优化图像分类模型
- 如何评估模型性能并进行预测
2. 环境准备与数据收集
2.1 PyTorch环境配置
首先需要搭建PyTorch开发环境。推荐使用Anaconda创建独立的Python环境:
bash复制conda create -n weather python=3.8
conda activate weather
然后安装PyTorch及其依赖。根据你的硬件配置选择合适的版本:
bash复制# 无GPU版本
conda install pytorch torchvision torchaudio cpuonly -c pytorch
# CUDA 12.x GPU版本
conda install pytorch torchvision torchaudio pytorch-cuda=12 -c pytorch -c nvidia
提示:可以通过
torch.cuda.is_available()检查CUDA是否可用
2.2 天气数据集获取
常用的天气识别数据集包括:
- Multi-class Weather Dataset (MWD)
- SWIMSEG - 包含晴天、多云、雨天、雾天等类别
- 自建数据集 - 通过网络爬虫收集各类天气图片
以MWD数据集为例,它包含4类天气图像:
- 晴天(Sunny)
- 雨天(Rainy)
- 多云(Cloudy)
- 雾天(Foggy)
数据集目录结构应组织为:
code复制weather_dataset/
├── train/
│ ├── sunny/
│ ├── rainy/
│ ├── cloudy/
│ └── foggy/
└── test/
├── sunny/
├── rainy/
├── cloudy/
└── foggy/
3. 数据预处理与增强
3.1 数据加载与转换
使用PyTorch的torchvision工具进行数据加载和预处理:
python复制from torchvision import transforms, datasets
# 定义数据转换
train_transform = transforms.Compose([
transforms.Resize(256),
transforms.RandomCrop(224),
transforms.RandomHorizontalFlip(),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225])
])
test_transform = transforms.Compose([
transforms.Resize(256),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225])
])
# 加载数据集
train_data = datasets.ImageFolder('weather_dataset/train', transform=train_transform)
test_data = datasets.ImageFolder('weather_dataset/test', transform=test_transform)
# 创建数据加载器
train_loader = torch.utils.data.DataLoader(train_data, batch_size=32, shuffle=True)
test_loader = torch.utils.data.DataLoader(test_data, batch_size=32, shuffle=False)
3.2 数据增强策略
针对天气识别任务,推荐使用以下数据增强方法:
- 随机水平翻转
- 随机旋转(-30°到30°)
- 颜色抖动(调整亮度、对比度和饱和度)
- 随机灰度化(以一定概率转为灰度图)
这些增强可以帮助模型更好地学习天气特征,提高泛化能力。
4. 模型构建与训练
4.1 卷积神经网络设计
我们构建一个简单的CNN模型用于天气分类:
python复制import torch.nn as nn
import torch.nn.functional as F
class WeatherCNN(nn.Module):
def __init__(self, num_classes=4):
super(WeatherCNN, self).__init__()
self.conv1 = nn.Conv2d(3, 32, kernel_size=3, padding=1)
self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding=1)
self.pool = nn.MaxPool2d(2, 2)
self.fc1 = nn.Linear(64 * 56 * 56, 128)
self.fc2 = nn.Linear(128, num_classes)
def forward(self, x):
x = self.pool(F.relu(self.conv1(x)))
x = self.pool(F.relu(self.conv2(x)))
x = x.view(-1, 64 * 56 * 56)
x = F.relu(self.fc1(x))
x = self.fc2(x)
return x
4.2 迁移学习方案
对于更复杂的天气识别任务,可以使用预训练模型进行迁移学习:
python复制from torchvision import models
model = models.resnet18(pretrained=True)
num_ftrs = model.fc.in_features
model.fc = nn.Linear(num_ftrs, 4) # 4 weather classes
4.3 模型训练流程
完整的训练代码如下:
python复制import torch.optim as optim
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = WeatherCNN().to(device)
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)
for epoch in range(20):
model.train()
running_loss = 0.0
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()
running_loss += loss.item()
print(f'Epoch {epoch+1}, Loss: {running_loss/len(train_loader)}')
5. 模型评估与优化
5.1 性能评估指标
使用准确率、混淆矩阵等指标评估模型:
python复制from sklearn.metrics import confusion_matrix, classification_report
model.eval()
all_preds = []
all_labels = []
with torch.no_grad():
for inputs, labels in test_loader:
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))
print(confusion_matrix(all_labels, all_preds))
5.2 优化技巧
- 学习率调度:使用
ReduceLROnPlateau或StepLR动态调整学习率 - 早停机制:在验证集性能不再提升时停止训练
- 模型集成:结合多个模型的预测结果提高准确率
6. 模型部署与应用
6.1 保存与加载模型
python复制# 保存模型
torch.save(model.state_dict(), 'weather_model.pth')
# 加载模型
model = WeatherCNN()
model.load_state_dict(torch.load('weather_model.pth'))
model.eval()
6.2 单张图片预测
python复制from PIL import Image
def predict_weather(image_path):
image = Image.open(image_path)
image = test_transform(image).unsqueeze(0)
with torch.no_grad():
output = model(image)
_, predicted = torch.max(output, 1)
return classes[predicted.item()]
7. 常见问题与解决方案
7.1 数据不平衡问题
天气数据集中某些类别(如雾天)样本可能较少,解决方案:
- 使用类别权重
- 过采样少数类
- 数据增强时对少数类使用更强的增强
7.2 模型过拟合
解决方法:
- 增加Dropout层
- 使用L2正则化
- 早停机制
- 简化模型结构
7.3 天气间的模糊边界
某些天气条件(如多云和阴天)可能难以区分,可以:
- 引入更细粒度的标注
- 使用多标签分类
- 增加模型容量
8. 进阶优化方向
- 使用更先进的模型架构如EfficientNet、Vision Transformer
- 引入注意力机制增强关键区域识别
- 结合时间序列数据(如连续帧)提高识别准确率
- 部署到移动端实现实时天气识别
通过这个项目,我们实现了从数据准备到模型部署的完整流程。实际应用中,可以根据具体需求调整模型结构和参数,持续优化识别性能。
