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