同一用户可能留下多条评论、交易或传感器记录。如果按行随机切分,模型可能在训练时见过这个人的特征,在测试时又遇到他的另一条记录。你想评估“能否处理新用户”,这种划分就会使答案偏乐观。这里用 GroupShuffleSplit 把一个用户的全部记录放在同一侧,并保留原始行编号方便核查。
先确定分组字段和测试目标
每行需要特征、标签、稳定样本编号和分组编号。分组字段可以是用户、患者、设备或文档来源,取决于哪些记录相互关联。缺失组号不能默认为每行一个新组:先查来源,无法恢复的样本单独隔离;错误的分组定义会让后面的断言看起来通过。

下面构造 12 位虚构用户,每位两行,共 24 行;标签只为演示切分。用户编号没有作为模型特征。这不是实际用户数据,也不会训练或评估一个模型。
生成分组切分并检查交集
以下示例在 Python 3.11、scikit-learn 1.7.2 的独立环境运行通过。先在练习目录创建虚拟环境并安装对应版本:
python -m venv .venv
# Windows PowerShell
.\.venv\Scripts\python.exe -m pip install "scikit-learn==1.7.2"
# 以下脚本统一用该环境的 Python 运行,例如:
.\.venv\Scripts\python.exe demo.py
from collections import Counter
import numpy as np
from sklearn.model_selection import GroupShuffleSplit
# 虚构数据:12 位用户,每人两条记录。用户编号独立保存,不作为模型特征。
users = np.repeat(np.arange(12), 2)
row_ids = np.array([f"row-{i:02d}" for i in range(24)])
X = np.column_stack([np.arange(24), np.arange(24) % 3])
y = np.repeat(np.arange(12) % 2, 2)
splitter = GroupShuffleSplit(n_splits=1, test_size=0.25, random_state=42)
train_idx, test_idx = next(splitter.split(X, y, groups=users))
train_users, test_users = set(users[train_idx]), set(users[test_idx])
assert train_users.isdisjoint(test_users)
assert set(train_idx).isdisjoint(test_idx)
assert len(train_idx) + len(test_idx) == len(y)
assert set(y[train_idx]) == {0, 1} and set(y[test_idx]) == {0, 1}
print("训练用户", sorted(train_users), "测试用户", sorted(test_users))
print("训练行数", len(train_idx), "测试行数", len(test_idx))
print("训练标签", dict(Counter(y[train_idx])))
print("测试标签", dict(Counter(y[test_idx])))
print("测试原始编号", row_ids[test_idx].tolist())
把脚本保存为 demo.py 后执行,当前环境输出:
训练用户 [np.int64(1), np.int64(2), np.int64(3), np.int64(4), np.int64(5), np.int64(6), np.int64(7), np.int64(8), np.int64(11)] 测试用户 [np.int64(0), np.int64(9), np.int64(10)]
训练行数 18 测试行数 6
训练标签 {np.int64(1): 10, np.int64(0): 8}
测试标签 {np.int64(0): 4, np.int64(1): 2}
测试原始编号 ['row-00', 'row-01', 'row-18', 'row-19', 'row-20', 'row-21']
关键验证是训练用户集合与测试用户集合的交集为空,训练/测试行编号不重叠,两侧行数之和等于输入。最后打印两边的类别分布;本例有两类,不代表所有真实切分都能保留每类。
比例针对组,不能当作行数比例
GroupShuffleSplit 的 test_size=0.25 指抽取组的比例。本例每组两行,所以也恰好拿到 6 条测试记录;若一个用户有 100 行、其他用户只有 1 行,行数比例可能明显不同。它不负责按标签分层,遇到某类只来自一个用户时,无法同时做到用户不交叉和两边都有该类。
n_splits=1 只生成一次划分。设成多次随机切分时,不同轮次的测试组可能重复;不要把它当作每组只出现一次的完整 K 折。需要组间交叉验证可研究 GroupKFold,兼顾分组与类别比例可研究 StratifiedGroupKFold,并核对各折实际分布。
把索引用到真实特征表
确认 X、y、users 与 row_ids 来自同一顺序的表,再用 X[train_idx]、y[train_idx] 训练,用 X[test_idx]、y[test_idx] 评估。将 row_ids 和组号随划分清单保存;编号不会因后续排序改变。预处理也要只在训练部分拟合,分组正确不代表预处理自动没有泄漏。
遇到长度不一致、组号为空、只有一个组或某一侧缺类别,应先修正数据或调整评估方案,不能删除断言继续。时间序列还需要按时间安排训练与测试,随机组切分不能保证训练只使用过去。
相关问答与依据
同一用户有不同标签怎么办?仍按用户分组;不要用标签拼成新组号,否则同一个用户会跨集合。若要评估老用户的未来行为,应重新定义时间与用户边界,不能直接拿本例当作正式方案。
分组后分数降低是否说明模型退步?仅凭分数不能判断。先确认新划分与上线的新用户场景一致,再比较同样评估定义下的模型。本例只验证分组程序,没有报告业务效果。
依据:GroupShuffleSplit 官方参数文档说明组比例与跨轮次重叠;交叉验证官方指南说明关联样本、组划分与时间序列方案。
Ai菜鸟网。发布者:AI小管家,转载请注明出处:https://www.alyyhw.com/30105.html