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

如何修改DistilBERT多任务模型的forward返回逻辑适配NER与跨度分类?

先修正代码中的明显错误

你当前代码里的logits_args= None写在了if labels is not None:分支外,这会导致无论是否传入标签,logits_args都会被置空,直接删掉这行代码。


一、处理TokenClassifierOutput的多任务logits问题

Hugging Face默认的TokenClassifierOutput仅支持单个logits字段,无法直接承载两个任务的输出。最清晰的方式是自定义多任务输出类,继承官方基础输出类来扩展字段:

from transformers.modeling_outputs import BaseModelOutputWithPastAndCrossAttentions
from typing import Optional, Tuple
import torch

class MultiTaskTokenClassifierOutput(BaseModelOutputWithPastAndCrossAttentions):
    loss: Optional[torch.FloatTensor] = None
    ner_logits: torch.FloatTensor = None  # NER任务的输出
    arg_logits: torch.FloatTensor = None  # 论元分类任务的输出
    hidden_states: Optional[Tuple[torch.FloatTensor]] = None
    attentions: Optional[Tuple[torch.FloatTensor]] = None

之后在return_dict=True的分支里,用这个自定义类返回结果:

return MultiTaskTokenClassifierOutput(
    loss=loss,
    ner_logits=logits,
    arg_logits=logits_args,
    hidden_states=outputs.hidden_states,
    attentions=outputs.attentions,
)

如果不想自定义类,也可以把两个logits打包成元组/字典塞进原logits字段,但自定义类的可读性和后续调用便利性更强。


二、修改if not return_dict:分支的逻辑

非return_dict模式下,模型会返回有序元组,只需把两个任务的logits按固定顺序加入元组即可:

if not return_dict:
    # 输出顺序:loss(若存在)、NER logits、论元分类logits、DistilBERT的其他输出(隐藏状态/注意力等)
    output = (logits, logits_args) + outputs[1:]
    return ((loss,) + output) if loss is not None else output

调用时只需按顺序解析元组就能拿到两个任务的结果。


额外提示

你当前的loss计算逻辑是用同一个labels同时计算两个任务的损失,这大概率不符合实际需求——NER和论元分类的标签应该是独立的。建议在forward函数参数中新增arg_labels,分开计算两个任务的损失:

# 修改forward参数,新增arg_labels
def forward(
    self,
    input_ids=None,
    attention_mask=None,
    head_mask=None,
    inputs_embeds=None,
    labels=None,  # NER的标签
    arg_labels=None,  # 论元分类的标签
    output_attentions=None,
    output_hidden_states=None,
    return_dict=None,
):
    # ... 其他代码不变 ...
    if labels is not None and arg_labels is not None:
        loss_fct = nn.BCEWithLogitsLoss()
        loss_spans = loss_fct(logits, labels.float())
        loss_args = loss_fct(logits_args, arg_labels.float())
        loss = loss_spans + loss_args

内容的提问来源于stack exchange,提问作者Glinty

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 14:37:51