基于LSTM的RNN训练时DataLoader迭代出现IndexError索引越界求助
解决PyTorch LSTM训练中DataLoader迭代的IndexError问题
问题背景
使用PyTorch训练基于LSTM的RNN时,在DataLoader迭代过程中触发「IndexError: index out of range」错误,已排查CSV文件、max_seq_length参数、数据预处理(填充/格式化)等环节,仍未解决。相关代码如下:
import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader, Dataset from torch.nn.utils.rnn import pad_sequence import pandas as pd import json class RNNModel(nn.Module): def __init__(self, vocab_size, embedding_dim, hidden_dim): super(RNNModel, self).__init__() self.embedding = nn.Embedding(vocab_size, embedding_dim) self.rnn = nn.LSTM(embedding_dim, hidden_dim, batch_first=True) self.fc = nn.Linear(hidden_dim, vocab_size) def forward(self, x): embedded = self.embedding(x) output, (h_n, c_n) = self.rnn(embedded) return self.fc(output) # Load tokenized vocabulary with open('cleaned_vocab.json', 'r', encoding='utf-8') as vocab_file: vocab = json.load(vocab_file) # Load processed data from CSV class CustomDataset(Dataset): def __init__(self, csv_path, max_seq_length): self.data = pd.read_csv(csv_path) self.max_seq_length = max_seq_length def __len__(self): return len(self.data) def __getitem__(self, idx): text = self.data.loc[idx, 'text'] tokens = [int(token) for token in text.split()] if len(tokens) > self.max_seq_length: tokens = tokens[:self.max_seq_length] padded_sequence = tokens + [0] * (self.max_seq_length - len(tokens)) input_sequence = torch.tensor(padded_sequence[:-1]) # Input sequence without last token target_sequence = torch.tensor(padded_sequence[1:]) # Target sequence without first token return input_sequence, target_sequence # Custom collate function class CustomCollate: def __init__(self, pad_idx): self.pad_idx = pad_idx def __call__(self, batch): input_seqs, target_seqs = zip(*batch) padded_input_seqs = pad_sequence(input_seqs, batch_first=True, padding_value=self.pad_idx) padded_target_seqs = pad_sequence(target_seqs, batch_first=True, padding_value=self.pad_idx) return padded_input_seqs, padded_target_seqs # Initialize custom dataset max_sequence_length = 30 # Define your desired maximum sequence length dataset = CustomDataset('processed_data.csv', max_sequence_length) # Create a dataloader with custom collate function dataloader = DataLoader(dataset, batch_size=32, shuffle=True, collate_fn=CustomCollate(0)) # Initialize the RNN model vocab_size = len(vocab) embedding_dim = 128 hidden_dim = 256 rnn_model = RNNModel(vocab_size, embedding_dim, hidden_dim) # Define loss function and optimizer criterion = nn.CrossEntropyLoss() optimizer = optim.Adam(rnn_model.parameters(), lr=0.001) # Training loop num_epochs = 10 for epoch in range(num_epochs): for input_batch, target_batch in dataloader: optimizer.zero_grad() # Forward pass output = rnn_model(input_batch) # Calculate loss and backpropagate loss = criterion(output.transpose(1, 2), target_batch) loss.backward() optimizer.step() # Save the trained model torch.save(rnn_model.state_dict(), 'rnn_model.pth') print("Training completed.")
错误排查与修复建议
1. 校验token索引是否超出vocab范围
nn.Embedding要求输入的token索引必须满足 0 <= token < vocab_size,如果CSV中存在大于等于vocab_size的token,会直接触发IndexError。在CustomDataset的__getitem__方法中添加校验:
def __getitem__(self, idx): text = self.data.loc[idx, 'text'] tokens = [int(token) for token in text.split()] # 新增:校验token范围 vocab_size = len(vocab) for token in tokens: if token < 0 or token >= vocab_size: raise ValueError(f"Invalid token {token} at sample {idx}, vocab size is {vocab_size}") # 后续原有逻辑...
2. 移除冗余的自定义Collate函数
你的Dataset已经将所有输入/目标序列处理为固定长度(max_sequence_length-1),CustomCollate中的pad_sequence属于多余操作,直接使用默认collate函数即可:
dataloader = DataLoader(dataset, batch_size=32, shuffle=True)
3. 过滤空序列或无效数据
如果CSV中存在空的text字段,处理后会生成无意义的全0序列,可能引发隐性错误。在Dataset初始化时提前过滤:
class CustomDataset(Dataset): def __init__(self, csv_path, max_seq_length): self.data = pd.read_csv(csv_path) # 过滤空文本行 self.data = self.data[self.data['text'].str.strip() != ''].reset_index(drop=True) self.max_seq_length = max_seq_length
4. 捕获错误并定位具体数据
在训练循环中添加异常捕获,打印出错时的batch信息,快速定位问题样本:
for epoch in range(num_epochs): for batch_idx, (input_batch, target_batch) in enumerate(dataloader): try: optimizer.zero_grad() output = rnn_model(input_batch) loss = criterion(output.transpose(1, 2), target_batch) loss.backward() optimizer.step() except IndexError as e: print(f"Error occurred at batch {batch_idx}:") print(f"Input batch shape: {input_batch.shape}") print(f"Target batch shape: {target_batch.shape}") print(f"First problematic input sample: {input_batch[0]}") raise e
5. 确认vocab索引的连续性
检查cleaned_vocab.json,确保vocab的索引是从0开始的连续整数。如果存在索引断层(比如跳过了某个数字),而CSV中恰好有对应断层的token值,也会触发索引越界错误。重新生成vocab时确保索引连续递增。
内容的提问来源于stack exchange,提问作者capdescx
相关产品推荐
相关产品推荐

