增加图像分割相应的代码
This commit is contained in:
@@ -0,0 +1,303 @@
|
||||
"""
|
||||
COCO格式数据转换为SegFormer训练格式
|
||||
|
||||
功能说明:
|
||||
1. 读取CVAT导出的COCO格式标注文件
|
||||
2. 将polygon标注转换为像素级mask(PNG图像)
|
||||
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()
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user