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

请求协助:基于RoBERTa与Jigsaw数据集微调多分类内容审核神经网络(PyTorch)

基于Jigsaw数据集将二分类模型转为多分类模型的实操步骤(PyTorch)

1. 确认多分类任务的类别配置

先明确目标任务的类别数量:比如Jigsaw Toxic Comment Classification Challenge包含6类(toxic、severe_toxic、obscene、threat、identity_hate、insult),后续所有调整都要匹配这个类别数。

2. 重新初始化适配多分类的模型

原代码加载的是预训练二分类模型,需要修改模型初始化时的num_labels参数,自动替换分类头为对应数量的输出层:

from transformers import AutoTokenizer, AutoModelForSequenceClassification

MODEL = "cardiffnlp/twitter-roberta-base-offensive"
tokenizer = AutoTokenizer.from_pretrained(MODEL)
# 这里指定num_labels为你的目标类别数,以Jigsaw Toxic的6类为例
model = AutoModelForSequenceClassification.from_pretrained(MODEL, num_labels=6)

3. 适配Jigsaw数据集的预处理逻辑

加载并处理Jigsaw数据集,将标签转换为模型可接受的格式:

from datasets import load_dataset
from torch.utils.data import DataLoader

def preprocess_text(examples):
    return tokenizer(examples["comment_text"], truncation=True, padding="max_length", max_length=128)

# 加载Jigsaw Toxic数据集(按需替换为你的目标Jigsaw子数据集)
dataset = load_dataset("jigsaw_toxicity_prediction")
# 批量处理文本
tokenized_dataset = dataset.map(preprocess_text, batched=True)

# 单标签多分类:直接使用类别索引;多标签则保留所有类别列并转为张量
# 以单标签为例(若为多标签,需调整标签列处理逻辑)
tokenized_dataset = tokenized_dataset.rename_column("toxic", "labels")
tokenized_dataset.set_format("torch", columns=["input_ids", "attention_mask", "labels"])

# 构建DataLoader
train_loader = DataLoader(tokenized_dataset["train"], batch_size=16, shuffle=True)

4. 调整损失函数与训练循环

根据任务类型(单标签/多标签)选择对应损失函数,修改训练逻辑:

import torch
import torch.nn as nn
from torch.optim import AdamW

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model.to(device)

# 单标签多分类用CrossEntropyLoss;多标签用BCEWithLogitsLoss
loss_fn = nn.CrossEntropyLoss()
optimizer = AdamW(model.parameters(), lr=5e-5)

model.train()
for batch in train_loader:
    batch = {k: v.to(device) for k, v in batch.items()}
    outputs = model(**batch)
    
    # 计算损失
    loss = loss_fn(outputs.logits, batch["labels"])
    
    # 反向传播与优化
    loss.backward()
    optimizer.step()
    optimizer.zero_grad()

5. 更新推理函数以输出多分类结果

修改原evaluate_text函数,适配多分类的输出展示:

def evaluate_text(text, class_names=["toxic", "severe_toxic", "obscene", "threat", "identity_hate", "insult"]):
    encoded_text = tokenizer(text, return_tensors='pt').to(device)
    
    model.eval()
    with torch.no_grad():
        outputs = model(**encoded_text)
    
    # 输出所有类别的概率
    probabilities = torch.softmax(outputs.logits, dim=1).squeeze().numpy()
    for cls, prob in zip(class_names, probabilities):
        print(f"{cls}: {prob:.4f}")
    
    # 输出预测的最可能类别
    top_idx = torch.argmax(outputs.logits, dim=1).item()
    print(f"\nTop predicted class: {class_names[top_idx]}")

# 测试
evaluate_text("Dude, that's amazing!")

关键注意事项

  • 任务类型区分:如果是多标签分类(一句话可属于多个类别),需将标签转为二进制张量,并使用BCEWithLogitsLoss,推理时用sigmoid替代softmax。
  • 类别不平衡处理:Jigsaw数据集存在严重的类别不平衡,可通过加权损失、数据增强或重采样缓解。
  • 微调充分性:替换分类头后,需在Jigsaw数据集上充分微调,避免预训练二分类的分类头参数影响多分类效果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 22:42:47