Skip to content

Latest commit

 

History

1 Commit

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

🌤️ Weather Image Classification

基于深度学习的天气图像分类竞赛项目,将天气图片自动分类为 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

LoRA(Low-Rank Adaptation)

仅在 ViT 最后两层的 Q/V 投影矩阵上添加低秩适配器(默认 rank=8),训练参数量极小(通常 < 1M),同时保持接近全量微调的效果。支持:

  • merge_lora() / unmerge_lora():将 LoRA 权重合并/分离到基础权重中
  • save_lora() / load_lora():单独保存/加载 LoRA 权重(文件极小)

测试时增强(TTA)

推理时对每张图片生成 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 False

查看训练曲线

tensorboard --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']

About

No description, website, or topics provided.

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages