增加图像分割相应的代码

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
+221
View File
@@ -0,0 +1,221 @@
"""
SegFormer-B0 最小化Demo
用于快速上手图像语义分割
"""
import os
import torch
import numpy as np
from PIL import Image
import matplotlib.pyplot as plt
from transformers import SegformerImageProcessor, SegformerForSemanticSegmentation
class SegFormerDemo:
def __init__(self, model_name="nvidia/segformer-b0-finetuned-ade-512-512"):
"""
初始化SegFormer-B0模型
Args:
model_name: 预训练模型名称
- "nvidia/segformer-b0-finetuned-ade-512-512" (ADE20K数据集,150类)
- "nvidia/segformer-b0-finetuned-cityscapes-1024-1024" (城市街景)
"""
print(f"加载SegFormer模型: {model_name}")
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"使用设备: {self.device}")
# 加载处理器和模型
self.processor = SegformerImageProcessor.from_pretrained(model_name)
self.model = SegformerForSemanticSegmentation.from_pretrained(model_name)
self.model.to(self.device)
self.model.eval()
# ADE20K数据集的类别数
self.num_classes = self.model.config.num_labels
print(f"模型加载完成! 支持 {self.num_classes} 个类别")
def segment_image(self, image_path):
"""
对图像进行语义分割
Args:
image_path: 图像路径
Returns:
pred_seg: 分割结果 (H, W) 每个像素的类别ID
original_image: 原始图像
"""
# 读取图像
image = Image.open(image_path).convert("RGB")
# 预处理
inputs = self.processor(images=image, return_tensors="pt")
inputs = {k: v.to(self.device) for k, v in inputs.items()}
# 推理
with torch.no_grad():
outputs = self.model(**inputs)
logits = outputs.logits
# 上采样到原始尺寸
upsampled_logits = torch.nn.functional.interpolate(
logits,
size=image.size[::-1], # (height, width)
mode="bilinear",
align_corners=False
)
# 获取每个像素的类别
pred_seg = upsampled_logits.argmax(dim=1)[0].cpu().numpy()
return pred_seg, np.array(image)
def visualize_segmentation(self, image, pred_seg, save_path=None):
"""
可视化分割结果
Args:
image: 原始图像 (H, W, 3)
pred_seg: 分割结果 (H, W)
save_path: 保存路径
"""
fig, axes = plt.subplots(1, 3, figsize=(15, 5))
# 原始图像
axes[0].imshow(image)
axes[0].set_title("Original Image")
axes[0].axis("off")
# 分割mask(彩色)
axes[1].imshow(pred_seg, cmap='tab20b')
axes[1].set_title("Segmentation Mask")
axes[1].axis("off")
# 叠加显示
axes[2].imshow(image)
axes[2].imshow(pred_seg, alpha=0.5, cmap='tab20b')
axes[2].set_title("Overlay")
axes[2].axis("off")
plt.tight_layout()
if save_path:
plt.savefig(save_path, dpi=150, bbox_inches='tight')
print(f"结果已保存到: {save_path}")
plt.show()
def extract_object_mask(self, pred_seg, target_classes):
"""
提取特定类别的mask
Args:
pred_seg: 分割结果 (H, W)
target_classes: 目标类别ID列表
Returns:
mask: 二值mask (H, W)
"""
mask = np.isin(pred_seg, target_classes)
return mask.astype(np.uint8) * 255
def extract_food_region(self, image, pred_seg):
"""
提取食物区域(基于ADE20K数据集)
ADE20K中食物相关类别(部分):
- 53: food, 54: table, 55: chair, 56: car, 57: person
- 116: plate, 117: banana, 118: apple, 119: sandwich, 120: orange
Args:
image: 原始图像 (H, W, 3)
pred_seg: 分割结果 (H, W)
Returns:
food_image: 提取的食物区域图像
food_mask: 食物区域mask
"""
# 食物相关类别(需要根据实际情况调整)
food_classes = [53, 116, 117, 118, 119, 120, 121, 122, 123]
# 提取mask
food_mask = self.extract_object_mask(pred_seg, food_classes)
# 应用mask
food_image = image.copy()
food_image[food_mask == 0] = 255 # 白色背景
return food_image, food_mask
def main():
"""
主函数:演示SegFormer-B0的基本用法
"""
print("=" * 50)
print("SegFormer-B0 快速上手Demo")
print("=" * 50)
# 1. 创建Demo实例
demo = SegFormerDemo()
# 2. 测试图像路径(请替换为您的图像路径)
test_image = "test_food3.jpg"
# 如果没有测试图像,创建一个示例
if not os.path.exists(test_image):
print(f"\n警告: 未找到测试图像 '{test_image}'")
print("请将您的食物图片重命名为 'test_food.jpg' 或修改代码中的路径")
print("\n提示: 您可以从以下位置获取测试图像:")
print(" - dataset目录中的训练图像")
print(" - 任何食物图片")
return
# 3. 进行分割
print(f"\n正在处理图像: {test_image}")
pred_seg, original_image = demo.segment_image(test_image)
# 4. 显示统计信息
unique_classes = np.unique(pred_seg)
print(f"\n检测到的类别数: {len(unique_classes)}")
print(f"类别ID: {unique_classes[:20]}...") # 显示前20个
# 5. 可视化结果
print("\n可视化分割结果...")
demo.visualize_segmentation(
original_image,
pred_seg,
save_path="segmentation_result.png"
)
# 6. 提取食物区域
print("\n提取食物区域...")
food_image, food_mask = demo.extract_food_region(original_image, pred_seg)
# 可视化食物区域
fig, axes = plt.subplots(1, 3, figsize=(15, 5))
axes[0].imshow(original_image)
axes[0].set_title("Original")
axes[0].axis("off")
axes[1].imshow(food_mask, cmap='gray')
axes[1].set_title("Food Mask")
axes[1].axis("off")
axes[2].imshow(food_image)
axes[2].set_title("Extracted Food Region")
axes[2].axis("off")
plt.tight_layout()
plt.savefig("food_extraction_result.png", dpi=150)
print("食物提取结果已保存到: food_extraction_result.png")
plt.show()
print("\n" + "=" * 50)
print("Demo运行完成!")
print("=" * 50)
if __name__ == "__main__":
main()
@@ -0,0 +1,303 @@
"""
COCO格式数据转换为SegFormer训练格式
功能说明:
1. 读取CVAT导出的COCO格式标注文件
2. 将polygon标注转换为像素级maskPNG图像)
3. 生成训练集和验证集的划分
4. 输出符合SegFormer训练要求的目录结构
输出目录结构:
data/segformer_format/
├── images/
│ ├── train/
│ │ ├── img1.jpg
│ │ └── img2.jpg
│ └── val/
│ └── img3.jpg
└── annotations/
├── train/
│ ├── img1.png # 像素值:0=背景, 1=食物区域
│ └── img2.png
└── val/
└── img3.png
"""
import os
import json
import numpy as np
from PIL import Image, ImageDraw
from pathlib import Path
from typing import List, Dict, Tuple
import shutil
class COCOToSegFormerConverter:
"""COCO格式到SegFormer格式的转换器"""
def __init__(self, coco_json_path: str, coco_images_dir: str, output_dir: str):
"""
初始化转换器
Args:
coco_json_path: COCO标注JSON文件路径(例如:instances_default.json
coco_images_dir: COCO图像所在目录
output_dir: 输出目录(将创建segformer_format文件夹)
"""
self.coco_json_path = coco_json_path
self.coco_images_dir = coco_images_dir
self.output_dir = Path(output_dir)
# 加载COCO标注
print(f"加载COCO标注文件: {coco_json_path}")
with open(coco_json_path, 'r', encoding='utf-8') as f:
self.coco_data = json.load(f)
print(f"✓ 图像数量: {len(self.coco_data['images'])}")
print(f"✓ 标注数量: {len(self.coco_data['annotations'])}")
print(f"✓ 类别数量: {len(self.coco_data['categories'])}")
# 创建输出目录结构
self._create_output_dirs()
def _create_output_dirs(self):
"""创建输出目录结构"""
dirs = [
self.output_dir / "images" / "train",
self.output_dir / "images" / "val",
self.output_dir / "annotations" / "train",
self.output_dir / "annotations" / "val",
]
for d in dirs:
d.mkdir(parents=True, exist_ok=True)
print(f"✓ 输出目录创建完成: {self.output_dir}")
def _polygon_to_mask(self, segmentation: List, image_size: Tuple[int, int]) -> np.ndarray:
"""
将COCO的polygon格式转换为像素级mask
Args:
segmentation: COCO的segmentation字段(polygon列表)
image_size: 图像尺寸 (width, height)
Returns:
mask: 二值mask数组 (H, W),1表示目标区域,0表示背景
"""
width, height = image_size
mask = Image.new('L', (width, height), 0) # 黑色背景
# COCO的segmentation可能包含多个polygon(例如:一个物体被遮挡分成多个部分)
for polygon in segmentation:
# polygon格式: [x1, y1, x2, y2, x3, y3, ...]
# 转换为坐标点列表: [(x1,y1), (x2,y2), ...]
points = [(polygon[i], polygon[i+1]) for i in range(0, len(polygon), 2)]
# 在mask上绘制填充的多边形
ImageDraw.Draw(mask).polygon(points, outline=1, fill=1)
return np.array(mask)
def _merge_annotations(self, image_id: int, image_size: Tuple[int, int]) -> np.ndarray:
"""
合并一张图像的所有标注为单一mask
由于用户标注时所有食材区域都是同一类别,我们需要将同一图像的多个标注合并
Args:
image_id: COCO图像ID
image_size: 图像尺寸 (width, height)
Returns:
merged_mask: 合并后的mask (H, W)
像素值: 0=背景(未标注区域), 1=食物区域
"""
width, height = image_size
merged_mask = np.zeros((height, width), dtype=np.uint8)
# 找到该图像的所有标注
annotations = [ann for ann in self.coco_data['annotations']
if ann['image_id'] == image_id]
# 将所有标注合并到同一个mask
for ann in annotations:
if 'segmentation' in ann and isinstance(ann['segmentation'], list):
# 转换polygon为mask
obj_mask = self._polygon_to_mask(ann['segmentation'], image_size)
# 合并到总mask(取并集)
merged_mask = np.maximum(merged_mask, obj_mask)
return merged_mask
def convert(self, train_ratio: float = 0.8, random_seed: int = 42):
"""
执行转换流程
Args:
train_ratio: 训练集占比(0.8表示80%训练,20%验证)
random_seed: 随机种子,确保每次划分一致
"""
print("\n" + "="*60)
print("开始转换数据集")
print("="*60)
# 创建图像ID到文件名的映射
id_to_image = {img['id']: img for img in self.coco_data['images']}
# 随机划分训练集和验证集
np.random.seed(random_seed)
image_ids = list(id_to_image.keys())
np.random.shuffle(image_ids)
split_idx = int(len(image_ids) * train_ratio)
train_ids = image_ids[:split_idx]
val_ids = image_ids[split_idx:]
print(f"\n数据集划分:")
print(f" 训练集: {len(train_ids)} 张图像")
print(f" 验证集: {len(val_ids)} 张图像")
# 处理训练集
print(f"\n处理训练集...")
self._process_split(train_ids, id_to_image, split='train')
# 处理验证集
print(f"\n处理验证集...")
self._process_split(val_ids, id_to_image, split='val')
# 保存数据集统计信息
self._save_dataset_info(train_ids, val_ids)
print("\n" + "="*60)
print("✓ 数据集转换完成!")
print("="*60)
print(f"\n输出目录: {self.output_dir}")
print("\n下一步: 运行 2_visualize_data.py 查看转换结果")
def _process_split(self, image_ids: List[int], id_to_image: Dict, split: str):
"""
处理训练集或验证集
Args:
image_ids: 图像ID列表
id_to_image: ID到图像信息的映射
split: 'train''val'
"""
for idx, img_id in enumerate(image_ids, 1):
img_info = id_to_image[img_id]
file_name = img_info['file_name']
width = img_info['width']
height = img_info['height']
print(f" [{idx}/{len(image_ids)}] {file_name}")
# 1. 复制图像到目标目录
src_image_path = Path(self.coco_images_dir) / file_name
dst_image_path = self.output_dir / "images" / split / file_name
if not src_image_path.exists():
print(f" ⚠️ 警告: 图像文件不存在 {src_image_path}")
continue
shutil.copy2(src_image_path, dst_image_path)
# 2. 生成mask并保存为PNG
mask = self._merge_annotations(img_id, (width, height))
# 保存mask(像素值即为类别ID:0=背景, 1=食物)
mask_filename = Path(file_name).stem + '.png' # 改为.png扩展名
mask_path = self.output_dir / "annotations" / split / mask_filename
# 使用PIL保存,确保像素值不被压缩
Image.fromarray(mask, mode='L').save(mask_path)
# 统计信息
food_pixels = np.sum(mask == 1)
total_pixels = mask.size
food_ratio = food_pixels / total_pixels * 100
print(f" ✓ 食物区域占比: {food_ratio:.1f}%")
def _save_dataset_info(self, train_ids: List[int], val_ids: List[int]):
"""保存数据集统计信息"""
info = {
"dataset_name": "Food Segmentation Dataset",
"num_classes": 2, # 背景 + 食物
"class_names": ["background", "food"],
"train_size": len(train_ids),
"val_size": len(val_ids),
"total_size": len(train_ids) + len(val_ids),
"image_format": "jpg/png",
"annotation_format": "png (pixel value = class id)",
}
info_path = self.output_dir / "dataset_info.json"
with open(info_path, 'w', encoding='utf-8') as f:
json.dump(info, f, indent=2, ensure_ascii=False)
print(f"\n✓ 数据集信息已保存: {info_path}")
def main():
"""
主函数:配置路径并执行转换
使用前请修改以下路径:
1. COCO_JSON_PATH: CVAT导出的JSON文件路径
2. COCO_IMAGES_DIR: CVAT导出的图像目录
3. OUTPUT_DIR: 输出目录(将创建segformer_format文件夹)
"""
# ==================== 配置区 ====================
# TODO: 请根据您的实际路径修改以下三个变量
# CVAT导出的COCO标注文件(通常名为instances_default.json
COCO_JSON_PATH = "d:/MyProjects/PythonProjects/FoodClassifier/SegFormer/data/raw_coco/annotations/instances_default.json"
# CVAT导出的图像目录
COCO_IMAGES_DIR = "d:/MyProjects/PythonProjects/FoodClassifier/SegFormer/data/raw_coco/images"
# 输出目录(将在此目录下创建segformer_format文件夹)
OUTPUT_DIR = "d:/MyProjects/PythonProjects/FoodClassifier/SegFormer/data/segformer_format"
# 训练集/验证集划分比例(0.8表示80%训练,20%验证)
TRAIN_RATIO = 0.8
# 随机种子(保证每次运行划分结果一致)
RANDOM_SEED = 42
# ===============================================
print("COCO数据集转换工具")
print("目标格式: SegFormer训练格式\n")
# 检查输入文件是否存在
if not os.path.exists(COCO_JSON_PATH):
print(f"❌ 错误: COCO标注文件不存在")
print(f" 路径: {COCO_JSON_PATH}")
print(f"\n请检查:")
print(f" 1. 是否已从CVAT导出COCO格式数据")
print(f" 2. 标注文件路径是否正确")
return
if not os.path.exists(COCO_IMAGES_DIR):
print(f"❌ 错误: 图像目录不存在")
print(f" 路径: {COCO_IMAGES_DIR}")
return
# 创建转换器并执行转换
converter = COCOToSegFormerConverter(
coco_json_path=COCO_JSON_PATH,
coco_images_dir=COCO_IMAGES_DIR,
output_dir=OUTPUT_DIR
)
converter.convert(train_ratio=TRAIN_RATIO, random_seed=RANDOM_SEED)
print("\n" + "="*60)
print("转换完成! 接下来的步骤:")
print("="*60)
print("1. 运行 2_visualize_data.py 检查转换结果")
print("2. 运行 3_train_minimal.py 开始训练")
if __name__ == "__main__":
main()
+389
View File
@@ -0,0 +1,389 @@
"""
数据集可视化工具
功能说明:
1. 可视化转换后的训练数据
2. 检查图像和mask是否正确对齐
3. 统计数据集的基本信息
4. 帮助发现标注错误
使用场景:
- 转换完成后,首先运行此脚本检查数据质量
- 训练前验证数据加载是否正确
"""
import os
import numpy as np
from PIL import Image
import matplotlib.pyplot as plt
from pathlib import Path
import random
class DatasetVisualizer:
"""数据集可视化工具"""
def __init__(self, data_root: str):
"""
初始化可视化工具
Args:
data_root: 数据集根目录(包含images和annotations文件夹)
"""
self.data_root = Path(data_root)
self.train_images_dir = self.data_root / "images" / "train"
self.train_masks_dir = self.data_root / "annotations" / "train"
self.val_images_dir = self.data_root / "images" / "val"
self.val_masks_dir = self.data_root / "annotations" / "val"
# 检查目录是否存在
if not self.data_root.exists():
raise FileNotFoundError(f"数据集目录不存在: {self.data_root}")
print(f"✓ 数据集根目录: {self.data_root}")
def get_dataset_stats(self):
"""获取数据集统计信息"""
print("\n" + "="*60)
print("数据集统计信息")
print("="*60)
# 统计图像数量
train_images = list(self.train_images_dir.glob("*.*"))
val_images = list(self.val_images_dir.glob("*.*"))
train_masks = list(self.train_masks_dir.glob("*.png"))
val_masks = list(self.val_masks_dir.glob("*.png"))
print(f"\n训练集:")
print(f" 图像数量: {len(train_images)}")
print(f" 标注数量: {len(train_masks)}")
print(f"\n验证集:")
print(f" 图像数量: {len(val_images)}")
print(f" 标注数量: {len(val_masks)}")
print(f"\n总计:")
print(f" 图像总数: {len(train_images) + len(val_images)}")
print(f" 标注总数: {len(train_masks) + len(val_masks)}")
# 统计mask中的类别分布
if len(train_masks) > 0:
print(f"\n正在分析mask内容...")
self._analyze_masks(train_masks[:3]) # 分析前3个mask
return {
'train': len(train_images),
'val': len(val_images),
'train_masks': len(train_masks),
'val_masks': len(val_masks)
}
def _analyze_masks(self, mask_paths: list):
"""
分析mask的像素值分布
Args:
mask_paths: mask文件路径列表
"""
for mask_path in mask_paths:
mask = np.array(Image.open(mask_path))
unique_values = np.unique(mask)
print(f"\n 文件: {mask_path.name}")
print(f" 尺寸: {mask.shape}")
print(f" 像素值: {unique_values}")
# 统计每个类别的像素数
for val in unique_values:
count = np.sum(mask == val)
ratio = count / mask.size * 100
class_name = "背景" if val == 0 else "食物"
print(f" {class_name}(类别{val}): {count}像素 ({ratio:.2f}%)")
def visualize_samples(self, split='train', num_samples=4, random_selection=True):
"""
可视化数据样本
Args:
split: 'train''val'
num_samples: 要可视化的样本数量
random_selection: 是否随机选择样本
"""
print(f"\n可视化{split}集样本...")
# 获取图像和mask路径
if split == 'train':
images_dir = self.train_images_dir
masks_dir = self.train_masks_dir
else:
images_dir = self.val_images_dir
masks_dir = self.val_masks_dir
# 获取所有图像文件
image_files = sorted(list(images_dir.glob("*.*")))
if len(image_files) == 0:
print(f" ⚠️ {split}集中没有图像文件")
return
# 选择要可视化的样本
if random_selection and len(image_files) > num_samples:
selected_files = random.sample(image_files, num_samples)
else:
selected_files = image_files[:num_samples]
# 创建画布
fig, axes = plt.subplots(num_samples, 3, figsize=(15, 5*num_samples))
if num_samples == 1:
axes = axes.reshape(1, -1)
for idx, image_path in enumerate(selected_files):
# 读取图像
image = Image.open(image_path).convert('RGB')
image_np = np.array(image)
# 读取对应的mask
mask_filename = image_path.stem + '.png'
mask_path = masks_dir / mask_filename
if not mask_path.exists():
print(f" ⚠️ 警告: 找不到mask文件 {mask_path.name}")
continue
mask = np.array(Image.open(mask_path))
# 显示原始图像
axes[idx, 0].imshow(image_np)
axes[idx, 0].set_title(f"原始图像\n{image_path.name}")
axes[idx, 0].axis('off')
# 显示mask(使用不同颜色)
# 0=黑色(背景), 1=白色(食物)
axes[idx, 1].imshow(mask, cmap='gray', vmin=0, vmax=1)
axes[idx, 1].set_title(f"分割Mask\n背景=黑, 食物=白")
axes[idx, 1].axis('off')
# 显示叠加效果
# 创建彩色mask用于叠加显示
colored_mask = np.zeros_like(image_np)
colored_mask[mask == 1] = [255, 0, 0] # 食物区域显示为红色
# 叠加显示
overlay = image_np.copy()
alpha = 0.4 # 透明度
overlay[mask == 1] = (
image_np[mask == 1] * (1 - alpha) +
colored_mask[mask == 1] * alpha
).astype(np.uint8)
axes[idx, 2].imshow(overlay)
axes[idx, 2].set_title("叠加显示\n红色=食物区域")
axes[idx, 2].axis('off')
# 打印统计信息
food_pixels = np.sum(mask == 1)
total_pixels = mask.size
food_ratio = food_pixels / total_pixels * 100
print(f" [{idx+1}] {image_path.name} - 食物占比: {food_ratio:.1f}%")
plt.tight_layout()
# 保存可视化结果
save_path = self.data_root / f"visualization_{split}.png"
plt.savefig(save_path, dpi=120, bbox_inches='tight')
print(f"\n✓ 可视化结果已保存: {save_path}")
plt.show()
def check_data_integrity(self):
"""
检查数据完整性
检查项目:
1. 每张图像是否有对应的mask
2. 图像和mask的尺寸是否匹配
3. mask的像素值是否在有效范围内
"""
print("\n" + "="*60)
print("数据完整性检查")
print("="*60)
issues = []
for split in ['train', 'val']:
print(f"\n检查{split}集...")
if split == 'train':
images_dir = self.train_images_dir
masks_dir = self.train_masks_dir
else:
images_dir = self.val_images_dir
masks_dir = self.val_masks_dir
image_files = list(images_dir.glob("*.*"))
for image_path in image_files:
# 检查1: mask文件是否存在
mask_filename = image_path.stem + '.png'
mask_path = masks_dir / mask_filename
if not mask_path.exists():
issues.append(f"{split}/{image_path.name}: 缺少mask文件")
continue
# 检查2: 尺寸是否匹配
image = Image.open(image_path)
mask = Image.open(mask_path)
if image.size != mask.size:
issues.append(
f"{split}/{image_path.name}: "
f"尺寸不匹配 (图像:{image.size}, mask:{mask.size})"
)
# 检查3: mask像素值是否有效
mask_np = np.array(mask)
unique_values = np.unique(mask_np)
# 有效值应该是0(背景)和1(食物)
invalid_values = [v for v in unique_values if v not in [0, 1]]
if invalid_values:
issues.append(
f"{split}/{image_path.name}: "
f"mask包含无效像素值 {invalid_values}"
)
# 输出检查结果
if len(issues) == 0:
print("\n✓ 数据完整性检查通过! 未发现问题")
else:
print(f"\n⚠️ 发现 {len(issues)} 个问题:")
for issue in issues:
print(f" - {issue}")
return len(issues) == 0
def show_class_distribution(self):
"""
显示类别分布统计
统计整个数据集中背景和食物的像素占比
"""
print("\n" + "="*60)
print("类别分布统计")
print("="*60)
for split in ['train', 'val']:
print(f"\n{split}集:")
if split == 'train':
masks_dir = self.train_masks_dir
else:
masks_dir = self.val_masks_dir
mask_files = list(masks_dir.glob("*.png"))
if len(mask_files) == 0:
print(f" 没有mask文件")
continue
# 统计所有mask的像素分布
total_background = 0
total_food = 0
for mask_path in mask_files:
mask = np.array(Image.open(mask_path))
total_background += np.sum(mask == 0)
total_food += np.sum(mask == 1)
total_pixels = total_background + total_food
print(f" 总像素数: {total_pixels:,}")
print(f" 背景像素: {total_background:,} ({total_background/total_pixels*100:.2f}%)")
print(f" 食物像素: {total_food:,} ({total_food/total_pixels*100:.2f}%)")
print(f" 类别平衡度: {min(total_background, total_food) / max(total_background, total_food):.3f}")
# 绘制饼图
fig, ax = plt.subplots(figsize=(8, 6))
ax.pie(
[total_background, total_food],
labels=['背景', '食物'],
autopct='%1.1f%%',
colors=['#808080', '#FF6B6B'],
startangle=90
)
ax.set_title(f'{split}集 - 类别分布')
save_path = self.data_root / f"class_distribution_{split}.png"
plt.savefig(save_path, dpi=120, bbox_inches='tight')
print(f" ✓ 分布图已保存: {save_path}")
plt.close()
def main():
"""
主函数:运行所有可视化和检查
"""
# ==================== 配置区 ====================
# TODO: 修改为您的数据集路径
DATA_ROOT = "d:/MyProjects/PythonProjects/FoodClassifier/SegFormer/data/segformer_format"
# 可视化参数
NUM_SAMPLES = 3 # 每个集合显示的样本数
RANDOM_SELECTION = False # True=随机选择, False=顺序选择前N个
# ===============================================
print("="*60)
print("数据集可视化工具")
print("="*60)
# 检查数据集路径
if not os.path.exists(DATA_ROOT):
print(f"\n❌ 错误: 数据集目录不存在")
print(f" 路径: {DATA_ROOT}")
print(f"\n请先运行 1_convert_coco_to_segformer.py 转换数据集")
return
# 创建可视化工具
visualizer = DatasetVisualizer(DATA_ROOT)
# 1. 显示数据集统计信息
stats = visualizer.get_dataset_stats()
# 2. 检查数据完整性
is_valid = visualizer.check_data_integrity()
if not is_valid:
print("\n⚠️ 请先修复数据问题再继续训练")
return
# 3. 显示类别分布
visualizer.show_class_distribution()
# 4. 可视化训练集样本
if stats['train'] > 0:
visualizer.visualize_samples(
split='train',
num_samples=min(NUM_SAMPLES, stats['train']),
random_selection=RANDOM_SELECTION
)
# 5. 可视化验证集样本
if stats['val'] > 0:
visualizer.visualize_samples(
split='val',
num_samples=min(NUM_SAMPLES, stats['val']),
random_selection=RANDOM_SELECTION
)
print("\n" + "="*60)
print("✓ 可视化完成!")
print("="*60)
print("\n如果数据没有问题,接下来可以:")
print(" 1. 运行 3_train_minimal.py 开始训练(CPU版本)")
print(" 2. 或运行 4_train_gpu.py 开始训练(GPU版本)")
if __name__ == "__main__":
main()
+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()
+344
View File
@@ -0,0 +1,344 @@
"""
SegFormer训练数据集类
功能说明:
1. 加载图像和对应的分割mask
2. 数据增强(翻转、旋转、颜色变换等)
3. 预处理和标准化
4. 批量数据加载
"""
import os
import numpy as np
from PIL import Image
import torch
from torch.utils.data import Dataset
from pathlib import Path
from typing import Optional, Tuple, List
import albumentations as A
from albumentations.pytorch import ToTensorV2
class FoodSegmentationDataset(Dataset):
"""
食物分割数据集
数据格式:
- 图像:RGB图像 (.jpg, .png等)
- 标注:PNG格式mask,像素值为类别ID (0=背景, 1=食物)
"""
def __init__(
self,
data_root: str,
split: str = 'train',
image_size: int = 512,
num_classes: int = 2,
augmentation: bool = True
):
"""
初始化数据集
Args:
data_root: 数据集根目录(包含images和annotations文件夹)
split: 'train''val'
image_size: 输入图像尺寸(将resize到此尺寸)
num_classes: 类别数(包括背景)
augmentation: 是否使用数据增强(仅训练集)
"""
self.data_root = Path(data_root)
self.split = split
self.image_size = image_size
self.num_classes = num_classes
self.augmentation = augmentation and (split == 'train')
# 构建图像和mask的路径
self.images_dir = self.data_root / "images" / split
self.masks_dir = self.data_root / "annotations" / split
# 获取所有图像文件
self.image_files = sorted(list(self.images_dir.glob("*.*")))
# 过滤:只保留有对应mask的图像
self.valid_samples = []
for img_path in self.image_files:
mask_path = self.masks_dir / (img_path.stem + '.png')
if mask_path.exists():
self.valid_samples.append((img_path, mask_path))
print(f"{split}集加载完成: {len(self.valid_samples)} 个样本")
# 构建数据增强pipeline
self.transform = self._build_transforms()
def _build_transforms(self):
"""
构建数据增强和预处理pipeline
使用albumentations库进行高效的数据增强
注意:对于分割任务,增强操作需要同时应用到图像和mask
"""
if self.augmentation:
# 训练集:激进的数据增强(因为数据量小)
transform = A.Compose([
# 1. 尺寸调整
A.Resize(self.image_size, self.image_size),
# 2. 几何变换(同时作用于图像和mask)
A.HorizontalFlip(p=0.5), # 50%概率水平翻转
A.VerticalFlip(p=0.3), # 30%概率垂直翻转
A.Rotate(limit=30, p=0.5), # ±30度旋转
A.ShiftScaleRotate(
shift_limit=0.1, # 平移±10%
scale_limit=0.2, # 缩放±20%
rotate_limit=20, # 旋转±20度
p=0.5
),
# 3. 颜色增强(仅作用于图像)
A.RandomBrightnessContrast(
brightness_limit=0.2,
contrast_limit=0.2,
p=0.5
),
A.HueSaturationValue(
hue_shift_limit=20,
sat_shift_limit=30,
val_shift_limit=20,
p=0.5
),
# 4. 模糊和噪声
A.OneOf([
A.GaussianBlur(blur_limit=(3, 5), p=1.0),
A.MedianBlur(blur_limit=5, p=1.0),
], p=0.3),
# 5. 标准化(使用ImageNet均值和标准差)
A.Normalize(
mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225],
),
# 6. 转换为Tensor
ToTensorV2(),
])
else:
# 验证集:仅resize和标准化
transform = A.Compose([
A.Resize(self.image_size, self.image_size),
A.Normalize(
mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225],
),
ToTensorV2(),
])
return transform
def __len__(self) -> int:
"""返回数据集大小"""
return len(self.valid_samples)
def __getitem__(self, idx: int) -> Tuple[torch.Tensor, torch.Tensor]:
"""
获取一个样本
Args:
idx: 样本索引
Returns:
image: 图像Tensor (3, H, W)
mask: 分割mask Tensor (H, W),像素值为类别ID
"""
img_path, mask_path = self.valid_samples[idx]
# 1. 读取图像和mask
image = np.array(Image.open(img_path).convert('RGB'))
mask = np.array(Image.open(mask_path))
# 2. 应用数据增强
# albumentations会自动将增强同时应用到image和mask
transformed = self.transform(image=image, mask=mask)
image = transformed['image'] # Tensor (3, H, W)
mask = transformed['mask'] # ndarray (H, W)
# 3. 将mask转换为Tensor
mask = torch.from_numpy(mask).long()
# 4. 检查mask的有效性
# 确保mask的值在[0, num_classes-1]范围内
if mask.max() >= self.num_classes:
print(f"警告: {mask_path.name} 包含无效类别ID {mask.max()}")
mask = torch.clamp(mask, 0, self.num_classes - 1)
return image, mask
def get_sample_info(self, idx: int) -> dict:
"""
获取样本的元信息(用于调试和可视化)
Args:
idx: 样本索引
Returns:
info: 包含文件名、路径等信息的字典
"""
img_path, mask_path = self.valid_samples[idx]
return {
'image_name': img_path.name,
'mask_name': mask_path.name,
'image_path': str(img_path),
'mask_path': str(mask_path),
}
def get_dataloaders(
data_root: str,
batch_size: int = 4,
image_size: int = 512,
num_workers: int = 0,
num_classes: int = 2
) -> Tuple[torch.utils.data.DataLoader, torch.utils.data.DataLoader]:
"""
创建训练集和验证集的DataLoader
Args:
data_root: 数据集根目录
batch_size: 批大小
image_size: 图像尺寸
num_workers: 数据加载进程数(CPU训练时设为0)
num_classes: 类别数
Returns:
train_loader: 训练集DataLoader
val_loader: 验证集DataLoader
"""
# 创建训练集
train_dataset = FoodSegmentationDataset(
data_root=data_root,
split='train',
image_size=image_size,
num_classes=num_classes,
augmentation=True # 训练集使用数据增强
)
# 创建验证集
val_dataset = FoodSegmentationDataset(
data_root=data_root,
split='val',
image_size=image_size,
num_classes=num_classes,
augmentation=False # 验证集不使用数据增强
)
# 创建DataLoader
train_loader = torch.utils.data.DataLoader(
train_dataset,
batch_size=batch_size,
shuffle=True, # 训练集打乱顺序
num_workers=num_workers,
pin_memory=False, # CPU训练时设为False
drop_last=True if len(train_dataset) > batch_size else False
)
val_loader = torch.utils.data.DataLoader(
val_dataset,
batch_size=batch_size,
shuffle=False, # 验证集不打乱
num_workers=num_workers,
pin_memory=False,
drop_last=False
)
print(f"\n✓ DataLoader创建完成")
print(f" 训练集: {len(train_dataset)} 样本, {len(train_loader)} 批次")
print(f" 验证集: {len(val_dataset)} 样本, {len(val_loader)} 批次")
return train_loader, val_loader
def test_dataset():
"""
测试数据集加载是否正常
用于开发调试,验证:
1. 数据能否正确加载
2. 数据增强是否正常工作
3. 数据的shape和类型是否正确
"""
import matplotlib.pyplot as plt
print("="*60)
print("数据集测试")
print("="*60)
# 配置
DATA_ROOT = "d:/MyProjects/PythonProjects/FoodClassifier/SegFormer/data/segformer_format"
# 创建数据集
dataset = FoodSegmentationDataset(
data_root=DATA_ROOT,
split='train',
image_size=512,
augmentation=True
)
# 测试读取第一个样本
image, mask = dataset[0]
info = dataset.get_sample_info(0)
print(f"\n样本信息:")
print(f" 文件名: {info['image_name']}")
print(f" 图像shape: {image.shape}") # 应该是 (3, 512, 512)
print(f" Mask shape: {mask.shape}") # 应该是 (512, 512)
print(f" Mask唯一值: {torch.unique(mask).numpy()}") # 应该是 [0, 1]
print(f" 图像数据范围: [{image.min():.3f}, {image.max():.3f}]")
# 可视化前3个样本(含数据增强效果)
fig, axes = plt.subplots(3, 3, figsize=(12, 12))
for i in range(3):
image, mask = dataset[i]
# 反标准化用于显示
mean = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)
std = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)
image_denorm = image * std + mean
image_denorm = torch.clamp(image_denorm, 0, 1)
# 转换为numpy用于显示
image_np = image_denorm.permute(1, 2, 0).numpy()
mask_np = mask.numpy()
# 显示图像
axes[i, 0].imshow(image_np)
axes[i, 0].set_title(f"样本{i+1}: 图像")
axes[i, 0].axis('off')
# 显示mask
axes[i, 1].imshow(mask_np, cmap='gray', vmin=0, vmax=1)
axes[i, 1].set_title(f"样本{i+1}: Mask")
axes[i, 1].axis('off')
# 显示叠加
overlay = image_np.copy()
red_mask = np.zeros_like(overlay)
red_mask[mask_np == 1] = [1, 0, 0]
overlay = overlay * 0.6 + red_mask * 0.4
axes[i, 2].imshow(overlay)
axes[i, 2].set_title(f"样本{i+1}: 叠加")
axes[i, 2].axis('off')
plt.tight_layout()
plt.savefig("dataset_test_result.png", dpi=120)
print(f"\n✓ 测试结果已保存: dataset_test_result.png")
plt.show()
print("\n" + "="*60)
print("✓ 数据集测试通过!")
print("="*60)
if __name__ == "__main__":
test_dataset()
+679
View File
@@ -0,0 +1,679 @@
"""
SegFormer Fine-tuning训练脚本(CPU版本)
功能说明:
1. 加载预训练的SegFormer模型
2. 在自定义数据集上Fine-tune
3. 支持两阶段训练(冻结Encoder + 全模型Fine-tune
4. 混合损失函数(CrossEntropy + Dice Loss
5. 完整的训练监控和模型保存
适用场景:
- 本地CPU开发调试
- 验证代码正确性
- 小数据集快速实验
pip install numpy==1.26.4
目前是2.2.6
"""
import os
import sys
import time
import json
import numpy as np
from pathlib import Path
from tqdm import tqdm
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.optim import AdamW
from torch.optim.lr_scheduler import CosineAnnealingLR
from transformers import SegformerForSemanticSegmentation
import matplotlib.pyplot as plt
# 导入自定义模块
from config import TrainConfig, get_cpu_config
from dataset import get_dataloaders
class DiceLoss(nn.Module):
"""
Dice Loss for semantic segmentation
Dice系数衡量预测和真实mask的重叠程度:
Dice = 2 * |A ∩ B| / (|A| + |B|)
Dice Loss = 1 - Dice
优点:
- 对类别不平衡不敏感
- 关注重叠区域,适合分割任务
"""
def __init__(self, smooth=1.0):
"""
Args:
smooth: 平滑因子,防止除零
"""
super(DiceLoss, self).__init__()
self.smooth = smooth
def forward(self, logits, targets):
"""
Args:
logits: 模型输出 (B, C, H, W)
targets: 真实标签 (B, H, W)
Returns:
loss: Dice Loss标量
"""
# 1. 将logits转换为概率
probs = F.softmax(logits, dim=1) # (B, C, H, W)
# 2. 将targets转换为one-hot编码
num_classes = logits.shape[1]
targets_one_hot = F.one_hot(targets, num_classes) # (B, H, W, C)
targets_one_hot = targets_one_hot.permute(0, 3, 1, 2).float() # (B, C, H, W)
# 3. 展平空间维度
probs = probs.view(probs.shape[0], probs.shape[1], -1) # (B, C, H*W)
targets_one_hot = targets_one_hot.view(targets_one_hot.shape[0], targets_one_hot.shape[1], -1)
# 4. 计算Dice系数
intersection = (probs * targets_one_hot).sum(dim=2) # (B, C)
union = probs.sum(dim=2) + targets_one_hot.sum(dim=2) # (B, C)
dice = (2.0 * intersection + self.smooth) / (union + self.smooth) # (B, C)
# 5. 对所有类别和batch取平均
dice_loss = 1.0 - dice.mean()
return dice_loss
class MixedLoss(nn.Module):
"""
混合损失:CrossEntropy + Dice Loss
结合两种损失的优点:
- CrossEntropy: 像素级分类,梯度稳定
- Dice Loss: 关注重叠区域,处理类别不平衡
"""
def __init__(self, num_classes, class_weights=None, dice_weight=0.5):
"""
Args:
num_classes: 类别数
class_weights: 类别权重(用于CrossEntropy
dice_weight: Dice Loss的权重
"""
super(MixedLoss, self).__init__()
# CrossEntropy Loss
if class_weights is not None:
class_weights = torch.tensor(class_weights, dtype=torch.float32)
self.ce_loss = nn.CrossEntropyLoss(weight=class_weights)
# Dice Loss
self.dice_loss = DiceLoss()
self.dice_weight = dice_weight
def forward(self, logits, targets):
"""
Args:
logits: 模型输出 (B, C, H, W)
targets: 真实标签 (B, H, W)
Returns:
loss: 混合损失
loss_dict: 各损失分量的字典
"""
ce = self.ce_loss(logits, targets)
dice = self.dice_loss(logits, targets)
total_loss = ce + self.dice_weight * dice
return total_loss, {'ce_loss': ce.item(), 'dice_loss': dice.item()}
class SegmentationMetrics:
"""
分割任务的评估指标
包括:
- mIoU (Mean Intersection over Union)
- Pixel Accuracy
- Per-class IoU
"""
def __init__(self, num_classes):
"""
Args:
num_classes: 类别数
"""
self.num_classes = num_classes
self.reset()
def reset(self):
"""重置所有统计"""
self.confusion_matrix = np.zeros((self.num_classes, self.num_classes))
def update(self, preds, targets):
"""
更新混淆矩阵
Args:
preds: 预测标签 (B, H, W) 或 (N,)
targets: 真实标签 (B, H, W) 或 (N,)
"""
preds = preds.flatten()
targets = targets.flatten()
# 过滤无效值
mask = (targets >= 0) & (targets < self.num_classes)
preds = preds[mask]
targets = targets[mask]
# 更新混淆矩阵
for pred, target in zip(preds, targets):
self.confusion_matrix[target, pred] += 1
def compute_miou(self):
"""
计算mIoU (Mean Intersection over Union)
IoU = TP / (TP + FP + FN)
mIoU = mean(IoU_per_class)
Returns:
miou: 平均IoU
iou_per_class: 每个类别的IoU
"""
# 对角线是TP(预测和真实都是该类别)
tp = np.diag(self.confusion_matrix)
# 每行之和是该类别的真实数量(TP + FN)
# 每列之和是预测为该类别的数量(TP + FP)
fp_fn = self.confusion_matrix.sum(axis=0) + self.confusion_matrix.sum(axis=1) - tp
# IoU = TP / (TP + FP + FN)
iou_per_class = tp / (fp_fn + 1e-10)
# 忽略无效类别(分母为0
valid_classes = fp_fn > 0
miou = iou_per_class[valid_classes].mean()
return miou, iou_per_class
def compute_pixel_accuracy(self):
"""
计算像素准确率
Pixel Accuracy = (TP + TN) / Total
Returns:
pixel_acc: 像素准确率
"""
total_correct = np.diag(self.confusion_matrix).sum()
total_pixels = self.confusion_matrix.sum()
pixel_acc = total_correct / (total_pixels + 1e-10)
return pixel_acc
def get_results(self):
"""
获取所有评估指标
Returns:
results: 包含所有指标的字典
"""
miou, iou_per_class = self.compute_miou()
pixel_acc = self.compute_pixel_accuracy()
return {
'miou': miou,
'pixel_accuracy': pixel_acc,
'iou_per_class': iou_per_class.tolist(),
}
class SegFormerTrainer:
"""SegFormer训练器"""
def __init__(self, config: TrainConfig):
"""
初始化训练器
Args:
config: 训练配置
"""
self.config = config
self.device = torch.device(config.device)
# 设置随机种子
self._set_seed(config.random_seed)
# 创建输出目录
self.output_dir = Path(config.output_dir)
self.output_dir.mkdir(parents=True, exist_ok=True)
# 加载数据
print("\n" + "="*60)
print("加载数据集")
print("="*60)
self.train_loader, self.val_loader = get_dataloaders(
data_root=config.data_root,
batch_size=config.batch_size,
image_size=config.image_size,
num_workers=config.num_workers,
num_classes=config.num_classes
)
# 加载模型
print("\n" + "="*60)
print("加载模型")
print("="*60)
self.model = self._build_model()
# 构建损失函数
self.criterion = self._build_criterion()
# 构建优化器和调度器
self.optimizer = self._build_optimizer()
self.scheduler = self._build_scheduler()
# 训练状态
self.current_epoch = 0
self.best_metric = 0.0
self.train_losses = []
self.val_metrics = []
print("\n" + "="*60)
print("✓ 训练器初始化完成")
print("="*60)
def _set_seed(self, seed):
"""设置随机种子"""
torch.manual_seed(seed)
np.random.seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
def _build_model(self):
"""
构建SegFormer模型
Returns:
model: SegFormer模型
"""
print(f"加载预训练模型: {self.config.pretrained_model}")
# 加载预训练模型
model = SegformerForSemanticSegmentation.from_pretrained(
self.config.pretrained_model,
num_labels=self.config.num_classes,
ignore_mismatched_sizes=True # 允许分类头尺寸不匹配
)
model.to(self.device)
# 统计参数量
total_params = sum(p.numel() for p in model.parameters())
trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
print(f"✓ 模型加载完成")
print(f" 总参数量: {total_params:,}")
print(f" 可训练参数: {trainable_params:,}")
return model
def _build_criterion(self):
"""构建损失函数"""
if self.config.use_mixed_loss:
criterion = MixedLoss(
num_classes=self.config.num_classes,
class_weights=self.config.class_weights,
dice_weight=self.config.dice_loss_weight
)
print(f"✓ 损失函数: Mixed Loss (CE + {self.config.dice_loss_weight}×Dice)")
else:
weights = None
if self.config.class_weights:
weights = torch.tensor(self.config.class_weights, dtype=torch.float32)
criterion = nn.CrossEntropyLoss(weight=weights)
print(f"✓ 损失函数: CrossEntropy")
return criterion.to(self.device)
def _build_optimizer(self):
"""构建优化器"""
if self.config.optimizer_type == 'adamw':
optimizer = AdamW(
self.model.parameters(),
lr=self.config.learning_rate,
weight_decay=self.config.weight_decay,
betas=self.config.adam_betas
)
print(f"✓ 优化器: AdamW (lr={self.config.learning_rate})")
else:
optimizer = torch.optim.SGD(
self.model.parameters(),
lr=self.config.learning_rate,
weight_decay=self.config.weight_decay,
momentum=self.config.sgd_momentum
)
print(f"✓ 优化器: SGD (lr={self.config.learning_rate})")
return optimizer
def _build_scheduler(self):
"""构建学习率调度器"""
if self.config.lr_scheduler_type == 'cosine':
scheduler = CosineAnnealingLR(
self.optimizer,
T_max=self.config.num_epochs,
eta_min=1e-7
)
print(f"✓ 学习率调度: Cosine Annealing")
else:
scheduler = None
print(f"✓ 学习率调度: None")
return scheduler
def _freeze_encoder(self):
"""冻结Encoder(仅训练Decoder"""
for name, param in self.model.named_parameters():
if 'decode_head' not in name:
param.requires_grad = False
trainable = sum(p.numel() for p in self.model.parameters() if p.requires_grad)
print(f" 冻结Encoder,可训练参数: {trainable:,}")
def _unfreeze_all(self):
"""解冻所有参数"""
for param in self.model.parameters():
param.requires_grad = True
trainable = sum(p.numel() for p in self.model.parameters() if p.requires_grad)
print(f" 解冻全模型,可训练参数: {trainable:,}")
def train_epoch(self):
"""
训练一个epoch
Returns:
avg_loss: 平均损失
"""
self.model.train()
epoch_loss = 0.0
epoch_ce_loss = 0.0
epoch_dice_loss = 0.0
pbar = tqdm(self.train_loader, desc=f"Epoch {self.current_epoch}/{self.config.num_epochs}")
for batch_idx, (images, masks) in enumerate(pbar):
# 1. 数据迁移到设备
images = images.to(self.device)
masks = masks.to(self.device)
# 2. 前向传播
outputs = self.model(pixel_values=images)
logits = outputs.logits
# 3. 上采样logits到mask尺寸
logits = F.interpolate(
logits,
size=masks.shape[-2:],
mode='bilinear',
align_corners=False
)
# 4. 计算损失
if isinstance(self.criterion, MixedLoss):
loss, loss_dict = self.criterion(logits, masks)
epoch_ce_loss += loss_dict['ce_loss']
epoch_dice_loss += loss_dict['dice_loss']
else:
loss = self.criterion(logits, masks)
# 5. 反向传播
self.optimizer.zero_grad()
loss.backward()
# 6. 梯度裁剪
if self.config.max_grad_norm > 0:
torch.nn.utils.clip_grad_norm_(
self.model.parameters(),
self.config.max_grad_norm
)
# 7. 更新参数
self.optimizer.step()
# 8. 累计损失
epoch_loss += loss.item()
# 9. 更新进度条
pbar.set_postfix({'loss': f'{loss.item():.4f}'})
# 计算平均损失
avg_loss = epoch_loss / len(self.train_loader)
if isinstance(self.criterion, MixedLoss):
avg_ce = epoch_ce_loss / len(self.train_loader)
avg_dice = epoch_dice_loss / len(self.train_loader)
print(f" Train Loss: {avg_loss:.4f} (CE: {avg_ce:.4f}, Dice: {avg_dice:.4f})")
else:
print(f" Train Loss: {avg_loss:.4f}")
return avg_loss
@torch.no_grad()
def validate(self):
"""
验证模型
Returns:
metrics: 评估指标字典
"""
self.model.eval()
metrics = SegmentationMetrics(self.config.num_classes)
val_loss = 0.0
pbar = tqdm(self.val_loader, desc="Validation")
for images, masks in pbar:
# 1. 数据迁移
images = images.to(self.device)
masks = masks.to(self.device)
# 2. 前向传播
outputs = self.model(pixel_values=images)
logits = outputs.logits
# 3. 上采样
logits = F.interpolate(
logits,
size=masks.shape[-2:],
mode='bilinear',
align_corners=False
)
# 4. 计算损失
if isinstance(self.criterion, MixedLoss):
loss, _ = self.criterion(logits, masks)
else:
loss = self.criterion(logits, masks)
val_loss += loss.item()
# 5. 计算预测
preds = logits.argmax(dim=1)
# 6. 更新指标
metrics.update(preds.cpu().numpy(), masks.cpu().numpy())
# 计算最终指标
results = metrics.get_results()
results['loss'] = val_loss / len(self.val_loader)
print(f"\n Validation Results:")
print(f" Loss: {results['loss']:.4f}")
print(f" mIoU: {results['miou']:.4f}")
print(f" Pixel Acc: {results['pixel_accuracy']:.4f}")
print(f" IoU per class: {[f'{x:.3f}' for x in results['iou_per_class']]}")
return results
def save_checkpoint(self, metrics, is_best=False):
"""
保存checkpoint
Args:
metrics: 评估指标
is_best: 是否是最佳模型
"""
checkpoint = {
'epoch': self.current_epoch,
'model_state_dict': self.model.state_dict(),
'optimizer_state_dict': self.optimizer.state_dict(),
'metrics': metrics,
'config': self.config.to_dict(),
}
# 保存最新checkpoint
checkpoint_path = self.output_dir / f"checkpoint_epoch_{self.current_epoch}.pt"
torch.save(checkpoint, checkpoint_path)
print(f" ✓ Checkpoint已保存: {checkpoint_path}")
# 保存最佳模型
if is_best:
best_path = self.output_dir / "best_model.pt"
torch.save(checkpoint, best_path)
print(f" ✓ 最佳模型已保存: {best_path}")
def train(self):
"""完整训练流程"""
print("\n" + "="*60)
print("开始训练")
print("="*60)
self.config.print_config()
# 两阶段训练
if self.config.use_two_stage_training:
print("\n【Stage 1: 冻结Encoder,训练Decoder】")
self._freeze_encoder()
stage1_end = self.config.stage1_epochs
else:
stage1_end = 0
# 主训练循环
for epoch in range(1, self.config.num_epochs + 1):
self.current_epoch = epoch
# 切换到Stage 2
if self.config.use_two_stage_training and epoch == stage1_end + 1:
print("\n" + "="*60)
print("【Stage 2: Fine-tune全模型】")
print("="*60)
self._unfreeze_all()
# 重建优化器(因为参数变了)
self.optimizer = self._build_optimizer()
self.scheduler = self._build_scheduler()
print(f"\n{'='*60}")
print(f"Epoch {epoch}/{self.config.num_epochs}")
print(f"{'='*60}")
# 训练
train_loss = self.train_epoch()
self.train_losses.append(train_loss)
# 学习率调度
if self.scheduler:
self.scheduler.step()
print(f" Learning Rate: {self.optimizer.param_groups[0]['lr']:.2e}")
# 验证
if epoch % self.config.eval_every_n_epochs == 0:
print(f"\n执行验证...")
val_metrics = self.validate()
self.val_metrics.append((epoch, val_metrics))
# 判断是否是最佳模型
current_metric = val_metrics[self.config.best_metric]
is_best = current_metric > self.best_metric
if is_best:
self.best_metric = current_metric
print(f" 🎉 新的最佳{self.config.best_metric}: {current_metric:.4f}")
# 保存checkpoint
if epoch % self.config.save_every_n_epochs == 0:
if self.config.save_best_only and not is_best:
print(f" 跳过保存(非最佳模型)")
else:
self.save_checkpoint(val_metrics, is_best=is_best)
print("\n" + "="*60)
print("✓ 训练完成!")
print("="*60)
print(f"最佳{self.config.best_metric}: {self.best_metric:.4f}")
# 保存训练曲线
self.plot_training_curves()
def plot_training_curves(self):
"""绘制训练曲线"""
fig, axes = plt.subplots(1, 2, figsize=(15, 5))
# 损失曲线
axes[0].plot(range(1, len(self.train_losses) + 1), self.train_losses)
axes[0].set_xlabel('Epoch')
axes[0].set_ylabel('Training Loss')
axes[0].set_title('Training Loss Curve')
axes[0].grid(True)
# mIoU曲线
if len(self.val_metrics) > 0:
epochs = [x[0] for x in self.val_metrics]
mious = [x[1]['miou'] for x in self.val_metrics]
axes[1].plot(epochs, mious, marker='o')
axes[1].set_xlabel('Epoch')
axes[1].set_ylabel('mIoU')
axes[1].set_title('Validation mIoU Curve')
axes[1].grid(True)
plt.tight_layout()
save_path = self.output_dir / "training_curves.png"
plt.savefig(save_path, dpi=120)
print(f"\n✓ 训练曲线已保存: {save_path}")
plt.close()
def main():
"""主函数"""
print("="*60)
print("SegFormer Fine-tuning训练脚本")
print("="*60)
# 获取配置
config = get_cpu_config()
# 根据需要修改配置
# config.num_epochs = 20
# config.image_size = 256
# config.eval_every_n_epochs = 5
# 创建训练器
trainer = SegFormerTrainer(config)
# 开始训练
trainer.train()
print("\n训练完成! 接下来可以:")
print(" 1. 运行 5_test_model.py 测试模型")
print(" 2. 查看输出目录中的训练曲线和checkpoint")
if __name__ == "__main__":
main()