PyTorch多类别Token分类模型CrossEntropyLoss维度不匹配问题
解决Token分类模型中CrossEntropyLoss维度不匹配问题
问题背景
构建了基于DistilBERT+BiLSTM的PyTorch Token分类模型(类别数1182),训练时触发维度不匹配错误:
RuntimeError: Expected target size [8, 1182], got [8, 256]
问题根源
- 错误使用预训练模型类:使用
AutoModelForTokenClassification(自带分类头),但需求是自定义后续的BiLSTM和分类层,应改用基础的AutoModel。 - BiLSTM输入维度配置错误:DistilBERT输出的
last_hidden_state形状为[batch_size, seq_len, hidden_dim],但BiLSTM的input_size被错误设置为序列长度(256),而非模型隐藏维度(768)。 - 缺失分类层定义:模型中直接调用
self.classification_layer但未在初始化时创建该线性层,无法完成最后一步维度映射。 - 多余的softmax操作:
CrossEntropyLoss内部已集成log_softmax计算,提前对输出做softmax会导致损失逻辑错误。 - 代码笔误:报错信息显示实际代码中使用了未转置的
y_pred计算损失,而非转置后的transposed_y_pred。
修复步骤
1. 修正模型结构
import torch import torch.nn as nn import torch.nn.functional as F from transformers import AutoModel # 替换为基础预训练模型类 import utilities as utils from global_constants import MAX_DOC_LENGTH class CustomTorchModel(nn.Module): def __init__(self, args_model_name_or_path): super().__init__() # 必须调用父类初始化方法 id_to_label, label_to_id = utils.unshelve_label_converters() label_qty = len(list(label_to_id)) # 使用不带分类头的AutoModel self.distilbert_layer = AutoModel.from_pretrained( args_model_name_or_path, num_labels=label_qty ) # 修正BiLSTM的input_size为模型隐藏维度 self.bilstm_layer = nn.LSTM( input_size=self.distilbert_layer.config.dim, hidden_size=self.distilbert_layer.config.dim, num_layers=1, batch_first=True, bidirectional=True ) # 添加分类层:将BiLSTM的双向输出(2*hidden_dim)映射到类别数 self.classification_layer = nn.Linear( 2 * self.distilbert_layer.config.dim, label_qty ) def forward(self, inputs): input_ids, attention_mask = inputs[0], inputs[1] distilbert_output = self.distilbert_layer(input_ids=input_ids, attention_mask=attention_mask) last_hidden_state = distilbert_output.last_hidden_state # 形状:[8,256,768] bilstm_output, _ = self.bilstm_layer(last_hidden_state) # 形状:[8,256,1536] output = self.classification_layer(bilstm_output) # 形状:[8,256,1182] return output # 移除F.softmax,交给CrossEntropyLoss处理
2. 确保训练逻辑的维度正确
def _update(engine, batch): model.train() optimizer.zero_grad() x, y = _prepare_batch(batch, device=device) y_pred = model(x) # 输出形状:[8,256,1182] # 转置为CrossEntropyLoss要求的格式:[batch, num_classes, seq_len] transposed_y_pred = torch.transpose(y_pred, 1, 2) # 形状:[8,1182,256] # 目标标签y形状为[8,256],每个元素是类别索引(长整型) loss = loss_fn(transposed_y_pred, y.long()) loss.backward() optimizer.step() return loss.item(), transposed_y_pred, y.long()
3. 验证维度匹配
修复后各关键张量维度:
- DistilBERT输出:
torch.Size([8, 256, 768]) - BiLSTM输出:
torch.Size([8, 256, 1536]) - 分类层输出:
torch.Size([8, 256, 1182]) - 转置后预测张量:
torch.Size([8, 1182, 256]) - 目标标签:
torch.Size([8, 256])
此时CrossEntropyLoss会自动对每个序列位置的分类结果计算损失,维度完全匹配。
内容的提问来源于stack exchange,提问作者clanofsol
相关产品推荐
相关产品推荐

