如何修改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
相关产品推荐
相关产品推荐

