训练BARTForSequenceClassification时数据维度不一致的问题解决
解决transformers 4.28.0微调BART-base二分类时评估阶段logits维度不匹配问题
核心问题定位
你遇到的报错本质是模型输出的logits结构不符合二分类任务的预期:
- 误用生成式BART模型(
BartForConditionalGeneration)做分类,导致模型返回包含解码器输出、隐藏状态等的Tuple,而非二维的分类logits; - 直接拼接Tuple元素时,因各元素维度差异(比如一个是
[样本数,2],另一个是[样本数,序列长度,词汇表大小]),引发维度不匹配错误。
分步解决方案
1. 替换为分类专用的BART模型
必须使用BartForSequenceClassification而非生成式模型,这是解决问题的关键。初始化代码示例:
from transformers import BartForSequenceClassification # 加载预训练BART-base并添加二分类头 model = BartForSequenceClassification.from_pretrained( "facebook/bart-base", num_labels=2, problem_type="single_label_classification" # 明确指定单标签二分类,避免歧义 )
这个模型的forward方法会直接返回SequenceClassifierOutput对象,其中的logits字段就是二维张量[batch_size, num_labels],完全符合分类任务的需求。
2. 修正compute_metrics函数的logits处理
如果确认用了分类模型,但logits仍为Tuple(极少数情况是自定义模型输出导致),可以在函数里先提取正确的分类logits:
from sklearn.metrics import f1_score import numpy as np def compute_metrics(eval_pred): logits, labels = eval_pred # 处理logits为Tuple的情况:筛选出维度为[样本数, 2]的分类logits if isinstance(logits, tuple): logits = [x for x in logits if x.ndim == 2][0] # 生成预测结果 predictions = np.argmax(logits, axis=-1) # 计算二分类F1分数 return {"f1": f1_score(labels, predictions, average="binary")}
3. 检查数据编码是否正确
确保数据处理时只保留分类任务所需的字段,不要混入生成任务的冗余字段:
def preprocess_function(examples): # 仅对输入文本编码,生成input_ids和attention_mask tokenized = tokenizer( examples["text"], truncation=True, padding="max_length", max_length=512 ) # 直接绑定labels,确保labels是一维数组(每个样本对应0/1) tokenized["labels"] = examples["label"] return tokenized # 映射到数据集 encoded_dataset = raw_dataset.map(preprocess_function, batched=True)
验证编码后的数据集,labels字段的形状应为[样本数],避免出现多维度的情况。
4. 验证Trainer配置
确保Trainer的eval_dataset是正确编码后的数据集,没有错误包含生成任务的字段(如decoder_input_ids)。transformers 4.28.0的Trainer会自动处理分类任务的输入输出,只要模型和数据正确,就能正常传递logits到compute_metrics。
内容的提问来源于stack exchange,提问作者Hardy Wen
相关产品推荐
相关产品推荐

