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

2层RNN训练出现batch size维度不匹配ValueError问题咨询

问题背景

搭建输出为11维的多分类RNN模型,输入采用预训练GloVe模型生成的词向量,训练过程出现批次大小维度不匹配报错:

  • 当设置batch_size=1时,报错信息为ValueError: Expected input batch_size (1) to match target batch_size (11).
  • 当调整batch_size=11时,报错变为ValueError: Expected input batch_size (11) to match target batch_size (121).

推测错误来源是text张量的形状为torch.Size([11, 300]),缺少序列长度维度,原以为未指定序列长度时会默认取值1,但不知道如何正确添加该维度。

训练循环代码

def train(model, device, train_loader, valid_loader, epochs, learning_rate):

  criterion = nn.CrossEntropyLoss()
  optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate)
  
  train_loss, validation_loss = [], []
  train_acc, validation_acc = [], []

  for epoch in range(epochs):
    #train
    model.train()
    running_loss = 0.
    correct, total = 0, 0
    steps = 0
    for idx, batch in enumerate(train_loader):
      text = batch["Sample"].to(device)
      target = batch['Class'].to(device)
      print(text.shape, target.shape)
      text, target = text.to(device), target.to(device)
      # add micro for coding training loop
      optimizer.zero_grad()
      print(text.shape)
      output, hidden = model(text.unsqueeze(1))
      #print(output.shape, target.shape, target.view(-1).shape)
      loss = criterion(output, target.view(-1))
      loss.backward()
      optimizer.step()
      steps += 1
      running_loss += loss.item()

      # get accuracy
      _, predicted = torch.max(output, 1)
      print(predicted)
      #predicted = torch.round(output.squeeze())
      total += target.size(0)
      correct += (predicted == target).sum().item()

    train_loss.append(running_loss/len(train_loader))
    train_acc.append(correct/total)

    print(f'Epoch: {epoch + 1}, 'f'Training Loss: {running_loss/len(train_loader):.4f}, 'f'Training Accuracy: {100*correct/total: .2f}%')

    # evaluate on validation data
    model.eval()
    running_loss = 0.
    correct, total = 0, 0

    with torch.no_grad():
      for idx, batch in enumerate(valid_loader):
        text = batch["Sample"].to(device)
        print(type(text), text.shape)
        target = batch['Class'].to(device)
        target = torch.autograd.Variable(target).long()
        text, target = text.to(device), target.to(device)

        optimizer.zero_grad()
        output = model(text)
        
        loss = criterion(output, target)
        running_loss += loss.item()

        # get accuracy
        _, predicted = torch.max(output, 1)
        #predicted = torch.round(output.squeeze())
        total += target.size(0)
        correct += (predicted == target).sum().item()

    validation_loss.append(running_loss/len(valid_loader))
    validation_acc.append(correct/total)

    print (f'Validation Loss: {running_loss/len(valid_loader):.4f}, 'f'Validation Accuracy: {100*correct/total: .2f}%')

  return train_loss, train_acc, validation_loss, validation_acc

训练启动代码

# Model hyperparamters
#vocab_size = len(word_array)
learning_rate = 1e-3
hidden_dim = 100
output_size = 11
input_size = 300
epochs = 10
n_layers = 2

# Initialize model, training and testing
set_seed(SEED)
vanilla_rnn_model = VanillaRNN(input_size, output_size, hidden_dim, n_layers)
vanilla_rnn_model.to(DEVICE)
vanilla_rnn_start_time = time.time()
vanilla_train_loss, vanilla_train_acc, vanilla_validation_loss, vanilla_validation_acc = train(vanilla_rnn_model,DEVICE,train_loader,valid_loader,epochs = epochs,learning_rate = learning_rate)

数据加载器构造代码

# Splitting dataset
# define a batch_size, I'll use 4 as an example
batch_size = 1

train_dset = CustomDataset(X2, y)  # create data set
train_loader = DataLoader(train_dset, batch_size=batch_size, shuffle=True) #load data with batch size
valid_dset = CustomDataset(X2, y)
valid_loader = DataLoader(valid_dset, batch_size=batch_size, shuffle=True)

g_seed = torch.Generator()
g_seed.manual_seed(SEED)

完整报错回溯

---------------------------------------------------------------------------
ValueError                                Traceback (most recent call last)
<ipython-input-23-bfd2f8f3456f> in <module>()
     19                                                                                                valid_loader,
     20                                                                                                epochs = epochs,
--- 21                                                                                                learning_rate = learning_rate)
     22 print("--- Time taken to train = %s seconds ---" % (time.time() - vanilla_rnn_start_time))
     23 #test_accuracy = test(vanilla_rnn_model, DEVICE, test_iter)

3 frames
<ipython-input-22-16748701034f> in train(model, device, train_loader, valid_loader, epochs, learning_rate)
     47       output, hidden = model(text.unsqueeze(1))
     48       #print(output.shape, target.shape, target.view(-1).shape)
--- 49       loss = criterion(output, target.view(-1))
     50       loss.backward()
     51       optimizer.step()

/usr/local/lib/python3.7/dist-packages/torch/nn/modules/module.py in _call_impl(self, *input, **kwargs)
   1049         if not (self._backward_hooks or self._forward_hooks or self._forward_pre_hooks or _global_backward_hooks
   1050                 or _global_forward_hooks or _global_forward_pre_hooks):
-> 1051             return forward_call(*input, **kwargs)
   1052         # Do not call functions when jit is used
   1053         full_backward_hooks, non_full_backward_hooks = [], []

/usr/local/lib/python3.7/dist-packages/torch/nn/modules/loss.py in forward(self, input: Tensor, target: Tensor) -> Tensor:
   1119     def forward(self, input: Tensor, target: Tensor) -> Tensor:
   1120         return F.cross_entropy(input, target, weight=self.weight,
-> 1121                                ignore_index=self.ignore_index, reduction=self.reduction)
   1122 
   1123 

/usr/local/lib/python3.7/dist-packages/torch/nn/functional.py in cross_entropy(input, target, weight, size_average, ignore_index, reduce, reduction)
   2822     if size_average is not None or reduce is not None:
   2823         reduction = _Reduction.legacy_get_string(size_average, reduce)
-> 2824     return torch._C._nn.cross_entropy_loss(input, target, weight, _Reduction.get_enum(reduction), ignore_index)
   2825 
   2826 

ValueError: Expected input batch_size (1) to match target batch_size (11).
解决方案

问题根源

  • 维度处理错误:你单样本的text形状为[11, 300],经过DataLoader组装后,batch_size=1时text的形状已经是[1, 11, 300](对应[batch_size, seq_len, feature_dim]的RNN标准输入格式),你额外调用text.unsqueeze(1)会新增多余维度,导致RNN输出的batch维度异常。
  • 标签维度错误:你的CustomDataset返回的Class标签形状为[11]而非单个值,batch_size=1时target形状为[1, 11],调用target.view(-1)后会被展平为长度11的张量,和输出的batch维度1不匹配,这也是batch_size=11时target维度变为121的直接原因。

修复步骤

  1. 先修改CustomDataset的__getitem__方法,确保每个样本返回的标签是单个整数,而非长度为11的序列。
  2. 修正训练阶段的输入和标签处理,删除多余的维度操作:
# 训练循环训练部分修改
for idx, batch in enumerate(train_loader):
    text = batch["Sample"].to(device)
    target = batch['Class'].to(device)
    optimizer.zero_grad()
    # 删除unsqueeze操作,text本身已经符合RNN输入要求
    output, hidden = model(text)
    # 删除target.view(-1)操作,target形状应为[batch_size]
    loss = criterion(output, target)
    # 后续逻辑保持不变
  1. 修正验证阶段的输入处理,和训练阶段保持一致:
# 训练循环验证部分修改
for idx, batch in enumerate(valid_loader):
    text = batch["Sample"].to(device)
    target = batch['Class'].to(device).long()
    # 保持和训练阶段相同的输入处理,补充hidden接收
    output, hidden = model(text)
    loss = criterion(output, target)
    # 后续逻辑保持不变

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 23:57:02