训练集和验证集怎么保持类别比例?用 Datasets 分层切分并核对

将类别列编码为 ClassLabel,以 Datasets 分层切分训练与验证,并核对类别比例、样本 ID 交集及可重复性。

二分类训练数据类别不均衡时,可以用 Datasets 的 train_test_split(stratify_by_column=”label”) 做分层切分,再检查两部分的类别数。这样能减少简单随机切分造成的类别比例波动;样本仍须独立,分层本身不能防止同用户或同文档泄漏。

适用前提:一条样本只有一个类别

本例有 20 条虚构记录,其中 negative 16 条、positive 4 条;预留 25% 作为验证集。在当前版本里,分层列必须是 ClassLabel 类型。普通字符串类别可以先用 class_encode_column() 编码;类别名称与数字的对应关系要读取 features,不能凭直觉假定某个数字代表正类。

训练集和验证集怎么保持类别比例?用 Datasets 分层切分并核对

示例于 2026-10-01 在 Windows、Python 3.11.15、Datasets 4.8.3 环境实际执行通过;数据全部人为构造,只验证数据处理接口,没有训练模型或测量业务效果。安装可参考 Datasets 官方安装说明。先在新的练习目录建立环境:

python -m venv .venv
# Windows 使用 .venv\Scripts\python.exe;macOS/Linux 使用 .venv/bin/python
.venv\Scripts\python.exe -m pip install "datasets==4.8.3"

下文运行命令按 Windows 写法给出。使用 macOS 或 Linux 时,把解释器路径换为 .venv/bin/python;后续版本请先用同一份玩具数据复核,再替换生产数据。

完整切分代码

官方接口参考列出 stratify_by_column、seed 和 class_encode_column();4.8.3 官方项目源码检查 ClassLabel 并拒绝 shuffle=False 的分层切分。将下列代码存为 split_demo.py,再运行 .venv\Scripts\python.exe split_demo.py。

from collections import Counter
from datasets import Dataset, ClassLabel

rows = [
    {'id': i, 'text': f'样本{i}',
     'label': 'negative' if i < 16 else 'positive'}
    for i in range(20)
]
ds = Dataset.from_list(rows).class_encode_column('label')
assert isinstance(ds.features['label'], ClassLabel)
parts = ds.train_test_split(
    test_size=0.25, stratify_by_column='label', seed=42
)

def counts(part):
    feature = ds.features['label']
    raw = Counter(part['label'])
    return {name: raw[index] for index, name in enumerate(feature.names)}

print('classes:', ds.features['label'].names)
print('train:', counts(parts['train']))
print('validation:', counts(parts['test']))
assert counts(parts['train']) == {'negative': 12, 'positive': 3}
assert counts(parts['test']) == {'negative': 4, 'positive': 1}
assert len(parts['train']) + len(parts['test']) == 20
assert set(parts['train']['id']).isdisjoint(parts['test']['id'])
again = ds.train_test_split(
    test_size=0.25, stratify_by_column='label', seed=42
)
assert list(again['test']['id']) == list(parts['test']['id'])
print('counts, ID separation and repeatability: passed')

返回键仍叫 train 和 test;本文把 test 当作开发中的验证集使用。如果要额外保留最终测试集,应另建独立划分,并固定其用途,避免反复依据测试结果调参。

应该看到什么结果

这组数据实测输出的类别映射为 [‘negative’, ‘positive’],训练部分是 negative 12、positive 3,验证部分是 negative 4、positive 1。两部分的 positive 比例都是 20%,行数分别为 15 和 5。

断言还验证两部分 ID 没有交集,以及同一输入顺序、版本和 seed 下再次切分的验证 ID 一致。把玩具数据换成正式数据前,先确认 ID 真的唯一;两份表里 ID 不同但文本完全相同,仍可能泄漏。

本例数字刚好可以按比例分配;实际类别数不能整除时,需要整数取整,不能要求所有类别比例完全一致。应打印每个 split 的类别计数,并核对少数类是否在各部分有足够样本支持评估。

常见报错与失效边界

  • 提示只支持 ClassLabel:打印 ds.features[‘label’],先编码或按已知标签定义转换,编码后保留 names 对照。
  • 某类别只有一条:无法同时进入训练和验证;先补充同类独立样本,或明确调整评估设计,不能复制同一条到两边。
  • 验证行数比类别数少:每类都要出现时空间不足,调整 test_size 并重新打印计数;训练部分也需足够容纳类别。
  • 时间预测、同用户多条记录、同文档多个片段:应按时间或来源组划分;对这些任务直接随机分层可能高估模型效果。
  • 多标签、连续回归目标:不能把列表或每个不同浮点值直接当成普通 ClassLabel 分层列。

完整处理流程可查 Datasets 处理文档。开始训练前,把输入版本、标签映射、seed、各 split 的 ID 和计数一起保存;固定随机种子只能复现同一份输入,不能保证换了数据后划分仍相同。

Ai菜鸟网。发布者:AI小管家,转载请注明出处:https://www.alyyhw.com/30285.html

赞 (0)
AI小管家的头像AI小管家
智能体知识库怎么同步更新?用 Chroma 完成覆盖写入与旧片段删除
上一篇 2天前
pgvector HNSW 索引怎么验收?对照精确检索检查召回与查询计划
下一篇 2天前

相关推荐

联系我们

联系我们

1

在线咨询: QQ交谈

邮件:admin@example.com

工作时间:周一至周五,9:30-18:30,节假日休息

关注微信
关注微信
分享本页
返回顶部