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

基于BERT的MBTI多标签分类:是否需自定义BCEWithLogitsLoss?

MBTI多标签分类中仅关注正样本的损失函数优化方案

我正在开展基于MBTI分类体系的多标签情感分析任务,具体流程为:划分数据后将批次数据输入BERT层生成嵌入向量;构建神经网络,基于BERT输出的嵌入向量进行分类,最终输出16维结果;将标签转换为0-15的索引后,生成one-hot编码的16维标签矩阵,当前使用BCEWithLogitsLoss计算损失。但我仅关注标签中的正样本(即值为1的位置),因此产生疑问:是否需要自定义损失函数,或采用其他解决方案?我尝试编写了自定义BCEWithLogitsLoss,但效果不佳,相关实现代码如下:

原模型与训练代码

tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')
bert_model = BertModel.from_pretrained('bert-base-uncased').to(device)

# 解冻BERT层以在训练中微调
for param in bert_model.parameters():
    param.requires_grad = True

def set_bert_required_grad(value:bool = True):
    for param in bert_model.parameters():
        param.requires_grad = value

class TextDataset(Dataset):
    def __init__(self, texts):
        self.texts = texts

    def __len__(self):
        return len(self.texts)
    
    def __getitem__(self, index):
        return self.texts[index]

class PersonalityDetectionModel(nn.Module):
    def __init__(self):
        super(PersonalityDetectionModel, self).__init__()

        self.dropout = nn.Dropout(0.4)

        self.fc = nn.Linear(BERT_VARIANTS_CLS_LAYER_SIZE, 512)
        self.attention = nn.MultiheadAttention(embed_dim=512, num_heads=16, dropout=0.2, device=device)

        self.fc1 = nn.Linear(512, 256)
        self.fc2 = nn.Linear(256, 128)
        self.fc3 = nn.Linear(128, 16)

        self.relu = nn.ReLU()

    def forward(self, posts):        
        posts = self.dropout(posts)
        posts = self.relu(self.fc(posts))

        posts = posts.unsqueeze(1)  # 为多头注意力添加序列维度
        posts, _ = self.attention(posts, posts, posts)
        posts = posts.squeeze(1)  # 移除序列维度

        posts = self.relu(self.fc1(posts))
        posts = self.dropout(posts)

        posts = self.relu(self.fc2(posts))
        posts = self.dropout(posts)
        
        posts = self.fc3(posts)
        return posts

def encode_batch(texts_batch):
    encoded_inputs = tokenizer(texts_batch, padding=True, truncation=True, return_tensors="pt").to(device)
    output = bert_model(**encoded_inputs)
    return output.last_hidden_state[:, 0, :]  # 取CLS token的嵌入

# One-hot编码转换
def labels_to_multilabel(local_labels, num_classes=16):
    multilabels = torch.zeros((local_labels.size(0), num_classes), device=local_labels.device)
    
    for idx, label in enumerate(local_labels):
        multilabels[idx][label] = 1
    
    return multilabels

batch_size = 32

model = PersonalityDetectionModel().to(device)

optimizer = torch.optim.Adam([
    {'params': model.parameters(), 'lr': 1e-3},
    {'params': bert_model.parameters(), 'lr': 5e-5}
])

criterion = nn.BCEWithLogitsLoss()

type_to_label = { 
    "INTJ": 0,
    "INTP": 1,
    "INFJ": 2,
    "INFP": 3,
    "ENTJ": 4,
    "ENTP": 5,
    "ENFJ": 6,
    "ENFP": 7,
    "ISTJ": 8,
    "ISFJ": 9,
    "ISTP": 10,
    "ISFP": 11,
    "ESTJ": 12,
    "ESFJ": 13,
    "ESTP": 14,
    "ESFP": 15
}

train_texts, val_texts, train_labels, val_labels = train_test_split(
    df['posts'].to_list(), 
    df['type'].map(type_to_label).values, 
    test_size=0.8, 
    random_state=1337,
    shuffle=True
)

train_dataset = TextDataset(train_texts)
val_dataset = TextDataset(val_texts)

train_dataloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)
val_dataloader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)

def train_model():
    set_bert_required_grad()

    model.train()
    bert_model.train()

    for epoch in range(1):
        total_correct = 0
        total_samples = 0
        for i, texts_batch in enumerate(train_dataloader):
            encoded_batch = encode_batch(texts_batch)
            local_labels = torch.tensor(train_labels[i * batch_size : (i + 1) * batch_size]).to(device)

            multi_labels = labels_to_multilabel(local_labels)

            optimizer.zero_grad()

            outputs = model(encoded_batch)

            loss = criterion(outputs, multi_labels)

            loss.backward()

            optimizer.step()

            # 用sigmoid将logits转为概率
            probs = torch.sigmoid(outputs)

            # 将概率转为二分类预测
            predicted_labels = (probs > 0.5).float()

            # 计算正样本位置的正确预测数
            correct_predictions = ((predicted_labels == 1) & (multi_labels == 1)).float().sum().item()

            # 统计总正样本数
            total_positives = multi_labels.sum().item()

            total_correct += correct_predictions
            total_samples += total_positives  # 只统计正样本位置

            if i % 10 == 0: 
                accuracy = total_correct / total_samples if total_samples > 0 else 0.0
                print(f"Epoch [{epoch+1}/10], Step [{i+1}/{len(train_dataloader)}], Loss: {loss.item():.4f}, Accuracy: {accuracy:.4f}")

        epoch_accuracy = total_correct / total_samples if total_samples > 0 else 0.0
        print(f"Epoch [{epoch+1}/10] Accuracy: {epoch_accuracy:.4f}")

尝试的自定义损失函数代码

class CustomBCEWithLogitsLoss(nn.Module):
    def __init__(self):
        super(CustomBCEWithLogitsLoss, self).__init__()

    def forward(self, outputs, targets):
        probs = torch.sigmoid(outputs)
        
        loss = -targets * torch.log(probs + 1e-12)  # 添加小epsilon避免log(0)
        
        loss = loss.sum() / (targets.sum() + 1e-12)  # 添加epsilon避免除以0
        
        return loss

问题分析与优化方案

你的核心需求是让模型只关注正样本的预测损失,不需要完全自定义损失函数,原生工具或调整后的自定义实现都能解决问题,以下是具体方案:

方案1:复用原生BCEWithLogitsLoss,聚焦正样本损失

原生BCEWithLogitsLoss支持reduction='none'参数,能输出每个样本每个标签的损失值,之后我们可以只筛选正样本位置的损失进行平均计算,完全符合你的需求:

# 初始化损失函数,设置reduction='none'保留每个位置的损失
criterion = nn.BCEWithLogitsLoss(reduction='none')

# 在训练循环中计算损失:
loss_per_position = criterion(outputs, multi_labels)
# 只保留正样本位置的损失,然后求平均
loss = loss_per_position[multi_labels == 1].mean()
# 处理批次中无正样本的边界情况
if torch.isnan(loss):
    loss = torch.tensor(0.0, device=device)

如果希望给正样本更高的关注度,可以结合pos_weight参数,比如把正样本损失权重设为10(根据数据集分布调整):

pos_weight = torch.tensor([10.0]*16).to(device)
criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight, reduction='none')
# 后续损失计算逻辑同上

方案2:修复自定义损失函数的数值稳定性问题

你之前的自定义损失手动计算sigmoid和log容易出现数值不稳定(比如接近0或1时的梯度异常),改用PyTorch内置的F.binary_cross_entropy_with_logits可以避免这个问题,同时优化损失归一化逻辑:

import torch.nn.functional as F

class CustomPositiveBCEWithLogitsLoss(nn.Module):
    def __init__(self):
        super().__init__()

    def forward(self, outputs, targets):
        # 用内置函数计算逐位置损失,避免手动sigmoid的数值问题
        loss_per_pos = F.binary_cross_entropy_with_logits(outputs, targets, reduction='none')
        # 筛选正样本位置的损失
        positive_losses = loss_per_pos[targets == 1]
        # 返回正样本损失的平均值,无正样本时返回0
        return positive_losses.mean() if len(positive_losses) > 0 else torch.tensor(0.0, device=outputs.device)

额外建议:回归任务本质,改用单标签分类

你的任务其实是单标签分类(每个样本对应唯一的MBTI类型),用多标签one-hot+BCE的方式属于过度设计,改用CrossEntropyLoss直接做单标签分类更高效,模型会自动聚焦正确类别(正样本):

# 移除labels_to_multilabel函数,直接使用0-15的原始标签
criterion = nn.CrossEntropyLoss()

# 训练循环中的损失计算:
local_labels = torch.tensor(train_labels[i * batch_size : (i + 1) * batch_size]).to(device)
outputs = model(encoded_batch)
loss = criterion(outputs, local_labels)

# 准确率计算调整为单标签逻辑:
_, predicted_labels = torch.max(outputs, dim=1)
correct_predictions = (predicted_labels == local_labels).sum().item()
total_samples += len(local_labels)

另外注意你的train_test_split设置了test_size=0.8,这意味着训练集仅占20%,可能导致模型训练不足,建议调整为test_size=0.2或0.3。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 14:39:50