PyTorch数据集dmap用法及相关代码功能含义问询
PyTorch数据集处理代码功能说明
代码整体作用
这是半监督学习场景下常用的数据集预处理逻辑,基于olive-oil-ml工具库实现数据集拆分和无标签数据集构造。
逐行解释
- 第一行代码
datasets = split_dataset(dataset(), splits=split)
调用split_dataset函数,将dataset()返回的完整原始数据集,按照splits参数指定的比例(如训练/验证/测试集7:2:1)切分,最终返回字典结构的结果,字典的键一般为train/val/test,分别对应切分后的训练集、验证集、测试集。 - 第二行代码
datasets['_unlab'] = dmap(lambda mb: mb[0], dataset())
这行是构造半监督学习需要的无标签数据集:dmap是数据集映射工具,作用是对第二个参数传入的数据集里的每一个样本,依次应用第一个参数指定的转换函数,返回转换后的新数据集lambda mb: mb[0]这个匿名函数的作用是,取每个原始样本的第一个元素:常规PyTorch数据集的每个样本默认是(特征张量, 标签值)的二元组,这个函数会只保留特征部分,扔掉标签- 最终生成的全量无标签数据集会被存入
datasets字典的_unlab键,供后续半监督训练调用
内容的提问来源于stack exchange,提问作者user25004
相关产品推荐
相关产品推荐

