TfBertForSequenceClassification多分类工作原理及实例解析
TfBertForSequenceClassification 多分类任务实现(5标签实例)
核心逻辑
TfBertForSequenceClassification是基于BERT封装的序列分类模型,默认支持二分类,只需通过num_labels参数指定分类数量(这里设为5),模型会自动在BERT的<[BOS_never_used_51bce0c785ca2f68081bfa7d91973934]>token输出后添加对应维度的分类层,将特征映射到5个标签的概率分布,最终通过交叉熵损失完成训练。
具体实现步骤(结合200句+5标签数据)
1. 数据预处理
先把原始文本和标签转换成模型可识别的格式:
- 标签编码:将Q、U、E、R、Y这5个标签转换为0-4的整数(比如Q→0、U→1、E→2、R→3、Y→4),用
LabelEncoder即可完成。 - 文本编码:用BERT配套的tokenizer对句子进行分词、截断/填充到固定长度,生成
input_ids、attention_mask等模型输入特征。
2. 模型初始化
直接指定分类数量为5,模型会自动构建适配的分类层:
from transformers import TfBertForSequenceClassification, BertTokenizer # 加载预训练tokenizer与模型(中文用bert-base-chinese,英文用bert-base-uncased) tokenizer = BertTokenizer.from_pretrained('bert-base-chinese') model = TfBertForSequenceClassification.from_pretrained('bert-base-chinese', num_labels=5)
3. 构建TensorFlow数据集
把预处理后的特征和标签转换成TF Dataset,方便批量训练:
import tensorflow as tf from sklearn.preprocessing import LabelEncoder import pandas as pd # 假设数据存储在csv文件,包含text(句子)和label(标签)两列 df = pd.read_csv('your_train_data.csv') # 标签转整数编码 le = LabelEncoder() df['label'] = le.fit_transform(df['label']) # 文本编码处理 encoded_inputs = tokenizer( df['text'].tolist(), padding=True, truncation=True, max_length=128, return_tensors='tf' ) # 构建可训练的数据集,打乱后按批次加载 dataset = tf.data.Dataset.from_tensor_slices(( {'input_ids': encoded_inputs['input_ids'], 'attention_mask': encoded_inputs['attention_mask']}, df['label'].values )).shuffle(200).batch(8)
4. 模型编译与训练
选择适配多分类的优化器和损失函数:
model.compile( optimizer=tf.keras.optimizers.Adam(learning_rate=2e-5), # 标签是整数编码,用SparseCategoricalCrossentropy;若为one-hot则用CategoricalCrossentropy loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True), metrics=['accuracy'] ) # 小样本(200条)训练,epochs设5-10即可 model.fit(dataset, epochs=8)
5. 预测与标签解码
训练完成后,对新句子进行预测,并将整数标签转回原始标签:
def predict_sentence_label(text): # 对单句编码 inputs = tokenizer(text, return_tensors='tf', padding=True, truncation=True) # 获取模型输出的logits,转成预测标签ID logits = model(inputs).logits pred_label_id = tf.argmax(logits, axis=1).numpy()[0] # 转回原始标签 return le.inverse_transform([pred_label_id])[0] # 测试示例 test_text = "这是一条需要分类的测试句子" print(predict_sentence_label(test_text)) # 输出Q/U/E/R/Y中的一个
关键注意点
- 分类层适配:
num_labels=5会让模型分类层输出5维logits,对应5个标签的未归一化概率。 - 小样本优化:200条数据属于小样本,建议用更小的学习率(如2e-5),也可以先冻结BERT预训练层,只训练分类头,再逐步解冻微调。
内容的提问来源于stack exchange,提问作者XYZ
相关产品推荐
相关产品推荐

