【问题标题】:Iterable dataset exhausts after a single epoch [IMBD dataset from torchtext]可迭代数据集在一个时期后耗尽[来自 torchtext 的 IMBD 数据集]
【发布时间】:2021-06-19 19:24:01
【问题描述】:

我想在情感分析任务上训练一个 RNN,对于这个任务,我使用了由 torchtext 提供的 IMDB 数据集,其中包含 50000 条电影评论,它是一个 Python 迭代器。我用了split=('train', 'test')

我首先使用torchtext.vocab.Vocab 构建了一个词汇表,并对每个句子进行了标记,然后进行了数字化。

为了将序列填充到相同的长度,我使用了torch.nn.utils.rnn.pad_sequence,还使用了collate_fnbatch_sampler。然后我使用 torch.utils.data.DataLoader 加载数据。

RNN 网络的实现很好,但数据加载器在一个 epoch 后就耗尽了,如下图所示。

我是否采用了正确的方法来加载这个可迭代数据集?以及为什么数据加载器在一个时期后耗尽,我该如何克服这个问题。

如果您想查看我的实现,请参阅共享的 colab 笔记本。

PS。我在关注来自github的torchtext官方changelog

你可以找到我的实现here

在附图中,您可以看到数据加载器在单个 epoch 后耗尽。

【问题讨论】:

    标签: python iterator pytorch torchtext


    【解决方案1】:

    问题是您的数据加载器是一个生成器,并且在完全迭代后耗尽。一种解决方案是在每个时期初始化数据加载器。二是不要使用批量采样器。整理功能应该做你想做的事。

    def collate_batch(batch):
        batch.sort(key=lambda x: len(x[0]), reverse=True)
        label_list, text_list, text_lengths = [], [], []
        
        for _text, _label in batch:
            label_list.append(_label)
            processed_text = torch.tensor(_text)
            text_list.append(processed_text)
            text_lengths.append(len(processed_text))
    
        return torch.tensor(label_list, dtype=torch.float32),
               pad_sequence(text_list, padding_value=3.0), 
               torch.tensor(text_lengths, dtype=torch.int64)
    

    【讨论】:

      猜你喜欢
      • 2021-09-03
      • 2021-04-30
      • 2016-10-13
      • 2016-12-25
      • 2019-04-24
      • 2020-08-03
      • 1970-01-01
      • 2018-02-07
      • 1970-01-01
      相关资源
      最近更新 更多