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

在PyTorch中如何根据指定索引列表从数据集批量提取对应样本

PyTorch指定索引取批次样本的实现方案

下面两种实现都是基于原生PyTorch组件,无额外依赖,兼顾简洁性和效率。

方案1:一次性取固定索引批次(最简便)

直接用Subset包装原Dataset,套普通DataLoader即可,不需要额外封装:

from torch.utils.data import DataLoader, Subset

# 替换为你自己的Dataset实例
my_dataset = YourCustomDataset()
list_idxs = [10, 109, 7, 12]

# 生成指定索引的子集,按索引顺序取数
subset = Subset(my_dataset, list_idxs)
# batch_size设为索引长度,一次返回所有指定样本
loader = DataLoader(subset, batch_size=len(list_idxs), num_workers=2)
batch = next(iter(loader))

方案2:封装为动态调用工具(符合getbatch调用要求)

如果需要频繁传入不同索引列表取批次,可以简单封装一个工具类,完全复用DataLoader的所有特性(多进程加载、自动批处理、pin_memory等):

from torch.utils.data import DataLoader, Dataset, Subset
from typing import List

class IndexedBatchLoader:
    def __init__(self, dataset: Dataset, **dataloader_kwargs):
        self.dataset = dataset
        # 支持传入所有DataLoader支持的参数,比如num_workers、pin_memory、collate_fn等
        self.dataloader_kwargs = dataloader_kwargs
    
    def getbatch(self, idxs: List[int]):
        subset = Subset(self.dataset, idxs)
        loader = DataLoader(subset, batch_size=len(idxs), **self.dataloader_kwargs)
        return next(iter(loader))

使用示例:

# 初始化加载器
loader = IndexedBatchLoader(my_dataset, num_workers=2, pin_memory=True)
# 任意传入索引列表即可得到对应批次
list_idxs = [10, 109, 7, 12]
batch = loader.getbatch(list_idxs)

轻量极简实现(适合单进程场景)

如果不需要多进程加载,也可以直接手动取样本后用默认的批处理函数拼接,开销更低:

from torch.utils.data import default_collate

def getbatch(dataset: Dataset, idxs: List[int]):
    return default_collate([dataset[i] for i in idxs])

# 调用方法
batch = getbatch(my_dataset, [10, 109, 7, 12])

注意事项

  • 所有方案返回的样本顺序和传入的索引列表顺序完全一致
  • 如果你的Dataset需要做重预处理,建议选择带多进程参数的方案,效率提升更明显
  • 不需要修改原有Dataset的任何代码,对所有符合PyTorch规范的Dataset都生效

内容的提问来源于stack exchange,提问作者thebesttony

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 02:18:00