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

基于自定义多标签数据集训练bert-base-uncased时遇到RuntimeError问题求助

基于自定义多标签数据集训练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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.22 15:14:48