为HuggingFace Transformer模型添加自定义层时触发IndexError:维度为2的张量索引过多
解决IndexError: too many indices for tensor of dimension 2的问题
这个问题我之前帮不少刚接触HuggingFace Transformers的新手踩过坑,核心原因是你选的预训练模型类用错了!
错误根源
你用AutoModelForSequenceClassification.from_pretrained(...)加载模型,但这个类本身已经自带了针对分类任务的输出层。它的输出outputs[0]是模型直接输出的分类logits,维度是[batch_size, num_labels](二维张量),而不是你以为的编码器最后一层的隐藏状态(三维张量:[batch_size, sequence_length, hidden_size])。
所以你试图用sequence_output[:,0,:]去索引一个二维张量,自然会触发“维度2的张量有太多索引”的错误。
修复方案
我们需要加载不带分类头的基础模型(AutoModel),这样才能拿到编码器的隐藏状态来接自定义分类层。修改步骤如下:
1. 替换模型加载类
把初始化里的AutoModelForSequenceClassification改成AutoModel,因为我们需要的是模型的主体编码器,而不是带分类头的完整分类模型。
2. 确认输出提取逻辑
AutoModel的outputs[0]才是编码器最后一层的隐藏状态,维度为[batch_size, seq_len, 768],这时候取第0个token(也就是<s>/cls token)的特征就完全没问题了。
修改后的完整代码
import torch import torch.nn as nn from transformers import AutoModel, AutoConfig, TokenClassifierOutput class CustomModel(nn.Module): def __init__(self, checkpoint, num_labels): super(CustomModel, self).__init__() self.num_labels = num_labels # 加载不带分类头的基础模型 self.model = AutoModel.from_pretrained( checkpoint, config=AutoConfig.from_pretrained( checkpoint, output_attentions=True, output_hidden_states=True ) ) self.dropout = nn.Dropout(0.1) self.classifier = nn.Linear(768, num_labels) def forward(self, input_ids=None, attention_mask=None, labels=None): # 提取编码器的输出,outputs[0]是最后一层隐藏状态(三维张量) outputs = self.model(input_ids=input_ids, attention_mask=attention_mask) # 对<cls> token的特征做dropout sequence_output = self.dropout(outputs[0][:, 0, :]) # 计算logits logits = self.classifier(sequence_output) loss = None if labels is not None: loss_fct = nn.CrossEntropyLoss() loss = loss_fct(logits.view(-1, self.num_labels), labels.view(-1)) return TokenClassifierOutput( loss=loss, logits=logits, hidden_states=outputs.hidden_states, attentions=outputs.attentions )
额外注意事项
- 如果你坚持要用
AutoModelForSequenceClassification,那你需要去掉它自带的分类头(比如self.model.classifier = nn.Identity()),然后取它的outputs.hidden_states[-1]作为最后一层隐藏状态,但这种方式不如直接用AutoModel简洁。 - 确认你的预训练模型的隐藏层维度是768(比如bert-base-uncased),如果用的是更大的模型(比如bert-large-uncased),需要把
nn.Linear(768, num_labels)改成nn.Linear(1024, num_labels)。
内容的提问来源于stack exchange,提问作者ntdev
相关产品推荐
相关产品推荐

