Files
FoodClassifier/SegFormer/scripts/2_visualize_data.py
T

390 lines
13 KiB
Python

"""
数据集可视化工具
功能说明:
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()