1. 项目概述
中文新闻文本分类是自然语言处理领域的基础任务之一,在舆情监控、内容推荐、信息检索等场景中有着广泛应用。本文将基于THUCNews数据集,从零开始实现三种不同技术路线的文本分类方案,并进行全面对比分析。
1.1 核心需求解析
我们需要解决的问题是:给定一段中文新闻文本,自动判断其所属的类别(如体育、财经、科技等)。这个任务看似简单,但实际面临以下挑战:
- 中文的语义理解比英文更复杂,存在分词、歧义等问题
- 新闻文本通常较长,需要有效捕捉关键信息
- 不同类别间的边界可能模糊(如科技与数码产品)
针对这些挑战,我们将实现三种典型方案:
- TextCNN:基于卷积神经网络的轻量级方案
- BiLSTM:擅长捕捉长距离依赖的双向LSTM方案
- BERT:当前最先进的预训练语言模型方案
提示:THUCNews数据集包含14个类别约74万条新闻文本,每条数据格式为"标签\t文本"。为简化实验,可以选择其中的10个主要类别。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 数据准备与预处理
2.1 数据集加载
首先我们需要构建一个PyTorch的Dataset类来加载和处理数据:
python复制import torch
from torch.utils.data import Dataset
class NewsDataset(Dataset):
def __init__(self, path, tokenizer=None, max_len=128):
self.samples = []
with open(path, encoding='utf-8') as f:
for line in f:
try:
y, x = line.strip().split('\t')
self.samples.append((int(y), x))
except:
continue # 跳过格式错误的行
self.tokenizer = tokenizer
self.max_len = max_len
def __len__(self):
return len(self.samples)
def __getitem__(self, idx):
y, x = self.samples[idx]
if self.tokenizer:
enc = self.tokenizer(x,
truncation=True,
padding='max_length',
max_length=self.max_len,
return_tensors='pt')
return enc['input_ids'].squeeze(0), enc['attention_mask'].squeeze(0), y
else:
return x, y
这个Dataset类的主要特点:
- 自动处理原始文本文件,过滤格式错误的数据
- 支持两种模式:原始文本模式(用于传统方法)和tokenizer模式(用于BERT等模型)
- 统一返回格式,便于后续处理
2.2 数据划分与批处理
通常我们会将数据划分为训练集、验证集和测试集:
python复制from torch.utils.data import random_split, DataLoader
# 加载完整数据集
full_dataset = NewsDataset('thucnews.txt')
# 按8:1:1划分
train_size = int(0.8 * len(full_dataset))
val_size = int(0.1 * len(full_dataset))
test_size = len(full_dataset) - train_size - val_size
train_set, val_set, test_set = random_split(
full_dataset, [train_size, val_size, test_size])
# 创建DataLoader
batch_size = 32
train_loader = DataLoader(train_set, batch_size=batch_size, shuffle=True)
val_loader = DataLoader(val_set, batch_size=batch_size)
test_loader = DataLoader(test_set, batch_size=batch_size)
注意:实际应用中建议使用固定划分而非随机划分,以确保结果可复现。THUCNews官方提供了标准划分方案。
3. 方案一:TextCNN实现
3.1 模型架构
TextCNN是最早将CNN应用于文本分类的经典模型,其核心思想是通过不同大小的卷积核捕捉n-gram特征。
python复制import torch.nn as nn
import torch.nn.functional as F
class TextCNN(nn.Module):
def __init__(self, vocab_size, embed_dim, num_classes):
super().__init__()
self.embed = nn.Embedding(vocab_size, embed_dim)
self.convs = nn.ModuleList([
nn.Conv2d(1, 100, (k, embed_dim)) for k in [3,4,5]
])
self.fc = nn.Linear(300, num_classes)
self.dropout = nn.Dropout(0.5)
def forward(self, x):
x = self.embed(x) # (B, L, D)
x = x.unsqueeze(1) # (B, 1, L, D)
x = [F.relu(conv(x)).squeeze(3) for conv in self.convs]
x = [F.max_pool1d(i, i.size(2)).squeeze(2) for i in x]
x = torch.cat(x, 1)
x = self.dropout(x)
return self.fc(x)
关键组件解析:
- 嵌入层:将词索引映射为稠密向量
- 多尺寸卷积核:并行使用3、4、5三种窗口大小
- 最大池化:提取每个特征通道的最显著特征
- Dropout:防止过拟合
3.2 训练与评估
TextCNN的训练相对简单:
python复制from torch.optim import Adam
# 初始化模型
vocab_size = 50000 # 根据实际词典大小调整
embed_dim = 300
num_classes = 10
model = TextCNN(vocab_size, embed_dim, num_classes).to(device)
# 优化器
optimizer = Adam(model.parameters(), lr=1e-3)
# 训练循环
for epoch in range(10):
train(model, train_loader, optimizer, device)
acc = evaluate(model, val_loader, device)
print(f'Epoch {epoch}, Val Acc: {acc:.4f}')
实测中,TextCNN在THUCNews上可以达到约90%的准确率,训练速度非常快(约5分钟/epoch)。
4. 方案二:BiLSTM实现
4.1 模型架构
BiLSTM通过双向LSTM捕捉文本的前后文依赖关系:
python复制class BiLSTM(nn.Module):
def __init__(self, vocab_size, embed_dim, hidden_dim, num_classes):
super().__init__()
self.embed = nn.Embedding(vocab_size, embed_dim)
self.lstm = nn.LSTM(embed_dim, hidden_dim,
batch_first=True,
bidirectional=True,
dropout=0.3)
self.fc = nn.Linear(hidden_dim*2, num_classes)
self.dropout = nn.Dropout(0.5)
def forward(self, x):
x = self.embed(x)
x, _ = self.lstm(x)
x = self.dropout(x)
# 取最后时刻的隐藏状态
x = x[:, -1, :]
return self.fc(x)
改进点:
- 双向LSTM:同时考虑前后文信息
- 序列输出:使用全部时间步的输出而非仅最后状态
- 增加Dropout:提升泛化能力
4.2 训练技巧
BiLSTM的训练需要注意以下几点:
python复制# 梯度裁剪防止爆炸
optimizer = Adam(model.parameters(), lr=1e-3)
max_grad_norm = 5.0
for batch in train_loader:
optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_grad_norm)
optimizer.step()
BiLSTM通常需要更长的时间训练(约15分钟/epoch),但准确率可以提升到92-93%。
5. 方案三:BERT实现
5.1 模型架构
BERT是目前最先进的预训练语言模型:
python复制from transformers import BertTokenizer, BertModel
class BertClassifier(nn.Module):
def __init__(self, num_classes):
super().__init__()
self.bert = BertModel.from_pretrained('bert-base-chinese')
self.fc = nn.Linear(768, num_classes)
self.dropout = nn.Dropout(0.1)
def forward(self, input_ids, attention_mask):
out = self.bert(input_ids=input_ids,
attention_mask=attention_mask)
cls = out.last_hidden_state[:, 0] # [CLS] token
cls = self.dropout(cls)
return self.fc(cls)
关键点:
- 使用预训练的BERT-base-chinese模型
- 取[CLS]位置的输出作为文本表示
- 较小的Dropout率(与原始BERT一致)
5.2 微调策略
BERT训练需要特殊处理:
python复制# 使用不同的学习率
no_decay = ['bias', 'LayerNorm.weight']
optimizer_grouped_parameters = [
{'params': [p for n, p in model.named_parameters()
if not any(nd in n for nd in no_decay)],
'weight_decay': 0.01},
{'params': [p for n, p in model.named_parameters()
if any(nd in n for nd in no_decay)],
'weight_decay': 0.0}
]
optimizer = AdamW(optimizer_grouped_parameters, lr=2e-5)
# 线性学习率预热
from transformers import get_linear_schedule_with_warmup
total_steps = len(train_loader) * epochs
scheduler = get_linear_schedule_with_warmup(
optimizer, num_warmup_steps=0, num_training_steps=total_steps)
BERT训练较慢(约30分钟/epoch),但准确率可达95%以上。
6. 方案对比与选型建议
6.1 全面对比
| 指标 | TextCNN | BiLSTM | BERT |
|---|---|---|---|
| 准确率 | 90% | 93% | 96% |
| 训练速度 | ⭐⭐⭐⭐ | ⭐⭐⭐ | ⭐⭐ |
| 推理速度 | ⭐⭐⭐⭐ | ⭐⭐⭐ | ⭐⭐ |
| 实现难度 | ⭐⭐ | ⭐⭐⭐ | ⭐⭐⭐⭐ |
| 数据需求 | ⭐⭐ | ⭐⭐ | ⭐⭐⭐⭐ |
| 可解释性 | ⭐⭐⭐ | ⭐⭐ | ⭐ |
6.2 选型指南
根据实际需求选择合适方案:
-
快速原型开发:选择TextCNN
- 代码简单,训练快速
- 适合课程项目或验证想法
-
平衡性能与效率:选择BiLSTM
- 比CNN更好的序列建模能力
- 适合生产环境中的实时系统
-
追求最佳效果:选择BERT
- 需要GPU资源
- 适合学术研究或关键业务场景
实操建议:可以先从TextCNN开始建立baseline,再逐步尝试更复杂的模型。实际部署时可以考虑模型蒸馏等技术,将BERT的知识迁移到小模型上。
7. 常见问题与解决方案
7.1 内存不足问题
问题现象:训练BERT时出现OOM错误
解决方案:
- 减小batch size(可小至8或4)
- 使用梯度累积:
python复制accumulation_steps = 4 for i, batch in enumerate(train_loader): loss.backward() if (i+1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad() - 尝试混合精度训练:
python复制from torch.cuda.amp import GradScaler, autocast scaler = GradScaler() with autocast(): outputs = model(inputs) loss = criterion(outputs, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()
7.2 类别不平衡问题
问题现象:某些类别准确率明显偏低
解决方案:
- 在DataLoader中设置sampler:
python复制from torch.utils.data import WeightedRandomSampler weights = 1.0 / torch.bincount(tensor_of_labels) sampler = WeightedRandomSampler(weights, len(weights)) - 使用带权重的损失函数:
python复制
class_weights = torch.FloatTensor([...]).to(device) criterion = nn.CrossEntropyLoss(weight=class_weights)
7.3 过拟合问题
问题现象:训练集准确率高但验证集不提升
解决方案:
- 增加Dropout比例
- 添加L2正则化:
python复制optimizer = Adam(model.parameters(), lr=1e-3, weight_decay=1e-4) - 使用早停机制(Early Stopping)
8. 进阶优化方向
对于希望进一步提升性能的开发者,可以考虑以下方向:
-
模型融合:
- 将多个模型的预测结果进行投票或平均
- 使用Stacking等元学习技术
-
数据增强:
- 同义词替换
- 回译(中→英→中)
- EDA(Easy Data Augmentation)
-
预训练微调:
- 尝试其他预训练模型(RoBERTa、ALBERT等)
- 领域自适应预训练(继续在新闻语料上预训练)
-
模型压缩:
- 知识蒸馏(用BERT训练小模型)
- 量化(FP16/INT8)
- 剪枝
我在实际项目中发现,简单的模型融合(TextCNN+BiLSTM)往往能带来1-2个百分点的提升,而计算成本增加不多。另外,针对中文文本,适当增加字级别的特征有时也能改善效果。
