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

同一数据集下多标签分类模型失效,多分类模型表现优异求助

图卷积神经网络多标签分类训练异常问题

我正在训练图卷积神经网络(Graph convolutional neural network),网络末端连接分类器:

  • 当以单分类(多分类)模式训练时,模型可取得91%的加权F1分数,效果优异:
    单(多分类)模型效果
  • 但将同一模型改为多标签分类训练时,训练时长大幅增加且模型失效:
    多标签分类模型失效情况

以下是适配单/多分类的分类器代码:

class Classifier(nn.Module):
def __init__(self, input_dim, hidden_size, tag_size, args, pred_type='SINGLE'):
    super(Classifier, self).__init__()
    self.emotion_att = MaskedEmotionAtt(input_dim)
    self.lin1 = nn.Linear(input_dim, hidden_size)
    self.drop = nn.Dropout(args.drop_rate)
    self.lin2 = nn.Linear(hidden_size, tag_size)
    self.pred_type = pred_type
    if args.class_weight:
        self.loss_weights = (torch.rand(tag_size) * 9 + 1).to(args.device)
        if self.pred_type == 'SINGLE':
            self.loss_func = nn.NLLLoss(self.loss_weights)
        elif self.pred_type == 'MULTI':
            self.loss_func = nn.BCELoss(self.loss_weights)
    else:
        if self.pred_type == 'SINGLE':
            self.loss_func = nn.NLLLoss()
        elif self.pred_type == 'MULTI':
            self.loss_func = nn.BCELoss()

def get_prob(self, h, text_len_tensor):
    # h_hat = self.emotion_att(h, text_len_tensor)
    hidden = self.drop(F.relu(self.lin1(h)))
    scores = self.lin2(hidden)
    if self.pred_type == 'SINGLE':
        log_prob = F.log_softmax(scores, dim=-1)
    elif self.pred_type == 'MULTI':
        log_prob = F.sigmoid(scores)
    else:
        return scores
    return log_prob

def forward(self, h, text_len_tensor):
    log_prob = self.get_prob(h, text_len_tensor)
    if self.pred_type == 'SINGLE':
        y_hat = torch.argmax(log_prob, dim=-1)
    elif self.pred_type == 'MULTI':
        y_hat = []
        for pred in log_prob:
            y_hat.append([1 if p > 0.5 else 0 for p in pred.flatten().tolist()])

        y_hat = torch.tensor(y_hat)
    try:
        y_hat
    except NameError:
        log.error('Prediction type should be one of these [\'SINGLE\', \'MULTI\']')
    return y_hat


def get_loss(self, h, label_tensor, text_len_tensor):
    log_prob = self.get_prob(h, text_len_tensor)
    if self.pred_type == 'SINGLE':
        loss = self.loss_func(log_prob, label_tensor)
    elif self.pred_type == 'MULTI':
        loss = self.loss_func(log_prob, label_tensor.float())
    return loss

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 09:54:53