如何基于索引列表获取Python arrow_dataset Dataset的子集
Arrow Dataset 按索引列表获取数据子集方法
针对datasets.arrow_dataset.Dataset类型的数据集对象,可通过以下两种方式按给定索引列表提取对应子集:
- 直接下标索引
直接将目标索引列表作为下标传入数据集对象即可,写法最简洁。注意:如果传入单个整数索引,返回的是单条样本的字典对象;传入索引列表时才会返回Dataset类型的子集:
# 假设ds是你的Dataset对象,target_idx是存储目标索引的列表 target_idx = [1, 5, 10, 23, 44] subset = ds[target_idx]
该方法支持传入包含重复值的索引列表,返回子集也会对应保留重复的样本条目。
- 内置
select()方法
这是官方推荐的显式子集提取方法,除列表外还支持传入生成器、numpy数组等任意可迭代的整数索引对象,大索引量场景下内存效率更高:
# 传入列表索引 subset = ds.select(target_idx) # 传入索引生成器,无需提前把全量索引加载到内存 def idx_gen(): # 示例:每隔50条取1条样本 for i in range(0, len(ds), 50): yield i subset_skip50 = ds.select(idx_gen)
注意:传入的所有索引值必须落在
[0, len(ds)-1]的合法区间内,否则会触发索引越界错误。如果需要按样本特征条件筛选而非固定索引提取,可使用ds.filter()方法实现。
内容的提问来源于stack exchange,提问作者scout4321
相关产品推荐
相关产品推荐

