393 lines
14 KiB
Python
393 lines
14 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
|
|
|
|
# 设置matplotlib中文字体
|
|
plt.rcParams['font.sans-serif'] = ['SimHei', 'Microsoft YaHei', 'Arial Unicode MS'] # 用来正常显示中文标签
|
|
plt.rcParams['axes.unicode_minus'] = False # 用来正常显示负号
|
|
|
|
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()
|