基于深度学习的天气图像分类竞赛项目,将天气图片自动分类为 cloudy(多云)、rainy(雨天)、snowy(雪天)、sunny(晴天) 四类。
Image (B, 3, 224, 224)
│
▼ ViT-B/16 [frozen + LoRA 微调]
Patch tokens (B, 196, 768)
│
▼ BiLSTM
Sequence (B, 196, H×2)
│
▼ Transformer Encoder [手写自注意力模块]
Contextualised (B, 196, H×2)
│
▼ Mean Pool → Linear → ReLU → Dropout → Linear
Logits (B, 4)
- ViT-B/16:使用 CLIP 预训练视觉编码器提取 patch 特征,默认冻结并注入 LoRA 低秩适配器进行高效微调
- BiLSTM:捕获 patch 序列的双向时序依赖
- Transformer Encoder:手写的多头自注意力 + 前馈网络堆叠,进一步建模全局上下文
- Classifier Head:LayerNorm → Linear → ReLU → Dropout → Linear,输出 4 类 logits
仅在 ViT 最后两层的 Q/V 投影矩阵上添加低秩适配器(默认 rank=8),训练参数量极小(通常 < 1M),同时保持接近全量微调的效果。支持:
merge_lora()/unmerge_lora():将 LoRA 权重合并/分离到基础权重中save_lora()/load_lora():单独保存/加载 LoRA 权重(文件极小)
推理时对每张图片生成 4 个版本(原图 + 亮度抖动 + 高斯噪声 + 条纹干扰),平均 logits 后输出最终预测,提升模型鲁棒性。
weather/
├── mainv2.py # 训练入口,命令行参数解析与主流程
├── infer.py # 推理接口(待实现)
├── model/
│ ├── __init__.py
│ └── modelv2.py # WeatherBaseline 模型 + LoRA + 手写 Transformer
├── utils/
│ ├── __init__.py
│ ├── weather_dataset.py # WeatherDataset + 数据增强 + DataLoader 工厂
│ ├── train.py # Trainer 训练循环 + 评估
│ └── others.py # 日志、早停、种子设置、Focal Loss
├── datas/
│ ├── row/ # 原始数据(按类别分子文件夹)
│ │ ├── cloudy/
│ │ ├── rainy/
│ │ ├── snowy/
│ │ └── sunny/
│ ├── clip-vit-base-patch16/ # CLIP ViT-B/16 预训练权重
│ ├── checkpoints/ # 模型检查点
│ └── logs/ # 训练日志 + TensorBoard 事件
└── README.md
- Python ≥ 3.8
- PyTorch ≥ 1.12
- CUDA(推荐)
pip install torch torchvision transformers scikit-learn tqdm tensorboard pillow将天气图片按类别放入 datas/row/ 目录下:
datas/row/
├── cloudy/
│ ├── cloudy_00001.jpg
│ └── ...
├── rainy/
├── snowy/
└── sunny/
支持的图片格式:.jpg、.jpeg、.png、.bmp。
# 默认参数训练
python mainv2.py
# 自定义超参数
python mainv2.py --batch_size 64 --lr 1e-4 --max_epochs 100 --patience 10
# 使用 ResNet 标准化统计量
python mainv2.py --model_type resnet
# 关闭 LoRA(完全冻结 ViT)
python mainv2.py --use_lora Falsetensorboard --logdir datas/logs/tensorboard| 参数 | 默认值 | 说明 |
|---|---|---|
--data_dir |
datas/row/train |
训练数据根目录 |
--model_type |
vit |
模型类型:vit / resnet |
--image_size |
224 |
输入图片尺寸 |
--batch_size |
32 |
批次大小 |
--val_ratio |
0.15 |
验证集比例 |
--test_ratio |
0.15 |
测试集比例 |
--lr |
1e-4 |
学习率 |
--weight_decay |
1e-4 |
权重衰减 |
--max_epochs |
50 |
最大训练轮数 |
--patience |
5 |
早停耐心值 |
--use_lora |
True |
是否启用 LoRA 微调 |
--lora_r |
8 |
LoRA 秩 |
--lora_alpha |
16 |
LoRA 缩放因子 |
--lstm_hidden |
256 |
BiLSTM 隐藏层维度 |
--lstm_layers |
2 |
BiLSTM 层数 |
--n_heads |
8 |
Transformer 注意力头数 |
--d_ff |
1024 |
Transformer 前馈维度 |
--n_tf_layers |
2 |
Transformer 编码器层数 |
--dropout |
0.3 |
Dropout 比率 |
--seed |
42 |
随机种子 |
--checkpoint_dir |
datas/checkpoints/best_modelv2_v1.pth |
模型保存路径 |
- 优化器:AdamW
- 学习率调度:Cosine Annealing
- 损失函数:CrossEntropyLoss(label smoothing = 0.15)
- 早停策略:基于验证集 Macro-F1,连续
patience轮无提升则停止 - 数据增强:随机水平翻转、随机旋转(±10°)、色彩抖动
- TensorBoard:记录 Loss、Accuracy、F1 曲线
- 分层抽样:按类别比例进行 3-way 分层划分(train / val / test)
- Accuracy:整体分类准确率
- Macro-F1:宏平均 F1 分数(各类别权重相等,适合类别不均衡场景)
推理接口定义在 infer.py 中(待完善),目前输出四类标签之一:
label = ['cloudy', 'rainy', 'snowy', 'sunny']