使用DistilBERT构建文本分类模型时遇RuntimeError问题求助
问题背景
使用BERT(dbmdz/bert-base-turkish-uncased)、RoBERTa(urakaytan/roberta-base-turkish-uncased)、DistilBERT(distilbert-base-uncased)构建文本分类模型时:
- BERT/RoBERTa调用
pooler_output时模型正常运行 - DistilBERT调用
last_hidden_state时触发错误:RuntimeError: Expected target size [32, 2], got [32]
错误原因
DistilBERT的last_hidden_state输出形状为[batch_size, seq_len, hidden_size](例如[32, 128, 768]),当前模型的全连接层直接作用于该张量后,最终输出形状为[batch_size, seq_len, 2],但训练时传入的标签是一维的[batch_size],两者形状不匹配,导致交叉熵损失计算失败。
而BERT/RoBERTa的pooler_output是CLS token经过处理后的结果,形状为[batch_size, hidden_size](例如[32,768]),经过全连接层后输出[batch_size,2],与标签形状匹配,因此可以正常运行。
解决方案
需要将DistilBERT的last_hidden_state转换为[batch_size, hidden_size]的张量,有两种常用方式:
方式1:提取CLS token输出
直接取last_hidden_state中第0个位置的token(即CLS token)的输出,这是BERT系列模型的标准分类用法:
class BERT_Arch(nn.Module): def __init__(self, bert): super(BERT_Arch, self).__init__() self.bert = bert self.dropout = nn.Dropout(0.1) self.relu = nn.ReLU() self.fc1 = nn.Linear(768,512) self.fc2 = nn.Linear(512,2) self.softmax = nn.LogSoftmax(dim=1) def forward(self, sent_id, mask): cls_hs = self.bert(sent_id, attention_mask=mask)["last_hidden_state"] # 提取CLS token的输出,形状变为[batch_size, 768] cls_hs = cls_hs[:, 0, :] x = self.fc1(cls_hs) x = self.relu(x) x = self.dropout(x) x = self.fc2(x) x = self.softmax(x) return x
方式2:序列均值池化(忽略padding)
对整个序列的隐藏状态做均值池化,同时忽略padding部分的影响,适合不需要依赖CLS token的场景:
class BERT_Arch(nn.Module): def __init__(self, bert): super(BERT_Arch, self).__init__() self.bert = bert self.dropout = nn.Dropout(0.1) self.relu = nn.ReLU() self.fc1 = nn.Linear(768,512) self.fc2 = nn.Linear(512,2) self.softmax = nn.LogSoftmax(dim=1) def forward(self, sent_id, mask): outputs = self.bert(sent_id, attention_mask=mask) last_hidden_state = outputs["last_hidden_state"] # 扩展mask维度,用于过滤padding部分 mask_expanded = mask.unsqueeze(-1).expand(last_hidden_state.size()) # 计算有效token的隐藏状态均值 sum_embeddings = torch.sum(last_hidden_state * mask_expanded, 1) sum_mask = torch.clamp(mask_expanded.sum(1), min=1e-9) # 避免除以0 cls_hs = sum_embeddings / sum_mask # 形状变为[batch_size, 768] x = self.fc1(cls_hs) x = self.relu(x) x = self.dropout(x) x = self.fc2(x) x = self.softmax(x) return x
验证修改
修改后,模型输出形状将变为[batch_size, 2],与标签[batch_size]的形状匹配,交叉熵损失可以正常计算,DistilBERT模型即可正常训练。
内容的提问来源于stack exchange,提问作者HappyDragneel

