You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

交叉验证(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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.30 23:27:04