【问题标题】:Broken Pipe Error while using Loading Text Data in Pytorch在 Pytorch 中使用加载文本数据时出现断管错误
【发布时间】:2020-10-16 19:43:23
【问题描述】:

我正在尝试预处理一些文本数据,但在创建 pytorch 数据加载器并循环检查它是否正常工作后,我收到了 Broken Pipe 错误。但是,在 Google Colab 中再次尝试时,代码可以正常工作,所以我认为这可能是我的设置有问题。

(Collat​​e 类没用,我只是还没有删除它。)

import numpy as np
import pandas as pd

data = pd.read_csv("imdb.csv")

import torch
import torch.nn as nn
import torch.optim as optim
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader

import spacy
spacy_eng = spacy.load("en")

class Vocabulary():
    def __init__(self, freq_threshold=4):
        self.word_to_index = {"<PAD>":0, "<SOS>":1, "<EOS>":2, "<UNK>":3}
        
        self.freq_threshold = freq_threshold
        self.max_length = 0
    
    def __len__(self):
        return len(self.word_to_index)
    
    @staticmethod
    def tokenizer_eng(text):
        return [tok.text.lower() for tok in spacy_eng.tokenizer(text)]
    
    def build_vocabulary(self, sentence_list):
        frequencies = {}
        idx = 4
        
        longest_length = 0
        
        for sentence in sentence_list:
            if len(sentence) > longest_length:
                self.max_length = len(sentence)
                longest_length = self.max_length
                
            for word in self.tokenizer_eng(sentence):
                if word not in frequencies:
                    frequencies[word] = 1
                else:
                    frequencies[word] += 1
                
                if frequencies[word] == self.freq_threshold:
                    self.word_to_index[word] = idx
                    idx += 1
        
        self.max_length += 25
    
    def numericalize(self, text):
        tokenized_text = self.tokenizer_eng(text)
        
        vector_text = []
        
        for token in tokenized_text:
            if token in self.word_to_index:
                vector_text.append(self.word_to_index[token])
            else:
                vector_text.append(self.word_to_index["<UNK>"])
        
        vector_text.append(self.word_to_index["<EOS>"])
        
        pad_length = self.max_length - len(vector_text)
        for i in range(0, pad_length):
            vector_text.append(self.word_to_index["<PAD>"])
            
        return vector_text


class IMDBDataset(Dataset):
    def __init__(self):
        data = pd.read_csv("imdb.csv").to_numpy()
        
        self.target = []
        for data_point in data[:, 2]:
            if data_point == "neg":
                self.target.append(0)
            else:
                self.target.append(1)
                
        self.text = data[:, 4]

        self.vocab = Vocabulary()
        self.vocab.build_vocabulary(self.text)
    
    def __len__(self):
        return self.text.shape[0]
    
    def __getitem__(self, idx):
        review = self.text[idx]
        
        vector_text = [self.vocab.word_to_index["<SOS>"]]
        vector_text += self.vocab.numericalize(review)
        
        target = self.target[idx]
        
        return torch.tensor(vector_text), torch.tensor(target)


class Collate:
    def __init__(self, pad_idx):
        self.pad_idx = pad_idx
    
    def __call__(self, batch):
        text = [item[0] for item in batch]
        text = nn.utils.rnn.pad_sequence(text, batch_first=False, padding_value=self.pad_idx)
        
        return text, batch[1]

def get_loader(batch_size=32, num_workers=4, shuffle=True, pin_memory=True):
    dataset = IMDBDataset()
    pad_idx = dataset.vocab.word_to_index["<PAD>"]
    
    loader = DataLoader(
        dataset=dataset,
        batch_size=batch_size,
        num_workers=num_workers,
        shuffle=shuffle,
        pin_memory=pin_memory,
        collate_fn=Collate(pad_idx=pad_idx) # Redundant now
    )
    
    return loader, dataset

train_dl, train_ds = get_loader()

for idx, (data, target) in  enumerate(train_dl):
    print(data.shape)

【问题讨论】:

  • 您好,请提供一个 minimal 可重现的示例:stackoverflow.com/help/minimal-reproducible-example。在这里也复制粘贴错误消息。如果它适用于 google colab 但不适用于您的本地计算机,那么代码没有问题的可能性不大,而是您的本地设置有问题,因此您可能想在问题中描述它

标签: python text deep-learning nlp pytorch


【解决方案1】:

不知道为什么它会起作用,但通过删除 get_loader() 函数并自己获取数据加载器,解决了这个问题。

train_dl = DataLoader(dataset, batch_size=32, shuffle=True)

【讨论】:

    【解决方案2】:

    你可以试试

    if __name__ == '__main__' and '__file__' in globals():
    

    【讨论】:

      【解决方案3】:

      我猜你正在使用 Windows。当您设置 num_workers &gt; 0 时,Pytorch 的数据加载器会出现此错误。所以要修复这个错误,设置num_workers = 0 或调用if __name__ == "__main__: 下的数据加载器(我无法解释为什么最后一个有效)。

      【讨论】:

        猜你喜欢
        • 2013-11-10
        • 2010-12-01
        • 2020-01-13
        • 2015-04-14
        • 1970-01-01
        • 2014-02-13
        • 2014-06-11
        • 1970-01-01
        • 1970-01-01
        相关资源
        最近更新 更多