简介:文本分类是自然语言处理(NLP)中的一项基础且核心的任务,其目标是将文本文档自动划分到预定义的类别中。其原理在于通过机器学习或深度学习模型,从文本中提取特征并学习类别间的决策边界。这项技术的价值在于能够自动化处理海量文本信息,极大地提升信息组织和检索的效率。在实际应用中,文本分类被广泛用于新闻分类、情感分析、垃圾邮件过滤、意图识别等场景。随着预训练语言模型的出现,如BERT(Bidirectional Encoder Representations from Transformers),文本分类的性能得到了显著提升。BERT通过在大规模语料上进行预训练,能够生成深度的上下文相关词向量,为下游任务提供了强大的语义理解基础。本文将聚焦于一个经典的中文文本分类实战项目,即利用BERT模型在THUCNews数据集上进行微调。THUCNews是一个由清华大学整理的大规模中文新闻数据集,以其类别平衡、文本规整的特点,成为评测模型性能和进行算法学习的理想基准。项目将详细阐述从环境搭建、数据预处理、模型构建、训练调优到最终部署的完整流程,并深入探讨在实践过程中可能遇到的挑战,如显存优化、过拟合处理等,旨在为开发者提供一个可复现、可优化的工程实践指南。
1. 从零开始:为什么选择THUCNews与BERT做中文文本分类
如果你正在寻找一个能跑通、有挑战性、且能真正学到东西的中文NLP实战项目,那么“基于THUCNews数据集的BERT文本分类”绝对是一个经典且理想的选择。这听起来可能像是一个教科书式的入门任务,但当你真正动手时,会发现从数据预处理到模型微调,再到最后的性能调优,每一步都藏着不少“坑”和“门道”。我最初接触这个组合时,也以为照着教程跑一遍就完事了,结果在数据编码、标签对齐、显存优化上接连碰壁。今天,我就把自己趟过的路、踩过的坑,以及最终沉淀下来的一套可复现、可优化的完整流程分享给你。
THUCNews是清华大学整理的一个大型中文新闻文本数据集,涵盖了10个类别(如体育、财经、房产等),总计约74万篇文档。它的价值在于“干净”且“规整”——类别平衡、文本长度适中、噪声相对较少,非常适合作为模型能力评测和算法学习的基准。而BERT(Bidirectional Encoder Representations from Transformers)作为Transformer编码器结构的代表,通过“掩码语言模型”和“下一句预测”任务进行预训练,能生成深度的上下文相关的词向量。将预训练好的BERT模型在THUCNews上进行微调,本质上就是让这个已经具备强大语言理解能力的“大脑”,快速学习如何针对新闻文本进行精准的类别判断。
这个项目适合所有希望深入理解如何将预训练模型应用于实际中文任务的开发者。无论你是想验证一个新想法的有效性,还是为公司的文本审核、新闻聚合系统搭建原型,这个流程都能提供一个坚实的起点。接下来,我会带你从环境搭建、数据剖析开始,一步步深入到模型微调、训练技巧和结果分析,过程中我会重点解释“为什么这么做”,而不仅仅是“怎么做”。
2. 环境搭建与数据初探:避开第一个“暗礁”
工欲善其事,必先利其器。在开始写第一行代码之前,环境的正确配置能避免后续无数莫名其妙的报错。同时,花时间真正理解你的数据,是模型成功的一半。
2.1 构建稳定可复现的Python环境
我强烈建议使用conda或venv创建独立的虚拟环境,这能保证项目依赖的纯净性。核心的库包括:
- 深度学习框架:
PyTorch或TensorFlow。我个人更倾向于PyTorch,因其动态图机制在调试和实验时更为灵活。安装时务必去官网根据你的CUDA版本选择正确的命令。 - Transformer库:
Hugging Face Transformers。这是我们的核心武器库,它提供了BERT等数千个预训练模型的简易加载和微调接口。 - 数据处理:
pandas用于数据操作,scikit-learn用于评估指标和数据集划分。 - 中文分词与处理:
jieba(虽然BERT用不到分词,但数据清洗时可能有用)。
你可以创建一个requirements.txt文件来管理依赖:
torch>=1.9.0 transformers>=4.15.0 pandas>=1.3.0 scikit-learn>=0.24.0 tqdm>=4.62.0通过pip install -r requirements.txt一键安装。这里有个小坑:transformers和torch的版本有时存在兼容性问题,如果遇到奇怪的错误,尝试指定稍旧但稳定的版本组合(如transformers==4.18.0,torch==1.10.0)往往是有效的解决方案。
2.2 深入解析THUCNews数据集结构
下载并解压THUCNews数据集后,你会发现它通常按类别文件夹组织,每个文件夹内是大量的纯文本文件。第一步不是急着全部读入,而是先做探索性数据分析(EDA)。
首先,查看数据分布:
import os data_path = './THUCNews' categories = os.listdir(data_path) for cat in categories: cat_path = os.path.join(data_path, cat) file_count = len([f for f in os.listdir(cat_path) if f.endswith('.txt')]) print(f'类别 [{cat}] 文件数量: {file_count}')这个简单的脚本能让你立刻看到各类别的样本数是否均衡。THUCNews在这个方面做得不错,但如果你发现某个类别样本极少(例如只有几百个),在划分训练验证集时就要考虑使用分层抽样,或后续使用过采样技术。
其次,审视文本内容:随机打开几个文件,看看文本的格式。你可能会发现一些问题:
- 无关字符:文件可能包含URL、邮箱地址、特殊符号(如“◆”、“★”)。
- 标题与正文混杂:很多新闻文件的第一行是标题,后面是正文,中间可能用换行或空格分隔。
- 文本长度差异巨大:短消息可能几十字,长报道可能数千字。
BERT模型有最大序列长度限制(通常是512个token)。你需要统计文本长度的分布,以决定一个合适的截断长度。计算文本的字符数(中文字符算一个)分布:
import numpy as np lengths = [] for cat in categories: cat_path = os.path.join(data_path, cat) for file in os.listdir(cat_path)[:100]: # 先抽样一部分看分布 with open(os.path.join(cat_path, file), 'r', encoding='utf-8') as f: text = f.read().strip() lengths.append(len(text)) print(f"平均长度: {np.mean(lengths):.0f}, 最大长度: {max(lengths)}, 最小长度: {min(lengths)}") print(f"95分位长度: {np.percentile(lengths, 95):.0f}") # 这个值对设定max_length非常关键在我的经验中,THUCNews的文本95%以上都在1000字符以内,考虑到BERT一个中文字符通常是一个token,设置max_length=256或384已经可以覆盖大部分有效信息,同时能极大减少计算和显存开销。这是一个重要的权衡:太短会丢失信息,太长则训练缓慢且容易过拟合。
3. 数据预处理流水线:为BERT准备“标准餐”
原始文本不能直接喂给BERT。我们需要构建一个高效、可复用的数据预处理流水线,将文本文件转化为模型可以消化的数值化张量。这个过程的核心是Dataset和DataLoader。
3.1 构建自定义Dataset类
我们创建一个NewsDataset类,继承自PyTorch的Dataset。它的核心任务是在__getitem__方法中,根据索引返回一条已经完成分词、编码、并添加了特殊标记的文本数据及其标签。
from torch.utils.data import Dataset from transformers import BertTokenizer import torch class NewsDataset(Dataset): def __init__(self, texts, labels, tokenizer, max_len): self.texts = texts self.labels = labels self.tokenizer = tokenizer self.max_len = max_len def __len__(self): return len(self.texts) def __getitem__(self, idx): text = str(self.texts[idx]) label = self.labels[idx] # 关键步骤:使用tokenizer进行编码 encoding = self.tokenizer.encode_plus( text, add_special_tokens=True, # 添加[CLS]和[SEP] max_length=self.max_len, padding='max_length', # 填充到max_length truncation=True, # 过长则截断 return_attention_mask=True, return_tensors='pt', # 返回PyTorch张量 ) return { 'input_ids': encoding['input_ids'].flatten(), 'attention_mask': encoding['attention_mask'].flatten(), 'labels': torch.tensor(label, dtype=torch.long) }这里有几个至关重要的细节:
- Tokenizer的选择:对于中文BERT,你应该使用
BertTokenizer,并指定对应的预训练模型,例如bert-base-chinese。tokenizer会自动处理中文的分词(基于字粒度),并添加[CLS](用于分类)和[SEP](分隔符)等特殊标记。 padding和truncation:我们设定padding='max_length',保证所有样本长度一致,便于批量训练。truncation=True确保超长文本被截断。截断策略通常是“从尾部截断”,因为新闻的摘要信息多在开头。attention_mask:这个张量非常重要,它告诉模型哪些位置是真实的token(值为1),哪些是填充的(值为0)。模型在计算注意力时会忽略填充位置。
3.2 划分数据集与创建DataLoader
在将数据装入Dataset之前,我们需要先划分训练集、验证集和测试集。通常按照8:1:1的比例。
from sklearn.model_selection import train_test_split # 假设all_texts和all_labels是之前加载的所有文本和标签列表 train_texts, temp_texts, train_labels, temp_labels = train_test_split( all_texts, all_labels, test_size=0.2, random_state=42, stratify=all_labels) val_texts, test_texts, val_labels, test_labels = train_test_split( temp_texts, temp_labels, test_size=0.5, random_state=42, stratify=temp_labels) # 初始化tokenizer from transformers import BertTokenizer tokenizer = BertTokenizer.from_pretrained('bert-base-chinese') MAX_LEN = 256 # 创建Dataset实例 train_dataset = NewsDataset(train_texts, train_labels, tokenizer, MAX_LEN) val_dataset = NewsDataset(val_texts, val_labels, tokenizer, MAX_LEN) test_dataset = NewsDataset(test_texts, test_labels, tokenizer, MAX_LEN)接下来,使用DataLoader进行批量加载。这里需要设置两个关键参数:
batch_size:根据你的GPU显存决定。对于bert-base-chinese,在12GB显存的GPU上,batch_size=16或32是安全的起点。collate_fn:虽然我们的Dataset已经完成了填充,使得一个batch内长度一致,无需自定义collate_fn,但了解这个概念有好处。如果Dataset返回的是不等长的序列,则需要一个collate_fn函数来动态填充同一个batch内的数据。
from torch.utils.data import DataLoader BATCH_SIZE = 32 train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True) val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False) test_loader = DataLoader(test_dataset, batch_size=BATCH_SIZE, shuffle=False)shuffle=True仅在训练集上使用,以打乱数据顺序,让模型学习更泛化。
4. 模型构建与微调策略:让BERT“学以致用”
有了准备好的数据,我们就可以开始构建模型了。微调BERT进行文本分类,通常是在BERT模型顶部添加一个简单的全连接分类层。
4.1 定义分类模型架构
我们创建一个继承自nn.Module的类。BertModel输出的是序列中每个token的表示,而对于文本分类任务,我们通常只使用[CLS]token的表示(其原始设计就是用于聚合整个序列的信息进行分类)。
import torch.nn as nn from transformers import BertModel class BertNewsClassifier(nn.Module): def __init__(self, n_classes): super(BertNewsClassifier, self).__init__() self.bert = BertModel.from_pretrained('bert-base-chinese') self.drop = nn.Dropout(p=0.3) # Dropout层防止过拟合 self.out = nn.Linear(self.bert.config.hidden_size, n_classes) # 输出层 def forward(self, input_ids, attention_mask): # 通过BERT模型,return_dict=True是为了使用更清晰的字典格式输出 outputs = self.bert( input_ids=input_ids, attention_mask=attention_mask, return_dict=True ) # 取[CLS] token的隐藏状态 (形状: [batch_size, hidden_size]) pooled_output = outputs.pooler_output # 或者用 outputs.last_hidden_state[:, 0] # 通过Dropout和分类层 output = self.drop(pooled_output) return self.out(output)关键点解析:
pooler_outputvslast_hidden_state[:, 0]:pooler_output是BERT模型内部已经对[CLS]token的输出做过一次线性变换和Tanh激活的结果。而last_hidden_state[:, 0]是[CLS]token最原始的最后一层隐藏状态。根据Hugging Face文档和许多实践,对于分类任务,直接使用pooler_output通常效果就很好,这也是默认做法。- Dropout:这是一个简单但强大的正则化技术。在训练时,它随机“关闭”一部分神经元,迫使网络不过度依赖某些特定的特征。
p=0.3或0.5是常见值,你可以将其视为一个可调的超参数。 - 冻结BERT参数:在微调初期,有时会先冻结BERT的大部分底层参数,只训练顶部的分类层,待训练稳定后再解冻所有参数进行精细微调。这对于小数据集尤其有效,可以防止过拟合。但THUCNews数据集规模尚可,通常直接进行全参数微调效果更好。
4.2 配置训练超参数与优化器
训练Transformer模型,优化器的选择和学习率的设置至关重要。
from transformers import AdamW, get_linear_schedule_with_warmup # 初始化模型,假设有10个新闻类别 model = BertNewsClassifier(n_classes=10) model = model.to(device) # 将模型移动到GPU EPOCHS = 4 # BERT微调通常3-5个epoch就足够,过多容易过拟合 optimizer = AdamW(model.parameters(), lr=2e-5, correct_bias=False) # AdamW是Adam优化器的权重衰减修正版,非常适合Transformer模型。 # 学习率lr=2e-5是一个经典的起点,也被称为“BERT黄金学习率”。 total_steps = len(train_loader) * EPOCHS # 学习率调度器:先线性预热,再线性衰减 scheduler = get_linear_schedule_with_warmup( optimizer, num_warmup_steps=0.1 * total_steps, # 前10%的步数用于预热 num_training_steps=total_steps ) criterion = nn.CrossEntropyLoss().to(device) # 多分类交叉熵损失为什么是这些设置?
- 学习率(2e-5):预训练模型已经学到了很好的通用特征,微调时只需要小幅调整。过大的学习率会“冲毁”这些预训练权重,导致模型性能下降甚至无法收敛。
- 学习率预热(Warmup):在训练开始时,模型权重是随机初始化的分类头和预训练好的BERT体。直接使用较大的学习率可能导致训练不稳定。预热阶段让学习率从0线性增加到设定值,给模型一个“热身”的过程。
- 偏差校正(correct_bias=False):在原始的AdamW论文中建议,当使用权重衰减时,应关闭Adam优化器的偏差校正。
Transformers库的AdamW默认已做此处理,显式写出是为了让你明白这个细节。
5. 训练循环与验证:监控模型“学习状态”
训练循环是模型学习的核心。我们需要编写代码来迭代数据、计算损失、反向传播、更新权重,并在验证集上评估模型表现,防止过拟合。
5.1 编写训练与验证函数
一个健壮的训练循环应包括损失计算、梯度清零、反向传播、梯度裁剪(可选)和参数更新。
def train_epoch(model, data_loader, criterion, optimizer, device, scheduler): model = model.train() # 设置为训练模式 total_loss = 0 correct_predictions = 0 for batch in data_loader: input_ids = batch['input_ids'].to(device) attention_mask = batch['attention_mask'].to(device) labels = batch['labels'].to(device) # 前向传播 outputs = model(input_ids=input_ids, attention_mask=attention_mask) loss = criterion(outputs, labels) _, preds = torch.max(outputs, dim=1) # 获取预测类别 correct_predictions += torch.sum(preds == labels) total_loss += loss.item() # 反向传播 loss.backward() # 梯度裁剪,防止梯度爆炸,对于BERT通常不是必须,但是个好习惯 nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() scheduler.step() optimizer.zero_grad() # 清空梯度,非常重要! avg_loss = total_loss / len(data_loader) accuracy = correct_predictions.double() / len(data_loader.dataset) return avg_loss, accuracy def eval_model(model, data_loader, criterion, device): model = model.eval() # 设置为评估模式 total_loss = 0 correct_predictions = 0 with torch.no_grad(): # 关闭梯度计算,节省内存和计算 for batch in data_loader: input_ids = batch['input_ids'].to(device) attention_mask = batch['attention_mask'].to(device) labels = batch['labels'].to(device) outputs = model(input_ids=input_ids, attention_mask=attention_mask) loss = criterion(outputs, labels) _, preds = torch.max(outputs, dim=1) correct_predictions += torch.sum(preds == labels) total_loss += loss.item() avg_loss = total_loss / len(data_loader) accuracy = correct_predictions.double() / len(data_loader.dataset) return avg_loss, accuracy注意事项:
model.train()和model.eval():这会影响Dropout、BatchNorm等层的行为。训练时必须用.train(),评估时必须用.eval()。optimizer.zero_grad():PyTorch的梯度是累加的,必须在每次反向传播前清空上一次的梯度,否则梯度会不断累积,导致训练出错。with torch.no_grad():在验证和测试时,我们不需要计算梯度。这个上下文管理器可以显著减少内存消耗并加速计算。
5.2 执行训练与早停策略
现在,我们可以运行多个epoch的训练,并在每个epoch后验证模型性能。
best_accuracy = 0 for epoch in range(EPOCHS): print(f'Epoch {epoch + 1}/{EPOCHS}') print('-' * 30) train_loss, train_acc = train_epoch( model, train_loader, criterion, optimizer, device, scheduler ) print(f'Train loss: {train_loss:.4f}, Train accuracy: {train_acc:.4f}') val_loss, val_acc = eval_model( model, val_loader, criterion, device ) print(f'Val loss: {val_loss:.4f}, Val accuracy: {val_acc:.4f}') # 简单的早停策略:保存验证集上性能最好的模型 if val_acc > best_accuracy: torch.save(model.state_dict(), 'best_model_state.bin') best_accuracy = val_acc print(f'>>> Best model saved with accuracy: {best_accuracy:.4f}') print()早停(Early Stopping):我们保存验证集上准确率最高的模型状态。如果连续多个epoch验证集性能不再提升,则可以提前终止训练,避免过拟合。这里实现的是最简单的版本,你可以增加一个patience参数来控制容忍的epoch数。
6. 模型评估与深入分析:不止于准确率
训练完成后,我们在独立的测试集上评估最终模型的性能。准确率是一个直观的指标,但对于分类问题,尤其是各类别样本量不完全相同时,我们需要更细致的分析。
6.1 在测试集上进行最终评估
加载之前保存的最佳模型,并在测试集上运行评估函数。
model.load_state_dict(torch.load('best_model_state.bin')) test_loss, test_acc = eval_model(model, test_loader, criterion, device) print(f'Test set performance: Loss = {test_loss:.4f}, Accuracy = {test_acc:.4f}')6.2 生成分类报告与混淆矩阵
sklearn.metrics提供了丰富的评估工具。
from sklearn.metrics import classification_report, confusion_matrix import seaborn as sns import matplotlib.pyplot as plt def get_predictions(model, data_loader, device): model = model.eval() predictions = [] real_values = [] with torch.no_grad(): for batch in data_loader: input_ids = batch['input_ids'].to(device) attention_mask = batch['attention_mask'].to(device) labels = batch['labels'].to(device) outputs = model(input_ids=input_ids, attention_mask=attention_mask) _, preds = torch.max(outputs, dim=1) predictions.extend(preds.cpu().tolist()) real_values.extend(labels.cpu().tolist()) return predictions, real_values y_pred, y_true = get_predictions(model, test_loader, device) # 打印详细的分类报告(精确率、召回率、F1分数) print(classification_report(y_true, y_pred, target_names=category_names)) # 假设category_names是类别名称列表 # 绘制混淆矩阵 cm = confusion_matrix(y_true, y_pred) plt.figure(figsize=(10, 8)) sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=category_names, yticklabels=category_names) plt.title('Confusion Matrix') plt.ylabel('True Label') plt.xlabel('Predicted Label') plt.tight_layout() plt.show()如何解读结果?
- 准确率(Accuracy):整体分类正确的比例。对于平衡数据集,这是一个好指标。
- 精确率(Precision):对于预测为某一类的样本,有多少是预测正确的。高精确率意味着“宁缺毋滥”。
- 召回率(Recall):对于真实为某一类的样本,有多少被成功预测出来。高召回率意味着“宁可错杀,不可放过”。
- F1分数(F1-Score):精确率和召回率的调和平均数,是综合衡量指标。
通过混淆矩阵,你可以直观地看到模型在哪些类别上容易混淆。例如,“财经”和“房产”新闻可能因为都涉及经济数据而容易分错,“体育”和“娱乐”可能因为都包含人物活动而混淆。这为你后续优化指明了方向:也许是这些类别的特征确实接近,也许你需要针对性地补充训练数据或进行特征工程。
7. 性能优化与实战技巧:从“能用”到“好用”
如果你的模型在测试集上表现已经不错(例如准确率>94%),那么恭喜你。但如果你想进一步提升,或者遇到了显存不足、训练慢的问题,下面这些实战技巧会很有帮助。
7.1 解决显存不足(OOM)问题
这是微调BERT时最常见的问题。bert-base-chinese模型约有1.1亿参数,即使batch_size=32,也可能在显存较小的GPU上告急。
策略一:梯度累积(Gradient Accumulation)其核心思想是:在硬件限制下,使用一个较小的batch_size进行前向传播和损失计算,但不立即更新权重(不执行optimizer.step())。而是连续计算多个小批量的梯度并累加,当累积的步数达到一个设定的“虚拟批量大小”时,再用累积的总梯度更新一次权重。
accumulation_steps = 4 # 虚拟批量大小 = batch_size * accumulation_steps optimizer.zero_grad() # 在累积循环开始前清空梯度 for step, batch in enumerate(train_loader): # ... 前向传播,计算损失 loss = loss / accumulation_steps # 损失按累积步数缩放 loss.backward() # 梯度累积 if (step + 1) % accumulation_steps == 0: optimizer.step() scheduler.step() optimizer.zero_grad() # 更新权重后清空梯度这样,你可以在batch_size=8的情况下,实现等效于batch_size=32的训练效果,同时显存占用大幅降低。
策略二:混合精度训练(Mixed Precision Training)使用torch.cuda.amp模块,让模型的部分计算使用16位浮点数(FP16),减少显存占用并加速计算。
from torch.cuda.amp import autocast, GradScaler scaler = GradScaler() # 梯度缩放,防止FP16下的梯度下溢 for batch in train_loader: optimizer.zero_grad() with autocast(): # 自动混合精度上下文 outputs = model(...) loss = criterion(...) scaler.scale(loss).backward() # 缩放损失,反向传播 scaler.step(optimizer) # 缩放梯度,更新权重 scaler.update() # 更新缩放因子 scheduler.step()混合精度训练通常能节省约50%的显存,并提升训练速度。对于BERT微调,效果非常显著。
7.2 尝试不同的BERT变体与池化策略
bert-base-chinese是一个很好的起点,但你可以尝试更强大或更高效的模型。
bert-large-chinese:参数更多,能力更强,但需要更多显存和计算时间。RoBERTa:去掉了BERT的“下一句预测”任务,使用动态掩码,在更大数据上训练更久,通常效果略优于BERT。Hugging Face上也有中文RoBERTa模型(如hfl/chinese-roberta-wwm-ext)。ALBERT:通过参数共享等技术大幅减少了参数量,训练和推理更快,但有时性能略有牺牲。Electra:使用“生成器-判别器”的预训练任务,效率更高。
此外,除了使用[CLS]token的输出,你还可以尝试其他池化策略:
- 均值池化(Mean Pooling):对序列中所有非填充token的最后一层隐藏状态取平均。
- 最大池化(Max Pooling):取所有token隐藏状态在每个维度上的最大值。
- 注意力池化(Attention Pooling):学习一个注意力权重,对token表示进行加权平均。
这些策略有时能捕捉到更丰富的序列信息,尤其是对于长文本。你可以通过修改模型forward函数中的pooled_output部分来轻松尝试。
7.3 超参数调优与正则化
如果模型在训练集上表现很好,但在验证集上表现不佳(过拟合),可以考虑以下方法:
- 增加Dropout率:将分类层前的Dropout从0.3提高到0.5。
- 减小学习率:尝试
1e-5甚至5e-6。 - 增加权重衰减(Weight Decay):在
AdamW优化器中,weight_decay参数默认是0.01,可以尝试增加到0.1。 - 使用更小的
max_length:如果文本信息冗余,更短的截断长度本身就是一种正则化。 - 标签平滑(Label Smoothing):在计算交叉熵损失时,不直接使用硬标签(0或1),而是使用平滑后的软标签(如0.9和0.1),可以减轻模型对训练数据的过度自信。
调优是一个系统性的实验过程,建议每次只改变一个变量,并在验证集上观察效果。可以使用wandb或TensorBoard等工具来跟踪实验。
8. 模型部署与推理:让模型“跑起来”
训练好的模型最终要用于预测新的数据。我们需要编写一个简洁的推理函数。
def predict_news_category(text, model, tokenizer, device, max_len=256): """预测单条新闻文本的类别""" model.eval() # 预处理输入文本 encoding = tokenizer.encode_plus( text, add_special_tokens=True, max_length=max_len, padding='max_length', truncation=True, return_attention_mask=True, return_tensors='pt', ) input_ids = encoding['input_ids'].to(device) attention_mask = encoding['attention_mask'].to(device) with torch.no_grad(): outputs = model(input_ids=input_ids, attention_mask=attention_mask) _, prediction = torch.max(outputs, dim=1) return prediction.cpu().item() # 返回预测的类别索引 # 示例使用 sample_text = "北京时间今晚,欧冠半决赛第二回合即将打响..." predicted_idx = predict_news_category(sample_text, model, tokenizer, device) print(f"预测类别: {category_names[predicted_idx]}")对于生产环境,你可能需要将模型转换为TorchScript或ONNX格式以提高推理效率,并封装成API服务(如使用FastAPI)。同时,考虑使用模型量化等技术进一步压缩模型大小,满足端侧部署的需求。
回顾整个项目,从数据准备到模型部署,最深的体会是:细节决定成败。Tokenizer的一个参数、学习率的一个数量级、数据预处理时的一个疏忽,都可能导致结果天差地别。这个基于THUCNews和BERT的项目,就像一把精密的瑞士军刀,每一个部件都值得你反复琢磨。它带给你的不仅是一个可运行的分类器,更是一套处理中文NLP任务的完整方法论。当你下次面对其他文本任务时,这套数据流水线、模型微调框架和调优思路,依然会是你最可靠的起点。
本文还有配套的精品资源,点击获取