请求协助:基于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
相关产品推荐
相关产品推荐

