PyTorch微调DistilBert二分类遇batch_size不匹配错误求助
问题分析与解决方案
错误Expected input batch_size (1) to match target batch_size (64)的核心原因是:模型输出的预测结果batch_size为1,而标签的batch_size为64,两者维度不匹配。问题出在模型输入的处理逻辑上。
问题根源
你的train_loader_pt返回的X应该是包含input_ids和attention_mask的元组(或字典),每个张量的形状为[64, seq_len](对应batch_size=64)。但当前模型的forward函数仅取了x[0]作为输入传入DistilBert,且没有正确解包输入参数,导致模型只处理了batch中的第一个样本,输出维度变为[1,2],与标签维度[64]冲突。
具体修改步骤
1. 修正模型的Forward函数
调整模型输入参数的处理逻辑,正确接收并传入DistilBert所需的input_ids和attention_mask:
from torch import nn device = "cuda" if torch.cuda.is_available() else "cpu" print(f"Using {device} device") class DistilBertClassification(nn.Module): def __init__(self): super(DistilBertClassification, self).__init__() self.dbert = dbert_pt self.dropout = nn.Dropout(p=0.1) self.linear1 = nn.Linear(768,64) self.ReLu = nn.ReLU() self.linear2 = nn.Linear(64,2) def forward(self, input_ids, attention_mask=None): # 传入input_ids和可选的attention_mask到DistilBert outputs = self.dbert(input_ids=input_ids, attention_mask=attention_mask) # 取<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>token的隐藏层输出 x = outputs["last_hidden_state"][:, 0, :] x = self.dropout(x) x = self.linear1(x) x = self.ReLu(x) logits = self.linear2(x) return logits model_pt = DistilBertClassification().to(device)
2. 修正训练循环中的模型调用
在训练时解包X的参数,传入模型:
from tqdm import tqdm for e in range(epochs): model_pt.train() train_loss = 0.0 train_accuracy = [] for X, y in tqdm(train_loader_pt): # 解包X中的input_ids和attention_mask,传入模型 prediction = model_pt(*X.cuda()) loss = criterion(prediction , y.cuda()) optimizer.zero_grad() loss.backward() optimizer.step() train_loss += loss.item() prediction_index = prediction.argmax(axis=1) accuracy = (prediction_index==y.cuda()) train_accuracy += accuracy train_accuracy = (sum(train_accuracy) / len(train_accuracy)).item()
3. 验证数据加载器输出(可选)
确保train_loader_pt返回的X中,input_ids的形状为[batch_size, seq_len],而非单个样本的列表。如果使用自定义collate_fn,需保证它将batch样本堆叠为二维张量:
def collate_fn(batch): input_ids = torch.stack([item[0] for item in batch]) attention_mask = torch.stack([item[1] for item in batch]) labels = torch.tensor([item[2] for item in batch]) return (input_ids, attention_mask), labels
验证效果
修改后,打印prediction.shape应该为torch.Size([64, 2]),与y.shape(torch.Size([64]))匹配,损失计算即可正常运行。
内容的提问来源于stack exchange,提问作者Akram H
相关产品推荐
相关产品推荐

