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

使用DistilBERT构建文本分类模型时遇RuntimeError问题求助

文本分类模型跨BERT系列架构适配问题解决

问题背景

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 18:45:00