GPT2微调报错:输入batch_size与目标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)当成另一个值,触发不匹配报错。
具体修复步骤
统一输入与标签的序列长度
修改分词函数,将标签截断/填充至与输入相同的序列长度,并把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避免不完整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

