Files
FoodClassifier/data_management/split_dataset.py
T
2026-02-06 03:00:50 -05:00

368 lines
15 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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
)