聊聊基于BERT模型实现多标签分类任务的实践与思考
概述
以预训练大模型为基座的神经网络模型,借助预训练带来的泛化能力,再结合微调后的领域针对性,已经成为NLP任务的常见解决方案。核心思路其实很直观:先用海量数据让模型学会“理解语言”,再通过少量领域数据让它学会“解决具体问题”。

GitHub上有一个轻量级的仓库——multi_label_classification,基于BERT预训练模型实现了多标签分类。通过对这个仓库源码的拆解,可以比较清晰地看到从模型定义到训练推理的完整链路。下面就来一步步看看里面的逻辑。
代码文件的注释如下:
模型类定义
基于BERT预训练模型,加上一个线性层和sigmoid激活函数,用来输出多标签分类的概率分布——这就是整个模型的核心骨架。
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import BertModel
# 基于BERT模型的多分类模型定义
class BertMultiLabelCls(nn.Module):
def __init__(self, hidden_size, class_num, dropout=0.1):
super(BertMultiLabelCls, self).__init__()
# 定义线性层(密集层)
# 输入的维度是hidden_size;输出的维度是多分类的个数。
self.fc = nn.Linear(hidden_size, class_num)
self.drop = nn.Dropout(dropout)
# 加载预训练模型
self.bert = BertModel.from_pretrained("bert-base-chinese")
def forward(self, input_ids, attention_mask, token_type_ids):
# 将输入的input_ids(文本token的ID),
# attention_mask(表示哪些token是重要的)
# token_type_ids(区分不同类型的token,如句子A和句子B)传递给BERT模型,获取模型的输出。
outputs = self.bert(input_ids, attention_mask, token_type_ids)
# 打印输出
# print(outputs)
# 从BERT模型的输出中选择第一个元素(通常是[CLS]标记的输出),然后通过dropout层。
cls = self.drop(outputs[1])
# 将经过dropout层的[CLS]标记的输出传递到全连接层self.fc,然后应用sigmoid激活函数,将输出转换为概率分布。
out = F.sigmoid(self.fc(cls))
# 返回最终的输出概率分布,用于多分类任务。
return out
这里有个细节值得注意:通常二分类问题用Sigmoid,多分类则用Softmax。但这里的模型用的是Sigmoid,原因其实很简单——我们处理的是多标签分类,每个标签之间并不是互斥的,所以需要为每个标签独立输出一个0-1之间的概率值。后面还会再提到这一点。
数据预处理
任何模型训练之前,数据预处理都是绕不开的一环。原始数据必须按照模型规定的格式整理好,才能送入训练流程。不同模型有不同的格式要求,比如alpaca格式、sharegpt格式等等,这里需要把原始JSON数据转成模型可接受的格式。
# -*- coding: utf-8 -*-
import json
import torch
from torch.utils.data import Dataset
from transformers import BertTokenizer
from data_preprocess import load_json
class MultiClsDataSet(Dataset):
def __init__(self, data_path, max_len=128, label2idx_path="./data/label2idx.json"):
self.label2idx = load_json(label2idx_path)
self.class_num = len(self.label2idx)
self.tokenizer = BertTokenizer.from_pretrained("bert-base-chinese")
self.max_len = max_len
self.input_ids, self.token_type_ids, self.attention_mask, self.labels = self.encoder(data_path)
def encoder(self, data_path):
texts = []
labels = []
with open(data_path, encoding="utf-8") as f:
for line in f:
line = json.loads(line)
texts.append(line["text"])
tmp_label = [0] * self.class_num
for label in line["label"]:
tmp_label[self.label2idx[label]] = 1
labels.append(tmp_label)
tokenizers = self.tokenizer(texts,
padding=True,
truncation=True,
max_length=self.max_len,
return_tensors="pt",
is_split_into_words=False)
input_ids = tokenizers["input_ids"]
token_type_ids = tokenizers["token_type_ids"]
attention_mask = tokenizers["attention_mask"]
return input_ids, token_type_ids, attention_mask,
torch.tensor(labels, dtype=torch.float)
def __len__(self):
return len(self.labels)
def __getitem__(self, item):
return self.input_ids[item], self.attention_mask[item],
self.token_type_ids[item], self.labels[item]
if __name__ == '__main__':
dataset = MultiClsDataSet(data_path="./data/train.json")
print(dataset.input_ids)
print(dataset.token_type_ids)
print(dataset.attention_mask)
print(dataset.labels)
# -*- coding: utf-8 -*-
"""
数据预处理
"""
import json
def load_json(data_path):
with open(data_path, encoding="utf-8") as f:
return json.loads(f.read())
def dump_json(project, out_path):
with open(out_path, "w", encoding="utf-8") as f:
json.dump(project, f, ensure_ascii=False)
def preprocess(train_data_path, label2idx_path, max_len_ratio=0.9):
"""
:param train_data_path:
:param label2idx_path:
:param max_len_ratio:
:return:
"""
labels = []
text_length = []
with open(train_data_path, encoding="utf-8") as f:
for data in f:
data = json.loads(data)
text_length.append(len(data["text"]))
labels.extend(data["label"])
labels = list(set(labels))
label2idx = {label: idx for idx, label in enumerate(labels)}
dump_json(label2idx, label2idx_path)
text_length.sort()
print("当设置max_len={}时,可覆盖{}的文本".format(text_length[int(len(text_length)*max_len_ratio)], max_len_ratio))
if __name__ == '__main__':
preprocess("./data/train.json", "./data/label2idx.json")
训练
训练的源码如下:
# -*- coding: utf-8 -*-
import torch
import torch.nn as nn
from torch.utils.data import DataLoader
from transformers import AdamW
import numpy as np
from data_preprocess import load_json
from bert_multilabel_cls import BertMultiLabelCls
from data_helper import MultiClsDataSet
from sklearn.metrics import accuracy_score
train_path = "./data/train.json"
dev_path = "./data/dev.json"
test_path = "./data/test.json"
label2idx_path = "./data/label2idx.json"
sa ve_model_path = "./model/multi_label_cls.pth"
label2idx = load_json(label2idx_path)
class_num = len(label2idx)
device = "cuda" if torch.cuda.is_a vailable() else "cpu"
lr = 2e-5
batch_size = 128
max_len = 128
hidden_size = 768
epochs = 10
# 预处理数据
train_dataset = MultiClsDataSet(train_path, max_len=max_len, label2idx_path=label2idx_path)
dev_dataset = MultiClsDataSet(dev_path, max_len=max_len, label2idx_path=label2idx_path)
# 从数据集中 批量 加载数据,批大小为batch_size
train_dataloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)
dev_dataloader = DataLoader(dev_dataset, batch_size=batch_size, shuffle=False)
# 计算准确率
def get_acc_score(y_true_tensor, y_pred_tensor):
y_pred_tensor = (y_pred_tensor.cpu() > 0.5).int().numpy()
y_true_tensor = y_true_tensor.cpu().numpy()
return accuracy_score(y_true_tensor, y_pred_tensor)
# 训练
def train():
model = BertMultiLabelCls(hidden_size=hidden_size, class_num=class_num)
# 启用 batch normalization 和 dropout 。
model.train()
model.to(device)
# 定义优化器
optimizer = AdamW(model.parameters(), lr=lr)
# 定义了一个二进制交叉熵损失函数(BCELoss),用于多标签分类问题,因为它可以处理多个标签。
criterion = nn.BCELoss()
dev_best_acc = 0.
# 按epoch训练,即训练轮数
for epoch in range(1, epochs):
# 启用 batch normalization 和 dropout 。
model.train()
# 按batch训练,即训练批次
for i, batch in enumerate(train_dataloader):
# 清空梯度
optimizer.zero_grad()
batch = [d.to(device) for d in batch]
# 获取批数据中的标签label数据
labels = batch[-1]
# 执行预训练模型的forward方法
logits = model(*batch[:3])
# 通过二分类交叉熵损失,计算模型返回值与标签实际值的损失概率
loss = criterion(logits, labels)
# 反向传播
loss.backward()
# 梯度更新
optimizer.step()
# 打印数据
if i % 100 == 0:
acc_score = get_acc_score(labels, logits)
print("Train epoch:{} step:{} acc: {} loss:{} ".format(epoch, i, acc_score, loss.item()))
# 验证集合
dev_loss, dev_acc = dev(model, dev_dataloader, criterion)
print("Dev epoch:{} acc:{} loss:{}".format(epoch, dev_acc, dev_loss))
if dev_acc > dev_best_acc:
dev_best_acc = dev_acc
torch.sa ve(model.state_dict(), sa ve_model_path)
# 测试
test_acc = test(sa ve_model_path, test_path)
print("Test acc: {}".format(test_acc))
# 验证
def dev(model, dataloader, criterion):
all_loss = []
# 切换成评估模式
model.eval()
true_labels = []
pred_labels = []
with torch.no_grad():
for i, batch in enumerate(dataloader):
input_ids, attention_mask, token_type_ids, labels = [d.to(device) for d in batch]
logits = model(input_ids=input_ids, attention_mask=attention_mask, token_type_ids=token_type_ids)
loss = criterion(logits, labels)
all_loss.append(loss.item())
true_labels.append(labels)
pred_labels.append(logits)
true_labels = torch.cat(true_labels, dim=0)
pred_labels = torch.cat(pred_labels, dim=0)
acc_score = get_acc_score(true_labels, pred_labels)
return np.mean(all_loss), acc_score
# 测试
def test(model_path, test_data_path):
test_dataset = MultiClsDataSet(test_data_path, max_len=max_len, label2idx_path=label2idx_path)
test_dataloader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)
model = BertMultiLabelCls(hidden_size=hidden_size, class_num=class_num)
model.load_state_dict(torch.load(model_path))
model.to(device)
# 切换成评估模式
model.eval()
true_labels = []
pred_labels = []
with torch.no_grad():
for i, batch in enumerate(test_dataloader):
input_ids, attention_mask, token_type_ids, labels = [d.to(device) for d in batch]
logits = model(input_ids=input_ids, attention_mask=attention_mask, token_type_ids=token_type_ids)
true_labels.append(labels)
pred_labels.append(logits)
true_labels = torch.cat(true_labels, dim=0)
pred_labels = torch.cat(pred_labels, dim=0)
acc_score = get_acc_score(true_labels, pred_labels)
return acc_score
if __name__ == '__main__':
train()
上面的代码中间出现了一些我们非常熟悉的面孔:epochs、batch_size、max_len、hidden_size。这些参数通常都会在一个配置文件里集中管理,并且大多数都有默认值。以ChatGLM3-6B的配置文件为例:
{
"_name_or_path": "THUDM/chatglm-6b",
"architectures": [
"ChatGLMModel"
],
"bos_token_id": 130004,
"eos_token_id": 130005,
"mask_token_id": 130000,
"gmask_token_id": 130001,
"pad_token_id": 3,
"hidden_size": 4096,
"inner_hidden_size": 16384,
"layernorm_epsilon": 1e-05,
"max_sequence_length": 2048,
"model_type": "chatglm",
"num_attention_heads": 32,
"num_layers": 28,
"position_encoding_2d": true,
"torch_dtype": "float16",
"use_cache": true,
"vocab_size": 130528
}
每个batch的训练流程其实非常标准:
- 清空梯度
- 获取实际的标签
- 通过预训练模型输出预测值
- 基于二分类交叉熵损失函数,计算实际值与预测值的损失
- 将损失反向传播
- 更新梯度
严格来说,多分类任务通常应该用CrossEntropyLoss或者NLLLoss(Negative Log Likelihood Loss),尤其是CrossEntropyLoss,它几乎就是多分类问题的标配。但这里偏偏用了二分类交叉熵损失函数——这个方法实际上是把多标签问题“拆”成了多个二分类子问题来处理,每个标签独立计算损失。这也是为什么前面的模型定义中用了Sigmoid而不是Softmax。
预测
推理预测的代码如下:
# -*- coding: utf-8 -*-
import torch
from data_preprocess import load_json
from bert_multilabel_cls import BertMultiLabelCls
from transformers import BertTokenizer
hidden_size = 768
class_num = 3
label2idx_path = "./data/label2idx.json"
sa ve_model_path = "./model/multi_label_cls.pth"
label2idx = load_json(label2idx_path)
idx2label = {idx: label for label, idx in label2idx.items()}
device = "cuda" if torch.cuda.is_a vailable() else "cpu"
tokenizer = BertTokenizer.from_pretrained("bert-base-chinese")
max_len = 128
model = BertMultiLabelCls(hidden_size=hidden_size, class_num=class_num)
model.load_state_dict(torch.load(sa ve_model_path))
model.to(device)
# 切换评估模式
model.eval()
def predict(texts):
# 加载分词器分词
outputs = tokenizer(texts, return_tensors="pt", max_length=max_len,
padding=True, truncation=True)
# 加载模型
logits = model(outputs["input_ids"].to(device),
outputs["attention_mask"].to(device),
outputs["token_type_ids"].to(device))
logits = logits.cpu().tolist()
# print(logits)
result = []
for sample in logits:
pred_label = []
for idx, logit in enumerate(sample):
if logit > 0.5:
pred_label.append(idx2label[idx])
result.append(pred_label)
return result
if __name__ == '__main__':
texts = ["中超-德尔加多扳平郭田雨绝杀 泰山2-1逆转亚泰", "今日沪深两市指数整体呈现震荡调整格局"]
result = predict(texts)
print(result)
推理预测的流程对用户来说非常简洁:只有两个核心组件——分词器和模型。分词器把输入文本处理成模型需要的各种维度数据,然后模型基于这些数据输出分类的概率分布。整个推理步骤清晰,关键在于如何设定阈值(这里的0.5)来判定哪个标签被激活。
总结
回过头来看,所谓“基于预训练大模型来做解决方案”,本质上就是接入一个已经具备强大语言理解能力的大模型,然后在它的基础上训练或微调出一个特定领域的“子任务”。这个子任务可能是分类、抽取、生成,不一而足——但目标很明确:让大模型在某个具体的任务上表现出更强的领域性和专业性。
说到底,它还是一个神经网络建模与训练的过程:以基座大模型为起点,再单独训练一个多分类(或多标签分类)的神经网络,以满足特定任务的需求。而模型的大部分基本信息,仍然是继承自基座大模型的。
如果没有接触过这类流程,建议先从基于PyTorch搭建一个简单的神经网络模型开始,理解建模、训练和推理的完整链路,然后再回头对比这里的实现——两者在本质上是一回事。
-
- 关于宇宙的好的网名有哪些
- 角色扮演 | 1
- 网名