基于自定义多标签数据集训练bert-base-uncased时遇到RuntimeError问题求助
嗨,我看了你遇到的这个RuntimeError问题,咱们一步步来排查解决:
首先,先看错误信息:RuntimeError: result type Float can't be cast to the desired output type Long,这个错误出现在binary_cross_entropy_with_logits函数里,说明损失计算时出现了类型不匹配的问题。结合你的代码,我发现几个需要调整的地方:
1. 模型配置需明确设置多标签分类类型
你继承了BertForSequenceClassification来实现多标签分类,但默认的模型配置里problem_type是单标签分类(single_label_classification),这会导致内部一些逻辑和多标签场景不兼容。在初始化模型的时候,需要显式指定problem_type为multi_label_classification:
model = BertForMultiLabelSequenceClassification.from_pretrained( model_checkpoint, num_labels=6, problem_type="multi_label_classification" # 新增这一行 )
2. 修正预测结果的处理逻辑
你的compute_metrics函数里直接对pred.predictions做round()是不对的,因为predictions是模型输出的logits(未经过激活函数的原始输出),应该先经过sigmoid激活,再根据阈值(比如0.5)来判断预测标签:
def compute_metrics(pred): labels = pred.label_ids # 先对logits做sigmoid激活,再用0.5作为阈值得到预测标签 preds = torch.sigmoid(torch.tensor(pred.predictions)).numpy() >= 0.5 precision, recall, f1, _ = precision_recall_fscore_support(labels, preds, average='weighted') acc = accuracy_score(labels, preds) return { 'accuracy': acc, 'f1': f1, 'precision': precision, 'recall': recall }
3. (可选)简化模型定义,减少自定义代码
其实你不需要自己继承BertForSequenceClassification来重写forward方法,Transformers库已经支持直接用AutoModelForSequenceClassification并指定problem_type来实现多标签分类,这样可以减少自定义代码带来的潜在问题:
# 替换你自定义的模型类,直接用下面的方式初始化 model = AutoModelForSequenceClassification.from_pretrained( model_checkpoint, num_labels=6, problem_type="multi_label_classification" )
这样模型会自动使用BCEWithLogitsLoss作为损失函数,代码更简洁也更稳定。
你可以先尝试这几个调整,应该能解决这个RuntimeError问题。
备注:内容来源于stack exchange,提问作者Emir Lise

