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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 17:15:18