基于BERT的MBTI多标签分类:是否需自定义BCEWithLogitsLoss?
我正在开展基于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

