增加图像分割相应的代码

This commit is contained in:
2025-12-09 18:24:28 +08:00
parent 069b8897e8
commit 56de053db1
9 changed files with 2256 additions and 4 deletions
+316
View File
@@ -0,0 +1,316 @@
"""
训练配置文件
功能说明:
1. 集中管理所有训练超参数
2. 区分CPU和GPU训练配置
3. 便于实验管理和超参数调优
使用方式:
from config import TrainConfig
config = TrainConfig()
# 根据需要修改配置
config.batch_size = 2
"""
from dataclasses import dataclass
from pathlib import Path
import torch
@dataclass
class TrainConfig:
"""
训练配置类
使用dataclass装饰器,自动生成__init__等方法
所有配置都有类型注解和默认值
"""
# ==================== 路径配置 ====================
# 数据集根目录
data_root: str = "d:/MyProjects/PythonProjects/FoodClassifier/SegFormer/data/segformer_format"
# 模型保存目录
output_dir: str = "d:/MyProjects/PythonProjects/FoodClassifier/SegFormer/outputs"
# 预训练模型名称
pretrained_model: str = "nvidia/segformer-b0-finetuned-ade-512-512"
# ==================== 模型配置 ====================
# 类别数(包括背景)
# 0: 背景, 1: 食物
num_classes: int = 2
# 输入图像尺寸
image_size: int = 512
# SegFormer模型版本
# B0: 最轻量 (3.7M参数)
# B1: 轻量 (13.7M)
# B2: 中等 (24.7M)
# B3: 较大 (44.6M)
# B4: 大型 (61.4M)
# B5: 超大 (81.9M)
model_variant: str = "b0"
# ==================== 训练配置 ====================
# 批大小
# CPU训练建议: 1-2
# GPU训练建议: 4-8 (4090可以用8-16)
batch_size: int = 1
# 总训练轮数
num_epochs: int = 50
# 学习率
learning_rate: float = 6e-5
# 权重衰减(L2正则化)
weight_decay: float = 0.01
# 学习率调度器类型
# 'cosine': 余弦退火
# 'linear': 线性衰减
# 'polynomial': 多项式衰减
lr_scheduler_type: str = "cosine"
# Warmup轮数(学习率逐步增加的轮数)
warmup_epochs: int = 5
# 梯度裁剪(防止梯度爆炸)
max_grad_norm: float = 1.0
# ==================== 优化器配置 ====================
# 优化器类型: 'adamw' 或 'sgd'
optimizer_type: str = "adamw"
# AdamW的beta参数
adam_betas: tuple = (0.9, 0.999)
# SGD动量
sgd_momentum: float = 0.9
# ==================== 损失函数配置 ====================
# 是否使用混合损失 (CrossEntropy + Dice)
use_mixed_loss: bool = True
# Dice Loss的权重
dice_loss_weight: float = 0.5
# 类别权重(用于处理类别不平衡)
# None: 自动计算
# List[float]: 手动指定 [背景权重, 食物权重]
class_weights: list = None # 例如: [0.3, 0.7]
# ==================== 数据加载配置 ====================
# 数据加载进程数
# CPU训练: 0 (避免进程间通信开销)
# GPU训练: 4-8
num_workers: int = 0
# 是否使用数据增强
use_augmentation: bool = True
# ==================== 训练策略 ====================
# 是否使用两阶段训练
# Stage 1: 冻结Encoder,只训练Decoder
# Stage 2: Fine-tune整个模型
use_two_stage_training: bool = True
# Stage 1的训练轮数(冻结Encoder
stage1_epochs: int = 10
# ==================== 验证和保存 ====================
# 验证频率(每N个epoch验证一次)
eval_every_n_epochs: int = 5
# 保存checkpoint频率(每N个epoch保存一次)
save_every_n_epochs: int = 10
# 是否只保存最佳模型
save_best_only: bool = True
# 最佳模型的评估指标: 'miou', 'loss', 'pixel_acc'
best_metric: str = "miou"
# ==================== 设备配置 ====================
# 设备: 'cpu', 'cuda', 'auto'
device: str = "auto"
# 是否使用混合精度训练(仅GPU
use_amp: bool = False
# ==================== 日志配置 ====================
# 打印频率(每N个batch打印一次)
print_every_n_batches: int = 5
# 是否保存训练日志
save_logs: bool = True
# 是否启用详细日志(包括每个batch的详细信息)
verbose: bool = True
# ==================== 随机种子 ====================
# 随机种子(确保可复现)
random_seed: int = 42
def __post_init__(self):
"""
初始化后的处理
在所有参数赋值后自动调用,用于:
1. 自动推断设备
2. 创建输出目录
3. 验证配置的合理性
"""
# 1. 自动推断设备
if self.device == "auto":
self.device = "cuda" if torch.cuda.is_available() else "cpu"
# 2. CPU训练时自动调整配置
if self.device == "cpu":
self.use_amp = False # CPU不支持混合精度
self.num_workers = 0 # CPU训练避免多进程开销
if self.batch_size > 2:
print(f"⚠️ CPU训练建议batch_size<=2,当前值 {self.batch_size} 可能很慢")
# 3. 创建输出目录
Path(self.output_dir).mkdir(parents=True, exist_ok=True)
# 4. 验证配置
assert self.num_classes >= 2, "类别数必须>=2"
assert self.batch_size > 0, "batch_size必须>0"
assert self.num_epochs > 0, "num_epochs必须>0"
assert self.learning_rate > 0, "learning_rate必须>0"
if self.use_two_stage_training:
assert self.stage1_epochs < self.num_epochs, \
"stage1_epochs必须小于num_epochs"
def to_dict(self) -> dict:
"""
将配置转换为字典(用于保存)
Returns:
config_dict: 配置字典
"""
return {
k: v for k, v in self.__dict__.items()
if not k.startswith('_')
}
def print_config(self):
"""打印所有配置信息"""
print("\n" + "="*60)
print("训练配置")
print("="*60)
print("\n【路径配置】")
print(f" 数据集: {self.data_root}")
print(f" 输出目录: {self.output_dir}")
print(f" 预训练模型: {self.pretrained_model}")
print("\n【模型配置】")
print(f" 类别数: {self.num_classes}")
print(f" 图像尺寸: {self.image_size}×{self.image_size}")
print(f" 模型版本: SegFormer-{self.model_variant.upper()}")
print("\n【训练配置】")
print(f" 批大小: {self.batch_size}")
print(f" 训练轮数: {self.num_epochs}")
print(f" 学习率: {self.learning_rate}")
print(f" 权重衰减: {self.weight_decay}")
print(f" 学习率调度: {self.lr_scheduler_type}")
print(f" Warmup轮数: {self.warmup_epochs}")
print("\n【训练策略】")
if self.use_two_stage_training:
print(f" 两阶段训练: 是")
print(f" Stage 1 (冻结Encoder): {self.stage1_epochs} epochs")
print(f" Stage 2 (全模型Fine-tune): {self.num_epochs - self.stage1_epochs} epochs")
else:
print(f" 两阶段训练: 否")
print("\n【损失函数】")
if self.use_mixed_loss:
print(f" 混合损失: CrossEntropy + {self.dice_loss_weight}×Dice")
else:
print(f" 损失函数: CrossEntropy")
if self.class_weights:
print(f" 类别权重: {self.class_weights}")
print("\n【设备配置】")
print(f" 设备: {self.device}")
print(f" 混合精度: {'' if self.use_amp else ''}")
print(f" 数据加载进程: {self.num_workers}")
print("\n【验证和保存】")
print(f" 验证频率: 每{self.eval_every_n_epochs}")
print(f" 保存频率: 每{self.save_every_n_epochs}")
print(f" 最佳指标: {self.best_metric}")
print("="*60 + "\n")
def get_cpu_config() -> TrainConfig:
"""
获取CPU训练的推荐配置
适用于:
- 本地开发调试
- 快速验证代码正确性
- 小数据集实验
"""
config = TrainConfig()
# CPU优化配置
config.device = "cpu"
config.batch_size = 1
config.num_workers = 0
config.use_amp = False
config.image_size = 256 # 降低分辨率加快训练
# 快速验证配置
config.num_epochs = 20
config.eval_every_n_epochs = 5
config.save_every_n_epochs = 10
return config
def get_gpu_config() -> TrainConfig:
"""
获取GPU训练的推荐配置
适用于:
- 4090等高性能GPU
- 正式训练
- 追求最佳性能
"""
config = TrainConfig()
# GPU优化配置
config.device = "cuda"
config.batch_size = 8
config.num_workers = 4
config.use_amp = True
config.image_size = 512
# 完整训练配置
config.num_epochs = 100
config.eval_every_n_epochs = 5
config.save_every_n_epochs = 10
return config
if __name__ == "__main__":
"""测试配置"""
print("CPU配置:")
cpu_config = get_cpu_config()
cpu_config.print_config()
print("\n\nGPU配置:")
gpu_config = get_gpu_config()
gpu_config.print_config()