增加图像分割相应的代码
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user