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

PyTorch自定义Dataset返回元组时的切片行为疑问

问题描述

我编写了一个EventDetectionDataset类,它的__getitem__方法返回tuple[list[str], list[str]]类型的数据。但执行train[:2]切片操作时,得到的是tuple[list[list[str]], list[list[str]]]类型结果,而非预期的list[tuple[list[str]], list[str]]。请问这个切片行为是不是把元组元素分别拼接成列表?为什么会有这样的设计?

数据集类代码

from torch.utils.data import Dataset
import json

def read_dataset(path: str) -> tuple[list[list[str]], list[list[str]]]:
    tokens_s, labels_s = [], []

    with open(path) as f:
        for line in f:
            data = json.loads(line)
            assert len(data["tokens"]) == len(data["labels"])
            tokens_s.append(data["tokens"])
            labels_s.append(data["labels"])

    assert len(tokens_s) == len(labels_s)

    return tokens_s, labels_s


class EventDetectionDataset(Dataset):
    def __init__(self, path: str) -> None:
        self.tokens, self.labels = read_dataset(path)

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

    def __getitem__(self, index) -> tuple[list[str], list[str]]:
        return self.tokens[index], self.labels[index]

执行切片的代码

train = EventDetectionDataset("path/to/data/train.jsonl")
train = train[:2]

实际返回结果

(
    [
        ['Hard', 'Rock', 'Hell', 'III', ':', 'The', 'Vikings', 'Ball', '.'],
        ['Casualties', 'and', 'damage', 'were', 'severe', 'on', 'both', 'sides', ',', 'and', 'the', 'defiance', 'of', 'the', 'French', 'ship', 'was', 'celebrated', 'in', 'both', 'countries', 'as', 'a', 'brave', 'defence', 'against', 'overwhelming', 'odds', '.']
    ],

     [
        ['O', 'O', 'O', 'O', 'O', 'O', 'O', 'O', 'O'],
        ['B-SCENARIO', 'O', 'B-CHANGE', 'O', 'O', 'O', 'O', 'O', 'O', 'O', 'O', 'O', 'O', 'O', 'O', 'O', 'O', 'B-ACTION', 'O', 'O', 'O', 'O', 'O', 'O', 'O', 'O', 'O', 'O', 'O']
    ]
)
解答

切片行为的本质

是的,这个切片操作确实是把元组的两个元素(tokens和labels)分别拼接成了大列表,最终返回一个包含两个大列表的元组。

设计原因

这是PyTorch Dataset类的默认行为:当处理切片或批量索引时,内部会遍历切片范围内的每个索引,调用__getitem__获取单个样本的元组,然后自动对所有返回的元组做转置处理——把所有元组的第一个元素收集成一个列表,第二个元素收集成另一个列表,最终组合成新的元组。

这种设计是为了适配PyTorch的数据加载逻辑:模型训练时通常需要批量格式的输入(比如所有样本的tokens组成一个batch,labels组成另一个batch),而非单个样本的元组列表。这样的批量处理方式能直接对接DataLoader,高效转换为模型可使用的张量格式,省去额外的格式转换步骤。

如何得到预期格式

如果确实需要由样本元组组成的列表,可以手动遍历索引:

train = EventDetectionDataset("path/to/data/train.jsonl")
expected_result = [train[i] for i in range(2)]

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 15:24:55