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

BertForSequenceClassification模型始终预测0的问题求助

问题:BERT二分类模型始终预测单一类别

我在最高法院判决预测数据集上使用BERT进行二分类任务时,遇到模型始终只预测一类的问题:

  • 原始数据中0标签占比约2/3,无论调整何种参数,模型准确率始终固定在67%;
  • 将标签分布调整为50/50的0和1后,准确率变为50%,说明模型完全在随机猜测单一类别。

数据预处理代码

import pandas as pd
import re
from sklearn.model_selection import train_test_split
from transformers import BertTokenizer
import torch

cases = pd.read_csv("justice.csv")
cases.drop(columns=['Unnamed: 0', 'ID', 'name', 'href', 'docket', 'term',  
                    'majority_vote', 'minority_vote', 'decision_type', 'disposition', 'issue_area'], inplace=True)
cases.dropna(inplace=True)

cases = cases.rename(columns={'first_party_winner': 'winning_party_idx'})
for i, row in cases.iterrows():
    if row['winning_party_idx'] == True:
        cases.loc[i, 'winning_party_idx'] = 0
    else:
        cases.loc[i, 'winning_party_idx'] = 1

# 创建镜像样本,交换双方以避免模型偏向第一方
mirrored_cases = cases.copy()
mirrored_cases['first_party'], mirrored_cases['second_party'] = mirrored_cases['second_party'], mirrored_cases['first_party']
mirrored_cases['winning_party_idx'] = (mirrored_cases['winning_party_idx'] == 0).astype(int)
mirrored_cases.reset_index(drop=True, inplace=True)

cases = pd.concat([cases, mirrored_cases])
cases.reset_index(drop=True, inplace=True)

cases['facts'] = cases['facts'].str.replace(r'<[^<]+?>', '', regex=True)
cases['facts'] = cases['facts'].apply(lambda x: re.sub(r'[^a-zA-Z0-9\'\s]', '', x))
#cases['facts'] = cases['facts'].str.lower()

def word_count(text):
  return len(text.split())

cases['facts_len'] = cases['facts'].apply(word_count)
cases['facts_len'].describe()

cases['facts'] = cases.loc[cases['facts_len'] <= 390, 'facts']
cases['facts'] = cases.apply(lambda x: f"{x['first_party']} [SEP] {x['second_party']} [SEP] {x['facts']}", axis=1)
cases = cases.drop(columns=['first_party', 'second_party', 'facts_len'])

train_facts, val_facts, train_winners,  val_winners = train_test_split(
    cases['facts'], cases['winning_party_idx'], test_size=0.20)

train_facts, val_facts = train_facts.tolist(), val_facts.tolist()
train_winners, val_winners = [str(i) for i in train_winners], [str(i) for i in val_winners]

tokenizer = BertTokenizer.from_pretrained('bert-base-cased')
train_encodings = tokenizer(train_facts, padding=True)
val_encodings = tokenizer(val_facts, padding=True)

class TextDataset(torch.utils.data.Dataset):
    def __init__(self, encodings, labels):
        self.encodings = encodings
        self.labels = labels

    def __getitem__(self, idx):
        item = {key: torch.tensor(val[idx]) for key, val in self.encodings.items()}
        item['labels'] = torch.tensor(int(self.labels[idx]))
        return item

    def __len__(self):
        return len(self.labels)

train_dataset = TextDataset(train_encodings, train_winners)
val_dataset = TextDataset(val_encodings, val_winners)

模型加载与训练代码

from transformers import BertForSequenceClassification, TrainingArguments, Trainer
import evaluate
import numpy as np

# 加载预训练模型
model = BertForSequenceClassification.from_pretrained('bert-base-cased', 
                                                      num_labels=2, 
                                                      hidden_dropout_prob=0.4,
                                                      attention_probs_dropout_prob=0.4)

training_args = TrainingArguments(
    output_dir="test_trainer", 
    logging_dir='logs', 
    evaluation_strategy="epoch",
    per_device_train_batch_size=32,  
    per_device_eval_batch_size=32,
    num_train_epochs=3,
    logging_steps=50,
)
metric = evaluate.load("accuracy")
def compute_metrics(eval_pred):
    logits, labels = eval_pred
    predictions = np.argmax(logits, axis=1)
    return metric.compute(predictions=predictions, references=labels)

trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=train_dataset,
    eval_dataset=val_dataset,
    compute_metrics=compute_metrics,
)

trainer.train()

问题分析与解决方案

核心原因

  1. 默认损失函数局限性:BertForSequenceClassification默认使用CrossEntropyLoss,该损失函数未考虑类别不平衡,会天然偏向多数类;当标签调整为均衡后模型仍瞎猜,说明数据预处理和训练设置存在其他问题。
  2. 数据预处理漏洞:
    • 处理文本长度时,cases['facts'] = cases.loc[cases['facts_len'] <= 390, 'facts']会将长文本样本的facts设为NaN,后续未过滤这些无效样本,导致模型输入缺失有效特征。
    • 分词未设置truncation=True,BERT最大输入长度为512,超长文本会引发输入异常,模型无法正常学习。
  3. 训练设置不足:
    • 未设置类别权重,无法抵消类别不平衡的影响;
    • 训练轮次仅3轮,模型未充分学习特征;默认学习率可能不适合当前任务。

修复方案

1. 修复数据预处理

  • 过滤超长文本样本,而非将facts设为NaN:
    cases = cases[cases['facts_len'] <= 390]
    
  • 分词时强制截断超长文本,符合BERT输入要求:
    train_encodings = tokenizer(train_facts, padding=True, truncation=True, max_length=512)
    val_encodings = tokenizer(val_facts, padding=True, truncation=True, max_length=512)
    

2. 使用带类别权重的损失函数

自定义Trainer的损失计算逻辑,给少数类更高权重:

from torch.nn import CrossEntropyLoss

class WeightedTrainer(Trainer):
    def compute_loss(self, model, inputs, return_outputs=False):
        labels = inputs.get("labels")
        outputs = model(**inputs)
        logits = outputs.get("logits")
        # 根据实际类别分布设置权重,示例中0类占比2/3,1类占1/3,权重设为1:2
        class_weights = torch.tensor([1.0, 2.0], device=model.device)
        loss_fct = CrossEntropyLoss(weight=class_weights)
        loss = loss_fct(logits.view(-1, self.model.config.num_labels), labels.view(-1))
        return (loss, outputs) if return_outputs else loss

# 替换原Trainer
trainer = WeightedTrainer(
    model=model,
    args=training_args,
    train_dataset=train_dataset,
    eval_dataset=val_dataset,
    compute_metrics=compute_metrics,
)

3. 优化训练参数

  • 增加训练轮次:num_train_epochs=5
  • 调整学习率:learning_rate=2e-5
  • 可选:加入梯度累积gradient_accumulation_steps=2,变相增大batch size

4. 完善评估指标

仅看准确率无法反映模型真实性能,补充精确率、召回率、F1分数:

metric = evaluate.combine(["accuracy", "precision", "recall", "f1"])

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 04:12:02