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

GPT2微调报错:输入batch_size与目标batch_size不匹配

解决GPT2微调时输入与目标batch_size不匹配的错误

问题概述

用GPT2微调简单问答数据集时触发错误:Expected input batch_size (28) to match target batch_size (456), Changing batch size increase the target batch size with GPT2 model,检查数据集形状看似正常,但报错依旧。

现有代码与输出

数据分词函数

def tokenize_data(total_marks, coding_feeddback):
    inputs = tokenizer(total_marks, truncation=True, padding=True, 
                                                     return_tensors="pt")
    labels = tokenizer(coding_feeddback, truncation=True, padding=True, 
                                                    return_tensors="pt")['input_ids']
    return inputs, labels

数据集准备

# 准备训练集与验证集
train_inputs, train_labels = tokenize_data(train_df['Question'].tolist(), 
 train_df['ans'].tolist())
val_inputs, val_labels = tokenize_data(val_df['Question'].tolist(), 
 val_df['ans'].tolist())

train_dataset = TensorDataset(train_inputs['input_ids'], train_labels)
val_dataset = TensorDataset(val_inputs['input_ids'], val_labels)

数据集尺寸验证代码

print('train input shape:',train_inputs['input_ids'].shape)
print('train label shape: ',train_labels.shape)
print('validation input shape: ',val_inputs['input_ids'].shape)
print('validation label shape: ',val_labels.shape)

输出结果

train input shape: torch.Size([76, 8])
train label shape:  torch.Size([76, 115])
validation input shape:  torch.Size([20, 8])
validation label shape:  torch.Size([20, 98])

DataLoader定义

batch_size = 4
train_dataloader = DataLoader(train_dataset, 
                   batch_size=batch_size, shuffle=True)
val_dataloader = DataLoader(val_dataset, batch_size=batch_size)

训练与验证循环

# 训练循环
model.train()
for epoch in range(num_epochs):
 for batch in train_dataloader:
    batch = [item.to(device) for item in batch]
    input_ids, labels = batch

    optimizer.zero_grad()
    
    print("indputIds:",len(input_ids))
    print("lebels:",len(labels))

    outputs = model(input_ids=input_ids, labels=labels)
    loss = outputs.loss
    logits = outputs.logits

    loss.backward()
    optimizer.step()

# 验证阶段
with torch.no_grad():
    model.eval()
    val_loss = 0.0
    for val_batch in val_dataloader:
        val_batch = [item.to(device) for item in val_batch]
        val_input_ids, val_labels = val_batch

        val_outputs = model(input_ids=val_input_ids, labels=val_labels)
        val_loss += val_outputs.loss.item()

    average_val_loss = val_loss / len(val_dataloader)
    print(f"Epoch: {epoch+1}, Validation Loss: {average_val_loss:.4f}")

model.train()

Batch维度示例

Batch 19
Inputs:
torch.Size([4, 8])
Targets:
torch.Size([4, 115])

验证集类似,目标尺寸为[4,98]

解决方案

核心问题

GPT2计算损失时要求输入序列与标签序列维度严格对齐(即[batch_size, seq_len]完全一致)。当前标签序列长度远大于输入,模型内部误将标签的总元素数(如4115=460)当成batch_size,输入的总元素数(如48=32)当成另一个值,触发不匹配报错。

具体修复步骤

  1. 统一输入与标签的序列长度
    修改分词函数,将标签截断/填充至与输入相同的序列长度,并把padding位置设为-100(PyTorch交叉熵损失会自动忽略该值):

    def tokenize_data(total_marks, coding_feeddback):
        inputs = tokenizer(total_marks, truncation=True, padding=True, return_tensors="pt")
        # 获取输入的序列长度
        input_seq_len = inputs['input_ids'].shape[1]
        # 处理标签,对齐输入长度
        labels_tokenized = tokenizer(
            coding_feeddback, 
            truncation=True, 
            padding='max_length', 
            max_length=input_seq_len, 
            return_tensors="pt"
        )
        labels = labels_tokenized['input_ids']
        # 标记padding位置,避免计算损失
        labels[labels == tokenizer.pad_token_id] = -100
        return inputs, labels
    
  2. 避免不完整Batch干扰
    在DataLoader中添加drop_last=True,丢弃最后一个样本数不足的Batch:

    train_dataloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, drop_last=True)
    val_dataloader = DataLoader(val_dataset, batch_size=batch_size, drop_last=True)
    

原理说明

GPT2的损失基于逐token交叉熵计算,要求输入和标签的维度完全匹配。当标签长度与输入不一致时,模型无法正确解析维度,会错误地将二维张量的总元素数当成batch_size,从而触发报错。通过对齐两者长度并标记无效padding位置,即可让模型正确计算损失。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.19 23:45:23