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

如何实现PyTorch多索引Dataset并适配DataLoader加载?

RNN分类任务中自定义Dataset的__len__与DataLoader加载问题

一、__len__方法能否返回多个长度?

不能。PyTorch的Dataset规范明确要求__len__必须返回单个整数,代表数据集的总样本数。你当前返回(序列数量, 序列长度)元组的写法不符合规范,会导致DataLoader在确定采样范围时抛出错误。

建议调整方案:

  • 将__len__改回返回序列总数:return len(self._all_flexions)
  • 把序列长度作为Dataset的类属性暴露,比如在__init__中添加self.sequence_length = self._sequence_length,后续需要时直接通过实例访问该属性。

二、如何用DataLoader加载该Dataset?

由于你的__getitem__需要接收序列索引+时间索引的组合(字典或元组形式),而默认Sampler只会生成单个整数索引,因此需要自定义采样逻辑,以下是两种可行实现:

方式一:自定义Sampler生成元组索引

自定义Sampler每次生成(sequence_idx, frames_slice)的元组,直接适配你现有__getitem__的逻辑:

import torch
from torch.utils.data import Sampler

class RNNSampler(Sampler):
    def __init__(self, num_sequences, sequence_length, total_frames_per_seq):
        self.num_sequences = num_sequences
        self.sequence_length = sequence_length
        # total_frames_per_seq:每个序列的总帧数列表,需提前统计
        self.total_frames = total_frames_per_seq

    def __iter__(self):
        for seq_idx in range(self.num_sequences):
            max_start = self.total_frames[seq_idx] - self.sequence_length
            if max_start <= 0:
                # 序列长度不足时取全部帧(可根据需求调整)
                yield (seq_idx, slice(0, self.total_frames[seq_idx]))
            else:
                # 随机采样起始帧(也可改为遍历所有可能片段)
                start_idx = torch.randint(0, max_start + 1, (1,)).item()
                yield (seq_idx, slice(start_idx, start_idx + self.sequence_length))

    def __len__(self):
        # 返回总样本数,这里按每个序列生成1个样本,可根据需求修改
        return self.num_sequences

使用示例:

# 实例化Dataset
dataset = MLDataWrangler(...)
# 统计每个序列的总帧数(根据你的数据结构调整)
total_frames = [self.read(path).data.zsig.shape[0] for path in dataset._all_flexions]
# 创建Sampler和DataLoader
sampler = RNNSampler(
    num_sequences=len(dataset._all_flexions),
    sequence_length=dataset._sequence_length,
    total_frames_per_seq=total_frames
)
dataloader = torch.utils.data.DataLoader(dataset, sampler=sampler, batch_size=4)

方式二:预生成所有样本索引映射

提前把所有(序列索引, 时间片段)的组合列出来,让__getitem__接收单个整数索引,无需自定义Sampler:

class MLDataWrangler(zrfr.ZRFReader, torch.utils.data.Dataset):
    def __init__(self, ...):
        # 原有初始化逻辑
        ...
        # 预生成所有样本的索引映射
        self.sample_map = []
        for seq_idx in range(len(self._all_flexions)):
            zrf_path = self._all_flexions[seq_idx]
            total_frames = self.read(zrf_path).data.zsig.shape[0]
            max_start = total_frames - self._sequence_length
            if max_start <= 0:
                self.sample_map.append((seq_idx, slice(0, total_frames)))
            else:
                # 遍历所有可能的时间片段(也可改为随机采样存储)
                for start in range(max_start + 1):
                    self.sample_map.append((seq_idx, slice(start, start + self._sequence_length)))

    def __len__(self) -> int:
        return len(self.sample_map)

    def __getitem__(self, idx) -> Tuple[np.ndarray, np.ndarray]:
        # 通过整数索引获取序列和时间片段
        sequence, frames = self.sample_map[idx]
        # 原有数据读取与处理逻辑
        zrf_path = self._all_flexions[sequence]
        video_pixels = self.read(zrf_path).data
        ...
        # 后续信号提取、标签处理逻辑不变
        ...

这种方式的优势是可以直接使用默认DataLoader,无需额外自定义Sampler,适合样本数量固定的场景。

注意事项

  • 若RNN需要固定长度输入,需确保所有采样的时间片段长度一致,避免DataLoader批量拼接时出错。
  • 自定义Sampler时,__len__返回的数值必须与实际生成的样本数一致,否则会出现采样不完整或重复的问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 23:50:23