交叉验证(cross validation)时如何基于患者文件夹拆分数据集,保证单患者数据不跨集?
解决方案
你需要的是按患者层级的分组拆分逻辑,完全可以用scikit-learn提供的分组拆分工具实现,天然满足「同一患者所有样本仅归属训练/测试其中一个集合」的要求,也能避免不同患者样本量不均带来的拆分偏差。
实现思路
- 首先遍历数据集目录,收集所有图片的路径、对应类别、所属患者ID三个核心信息,分组键采用
类别_患者ID的格式,避免不同类别下出现同名患者被误判为同一组。 - 用
GroupKFold做交叉验证拆分,拆分维度是患者组,不是单个图片。如果只需要单次训练/测试集拆分,替换为GroupShuffleSplit即可。 - 拆分后支持生成索引txt文件(省空间,训练时直接读取),也可以直接把文件复制到对应拆分目录。
代码实现
import os import shutil from sklearn.model_selection import GroupKFold, GroupShuffleSplit # 自定义配置参数 DATA_ROOT = "替换为你的原始数据集根目录路径" OUTPUT_ROOT = "替换为拆分结果的输出目录路径" # 交叉验证折数,按需调整 N_FOLD = 5 # 单次拆分的测试集比例,仅用GroupShuffleSplit时生效 TEST_RATIO = 0.2 RANDOM_SEED = 42 # 第一步:遍历数据集,收集全量样本信息 sample_paths = [] group_ids = [] labels = [] for class_name in os.listdir(DATA_ROOT): class_dir = os.path.join(DATA_ROOT, class_name) if not os.path.isdir(class_dir): continue # 遍历当前类别下的所有患者目录 for patient_id in os.listdir(class_dir): patient_dir = os.path.join(class_dir, patient_id) if not os.path.isdir(patient_dir): continue # 拼接全局唯一的分组ID,避免不同类别同名患者冲突 global_patient_id = f"{class_name}_{patient_id}" # 遍历患者下所有图片 for img_name in os.listdir(patient_dir): if img_name.lower().endswith((".png", ".jpg", ".jpeg")): sample_paths.append(os.path.join(class_dir, patient_dir, img_name)) group_ids.append(global_patient_id) labels.append(class_name) # 第二步:执行拆分 # 场景1:K折交叉验证拆分 gkf = GroupKFold(n_splits=N_FOLD) for fold_idx, (train_idx, test_idx) in enumerate(gkf.split(sample_paths, labels, group_ids)): fold_output_dir = os.path.join(OUTPUT_ROOT, f"fold_{fold_idx + 1}") os.makedirs(fold_output_dir, exist_ok=True) # 输出训练集索引 with open(os.path.join(fold_output_dir, "train.txt"), "w", encoding="utf-8") as f: for idx in train_idx: f.write(f"{sample_paths[idx]} {labels[idx]}\n") # 输出测试集索引 with open(os.path.join(fold_output_dir, "test.txt"), "w", encoding="utf-8") as f: for idx in test_idx: f.write(f"{sample_paths[idx]} {labels[idx]}\n") # 打印校验信息,确认患者无重叠 train_patients = set([group_ids[i] for i in train_idx]) test_patients = set([group_ids[i] for i in test_idx]) print(f"第{fold_idx+1}折:训练集患者数{len(train_patients)},测试集患者数{len(test_patients)},重叠患者数{len(train_patients & test_patients)}") # 场景2:单次训练/测试集拆分,不需要交叉验证时用这个替换上面的拆分逻辑即可 # gss = GroupShuffleSplit(n_splits=1, test_size=TEST_RATIO, random_state=RANDOM_SEED) # train_idx, test_idx = next(gss.split(sample_paths, labels, group_ids)) # 后续输出txt/复制文件的逻辑和交叉验证场景一致,这里省略
注意事项
- 代码默认生成索引txt文件,如果你需要直接复制图片到对应训练/测试目录,只需要在拆分后遍历索引,按照原来的
类别/患者/图片的目录结构复制文件即可,逻辑和写txt完全一致。 - 拆分后打印的重叠患者数必须为0,才符合你的需求。
内容的提问来源于stack exchange,提问作者Rawan
相关产品推荐
相关产品推荐

