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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 00:05:54