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

PyTorch训练验证数据拆分报错:长度总和不匹配问题求助

问题:PyTorch训练集与验证集拆分触发ValueError错误

完整代码

import torch
import torch.nn as nn
import torch.optim as optim
import torch.nn.functional as F

class ChatBot(nn.Module):
    def __init__(self, input_size, hidden_size, num_layers, output_size):
        super().__init__()
        self.hidden_size = hidden_size
        self.num_layers = num_layers
        self.lstm = nn.LSTM(input_size, hidden_size, num_layers, batch_first=True)
        self.fc = nn.Linear(hidden_size, output_size)
    
    def forward(self, x, hidden):
        out, hidden = self.lstm(x, hidden)
        out = self.fc(out[:, -1, :])
        return out, hidden


    def init_hidden(self, batch_size):
        weight = next(self.parameters()).data
        hidden = (weight.new(self.num_layers, batch_size, self.hidden_size).zero_(),
              weight.new(self.num_layers, batch_size, self.hidden_size).zero_())
        return hidden

class ChatDataset(torch.utils.data.Dataset):
    def __init__(self, data):
        self.data = data
    
    def __len__(self):
        return len(self.data)


    def __getitem__(self, index):
        return self.data[index]


def train(model, train_loader, loss_fn, optimizer, device):
    model.train()
    for inputs, targets in train_loader:
        inputs = inputs.to(device)
        targets = targets.to(device)
    
        hidden = model.init_hidden(inputs.size(0))
        hidden = tuple([each.data for each in hidden])
    
        optimizer.zero_grad()
        outputs, _ = model(inputs, hidden)
        loss = loss_fn(outputs.view(-1), targets.view(-1))
        loss.backward()
        optimizer.step()
    
def evaluate(model, val_loader, loss_fn, device):
    model.eval()
    total_loss = 0
    with torch.no_grad():
        for inputs, targets in val_loader:
            inputs = inputs.to(device)
            targets = targets.to(device)
        
            hidden = model.init_hidden(inputs.size(0))
            hidden = tuple([each.data for each in hidden])
        
            outputs, _ = model(inputs, hidden)
            total_loss += loss_fn(outputs, targets).item()
    return total_loss / len(val_loader)

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

input_size = 500
hidden_size = 128
num_layers = 2
output_size = 500

model = ChatBot(input_size, hidden_size, num_layers, output_size)
model = model.to(device)

data = [("Hi, how are you?", "I'm doing well, thank you for asking."),
("What's your name?", "I'm a chatbot, I don't have a name."),
("What's the weather like?", "I'm not sure, I don't have access to current weather information."),
("What's the time?", "I'm not sure, I don't have access to the current time.")]

dataset = ChatDataset(data)

train_dataset, val_dataset = torch.utils.data.random_split(dataset, [int(0.8 * len(dataset)), int(0.2 * len(dataset))])

train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=32, shuffle=True)
val_loader = torch.utils.data.DataLoader(val_dataset, batch_size=32, shuffle=False)

loss_fn = nn.MSELoss()
optimizer =  optim.Adam(model.parameters(), lr=0.001)

num_epochs = 100

for epoch in range(num_epochs):
   train(model, train_loader, loss_fn, optimizer, device)
   val_loss = evaluate(model, val_loader, loss_fn, device)
   print("Epoch [{}/{}], Validation Loss: {:.4f}".format(epoch+1, num_epochs, val_loss))

torch.save(model.state_dict(), 'chatbot_model.pt')

运行错误

ValueError
Traceback (most recent call last)
<ipython-input-8-ae2a6dd1bc7c> in <module>
 78 dataset = ChatDataset(data)
 79 
---> 80 train_dataset, val_dataset = torch.utils.data.random_split(dataset, [int(0.8 * len(dataset)), int(0.2 * len(dataset))])
 81 
 82 train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=32, shuffle=True)
/usr/local/lib/python3.8/dist-packages/torch/utils/data/dataset.py in random_split(dataset, lengths, generator)
345     # Cannot verify that dataset is Sized
346     if sum(lengths) != len(dataset):    # type: ignore[arg-type]
--> 347         raise ValueError("Sum of input lengths does not equal the length of the input dataset!")
348 
349     indices = randperm(sum(lengths), generator=generator).tolist()  # type: ignore[call-overload]
ValueError: Sum of input lengths does not equal the length of the input dataset!

错误原因

你的数据集长度是4,计算拆分长度时:

  • int(0.8 * 4) 得到3(int()会直接截断小数部分)
  • int(0.2 * 4) 得到0
    两者总和是3,不等于原数据集长度4,触发了random_split的长度校验逻辑。

解决方案

有三种可行的修改方式:

方式1:手动指定拆分长度

直接根据数据集实际长度分配,比如训练集3条,验证集1条:

train_dataset, val_dataset = torch.utils.data.random_split(dataset, [3, 1])

方式2:使用浮点数比例(PyTorch 1.10+支持)

无需转成整数,直接传入比例列表,random_split会自动计算正确的长度:

train_dataset, val_dataset = torch.utils.data.random_split(dataset, [0.8, 0.2])

方式3:确保总和匹配原数据集长度

先计算训练集长度,验证集长度用原长度减去训练集长度,避免截断误差:

train_len = int(0.8 * len(dataset))
val_len = len(dataset) - train_len
train_dataset, val_dataset = torch.utils.data.random_split(dataset, [train_len, val_len])

额外注意点

你的数据集只有4条样本,后续训练时batch_size=32会导致每个batch直接包含全部样本,建议根据数据集大小调整batch_size,比如设为1或2。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 17:05:22