diff --git a/data_management/split_dataset.py b/data_management/split_dataset.py new file mode 100644 index 0000000..a4b96c3 --- /dev/null +++ b/data_management/split_dataset.py @@ -0,0 +1,367 @@ +"""cui +数据集分割脚本 +1. 将train目录下的图片按70:15:15分割到train/val/test +2. 对样本量过少的类别进行标记 +""" +import os +import shutil +import random +from pathlib import Path +from typing import List, Tuple + + +def split_dataset( + dataset_root: str, + train_ratio: float = 0.70, + val_ratio: float = 0.15, + test_ratio: float = 0.15, + min_samples_threshold: int = 10, + mark_insufficient: bool = True, + max_total: int = 120, + seed: int = 42 +): + """ + 分割数据集 + + Args: + dataset_root: 数据集根目录(包含train目录) + train_ratio: 训练集比例(默认0.70) + val_ratio: 验证集比例(默认0.15) + test_ratio: 测试集比例(默认0.15) + min_samples_threshold: 最小样本数阈值,低于此值的类别会被标记(默认10) + mark_insufficient: 是否在文件夹名称后添加"样本数太少"标记(默认True) + max_total: 如果总数超过100,限制到max_total张(默认120) + seed: 随机种子(默认42) + """ + # 验证比例 + if abs(train_ratio + val_ratio + test_ratio - 1.0) > 0.001: + raise ValueError(f"比例总和必须为1.0,当前为: {train_ratio + val_ratio + test_ratio}") + + dataset_root = Path(dataset_root) + train_dir = dataset_root / 'train' + val_dir = dataset_root / 'val' + test_dir = dataset_root / 'test' + + if not train_dir.exists(): + raise ValueError(f"训练目录不存在: {train_dir}") + + # 创建val和test目录 + val_dir.mkdir(parents=True, exist_ok=True) + test_dir.mkdir(parents=True, exist_ok=True) + + # 设置随机种子 + random.seed(seed) + + print("="*80) + print(f"数据集分割工具") + print("="*80) + print(f"数据集根目录: {dataset_root}") + print(f"分割比例: 训练集={train_ratio*100:.0f}%, 验证集={val_ratio*100:.0f}%, 测试集={test_ratio*100:.0f}%") + print(f"最小样本数阈值: {min_samples_threshold}") + print(f"样本不足标记: {'启用' if mark_insufficient else '禁用'}") + print("="*80 + "\n") + + # 检查val和test目录是否已有数据 + val_has_data = val_dir.exists() and any(val_dir.iterdir()) + test_has_data = test_dir.exists() and any(test_dir.iterdir()) + + if val_has_data or test_has_data: + print("⚠ 检测到验证集或测试集已有数据") + print("将检查所有类别的比例,如果不符合70:15:15将重新分配\n") + + # 获取所有类别文件夹 + class_folders = [f for f in train_dir.iterdir() if f.is_dir()] + total_classes = len(class_folders) + + print(f"找到 {total_classes} 个类别\n") + + stats = { + 'total_classes': total_classes, + 'processed_classes': 0, + 'skipped_classes': 0, + 'sufficient_samples': 0, + 'insufficient_samples': 0, + 'total_images': 0, + 'train_images': 0, + 'val_images': 0, + 'test_images': 0, + 'marked_classes': [] + } + + # 处理每个类别 + for class_folder in sorted(class_folders): + class_name = class_folder.name + + # 收集所有图片(从train/val/test) + train_images = [] + for ext in ['*.jpg', '*.jpeg', '*.png', '*.JPG', '*.JPEG', '*.PNG']: + train_images.extend(list(class_folder.glob(ext))) + + val_class_dir = val_dir / class_name + val_images = [] + if val_class_dir.exists(): + for ext in ['*.jpg', '*.jpeg', '*.png', '*.JPG', '*.JPEG', '*.PNG']: + val_images.extend(list(val_class_dir.glob(ext))) + + test_class_dir = test_dir / class_name + test_images = [] + if test_class_dir.exists(): + for ext in ['*.jpg', '*.jpeg', '*.png', '*.JPG', '*.JPEG', '*.PNG']: + test_images.extend(list(test_class_dir.glob(ext))) + + # 合并所有图片 + all_images = train_images + val_images + test_images + num_images = len(all_images) + stats['total_images'] += num_images + + if num_images == 0: + print(f"⚠ {class_name}: 0张图片,跳过") + continue + + # 检查当前比例是否正确 + current_train = len(train_images) + current_val = len(val_images) + current_test = len(test_images) + + # 计算期望的数量 + expected_train = int(num_images * train_ratio) + expected_val = int(num_images * val_ratio) + expected_test = num_images - expected_train - expected_val + + # 检查比例是否接近(允许1-2张的误差) + ratio_correct = ( + abs(current_train - expected_train) <= 2 and + abs(current_val - expected_val) <= 2 and + abs(current_test - expected_test) <= 2 + ) + + # 如果总数超过100,即使比例正确也需要限制 + if ratio_correct and num_images > 0 and num_images <= 100: + # 比例正确且总数<=100,跳过 + stats['skipped_classes'] += 1 + is_insufficient = num_images < min_samples_threshold + if is_insufficient: + stats['insufficient_samples'] += 1 + else: + stats['sufficient_samples'] += 1 + print(f"✓ {class_name}: 总数={num_images}, 比例正确 (train={current_train}, val={current_val}, test={current_test}),跳过") + continue + + # 如果比例正确但总数>100,需要限制 + if ratio_correct and num_images > 100: + print(f"处理 {class_name}: 总数={num_images} > 100, 比例正确但需要限制到{max_total}张") + stats['processed_classes'] += 1 + # 随机选择max_total张 + random.shuffle(all_images) + image_files = all_images[:max_total] + num_images = len(image_files) + # 重新计算期望数量 + expected_train = int(num_images * train_ratio) + expected_val = int(num_images * val_ratio) + expected_test = num_images - expected_train - expected_val + print(f" 将重新分配为: train={expected_train}, val={expected_val}, test={expected_test}") + elif not ratio_correct: + # 比例不正确,需要重新分配 + stats['processed_classes'] += 1 + print(f"处理 {class_name}: 总数={num_images}, 当前比例不正确 (train={current_train}, val={current_val}, test={current_test})") + + # 如果总数超过100,限制到max_total + if num_images > 100: + print(f" 总数={num_images} > 100, 限制到{max_total}张") + random.shuffle(all_images) + image_files = all_images[:max_total] + num_images = len(image_files) + # 重新计算期望数量 + expected_train = int(num_images * train_ratio) + expected_val = int(num_images * val_ratio) + expected_test = num_images - expected_train - expected_val + else: + image_files = all_images + print(f" 将重新分配为: train={expected_train}, val={expected_val}, test={expected_test}") + + + # 检查样本数是否足够 + is_insufficient = num_images < min_samples_threshold + + if is_insufficient: + stats['insufficient_samples'] += 1 + if mark_insufficient: + # 标记样本不足的类别 + marked_name = f"{class_name}_样本数太少" + stats['marked_classes'].append((class_name, num_images, marked_name)) + else: + stats['sufficient_samples'] += 1 + + # 确定最终的类别名称(如果需要标记) + final_class_name = f"{class_name}_样本数太少" if (is_insufficient and mark_insufficient) else class_name + + # 如果还没打乱(总数<=100的情况),现在打乱 + if num_images == len(all_images): + random.shuffle(image_files) + + # 计算分割点 + n_train = int(num_images * train_ratio) + n_val = int(num_images * val_ratio) + # 剩余的全部给test,确保所有图片都被分配 + + # 分割图片(重新分配) + new_train_images = image_files[:n_train] + new_val_images = image_files[n_train:n_train + n_val] + new_test_images = image_files[n_train + n_val:] + + # 创建目标目录 + train_target = train_dir / final_class_name + val_target = val_dir / final_class_name + test_target = test_dir / final_class_name + + train_target.mkdir(exist_ok=True) + val_target.mkdir(exist_ok=True) + test_target.mkdir(exist_ok=True) + + # 先移动图片到正确的位置(在删除之前) + moved_to_train = 0 + for img in new_train_images: + if not img.exists(): + continue + target = train_target / img.name + if img.parent != train_target: + if target.exists(): + target.unlink() + try: + shutil.move(str(img), str(target)) + moved_to_train += 1 + except Exception as e: + print(f" ⚠ 移动失败 {img}: {e}") + + moved_to_val = 0 + for img in new_val_images: + if not img.exists(): + continue + target = val_target / img.name + if img.parent != val_target: + if target.exists(): + target.unlink() + try: + shutil.move(str(img), str(target)) + moved_to_val += 1 + except Exception as e: + print(f" ⚠ 移动失败 {img}: {e}") + + moved_to_test = 0 + for img in new_test_images: + if not img.exists(): + continue + target = test_target / img.name + if img.parent != test_target: + if target.exists(): + target.unlink() + try: + shutil.move(str(img), str(target)) + moved_to_test += 1 + except Exception as e: + print(f" ⚠ 移动失败 {img}: {e}") + + # 然后清理所有目录(删除不在新分配列表中的图片) + # 收集所有原始图片路径 + original_images = set(train_images + val_images + test_images) + new_images = set(new_train_images + new_val_images + new_test_images) + images_to_delete = original_images - new_images + + for img in images_to_delete: + if img.exists(): + try: + img.unlink() + except Exception as e: + pass + + # 处理样本不足标记(如果需要) + if mark_insufficient and is_insufficient: + # 需要重命名train目录的文件夹 + if class_folder != train_target: + # 如果原文件夹还存在且不为空,需要移动剩余文件 + if class_folder.exists() and any(class_folder.iterdir()): + for item in class_folder.iterdir(): + if item.is_file(): + shutil.move(str(item), train_target / item.name) + # 删除原文件夹(如果为空) + try: + if class_folder.exists() and not any(class_folder.iterdir()): + class_folder.rmdir() + except Exception as e: + print(f" 警告: 删除原文件夹失败: {e}") + + # 更新统计 + stats['train_images'] += len(new_train_images) + stats['val_images'] += len(new_val_images) + stats['test_images'] += len(new_test_images) + + # 打印进度 + status = "⚠样本不足" if is_insufficient else "✓" + print(f"{status} {class_name}: 总计={num_images}, 训练={len(new_train_images)}, 验证={len(new_val_images)}, 测试={len(new_test_images)}") + + # 打印总结 + print("\n" + "="*80) + print("分割完成!") + print("="*80) + print(f"总类别数: {stats['total_classes']}") + print(f"已处理: {stats['processed_classes']} 个类别(比例不正确,已重新分配)") + print(f"已跳过: {stats['skipped_classes']} 个类别(比例正确)") + print(f"样本充足: {stats['sufficient_samples']} 个类别") + print(f"样本不足: {stats['insufficient_samples']} 个类别") + print(f"\n总图片数: {stats['total_images']}") + print(f"训练集: {stats['train_images']} 张 ({stats['train_images']/stats['total_images']*100:.1f}%)") + print(f"验证集: {stats['val_images']} 张 ({stats['val_images']/stats['total_images']*100:.1f}%)") + print(f"测试集: {stats['test_images']} 张 ({stats['test_images']/stats['total_images']*100:.1f}%)") + + if stats['marked_classes']: + print(f"\n已标记的类别(样本数 < {min_samples_threshold}):") + for old_name, count, new_name in stats['marked_classes']: + print(f" {old_name} ({count}张) -> {new_name}") + + print("\n✓ 完成!") + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser(description='数据集分割工具') + parser.add_argument('--dataset', type=str, + default='dataset/WholeIngredientRecognition', + help='数据集根目录(默认: dataset/WholeIngredientRecognition)') + parser.add_argument('--train-ratio', type=float, default=0.70, + help='训练集比例(默认: 0.70)') + parser.add_argument('--val-ratio', type=float, default=0.15, + help='验证集比例(默认: 0.15)') + parser.add_argument('--test-ratio', type=float, default=0.15, + help='测试集比例(默认: 0.15)') + parser.add_argument('--min-samples', type=int, default=10, + help='最小样本数阈值,低于此值会被标记(默认: 10)') + parser.add_argument('--no-mark', action='store_true', + help='不标记样本不足的类别') + parser.add_argument('--seed', type=int, default=42, + help='随机种子(默认: 42)') + parser.add_argument('--max-total', type=int, default=120, + help='如果总数超过100,限制到max_total张(默认: 120)') + + args = parser.parse_args() + + # 获取项目根目录 + script_dir = Path(__file__).parent + project_root = script_dir.parent + + # 处理相对路径 + if not os.path.isabs(args.dataset): + dataset_path = project_root / args.dataset + else: + dataset_path = Path(args.dataset) + + split_dataset( + dataset_root=str(dataset_path), + train_ratio=args.train_ratio, + val_ratio=args.val_ratio, + test_ratio=args.test_ratio, + min_samples_threshold=args.min_samples, + mark_insufficient=not args.no_mark, + max_total=args.max_total, + seed=args.seed + )