AI人工智能--文本分类任务-第七周(小白)
一、文本分类基础认知
1.1 核心定义
文本分类是预先设定文本类别集合,对单篇文本预测其所属类别的 NLP 任务。核心逻辑是建立 “文本内容 - 类别标签” 的映射关系,本质是通过算法学习文本特征与类别间的关联模式。
1.2 典型案例
| 分类类型 | 文本示例 | 对应类别 | 应用场景 |
|---|---|---|---|
| 情感分析 | “这家饭店太难吃了” | 负类 | 电商评论情感判断、舆情情绪分析 |
| 情感分析 | “这家菜很好吃” | 正类 | 同上 |
| 领域分类 | “今日 A 股行情大好” | 经济 | 新闻频道划分、资讯推荐领域定位 |
| 领域分类 | “今日湖人击败勇士” | 体育 | 同上 |
二、文本分类核心应用场景
文本分类广泛渗透于信息处理各环节,核心场景可分为以下四类:
2.1 资讯内容打标签
-
核心目标:为新闻、文章等资讯分配多级类别标签,实现内容结构化。
-
案例:新浪网的分类体系(一级分类:新闻、体育、财经、娱乐等;二级分类:体育下的 NBA、英超、中超,财经下的股票、基金、外汇),用户可通过分类快速定位感兴趣内容。
-
价值:提升资讯检索效率,支撑个性化推荐(如向体育爱好者推送 NBA 新闻)。
2.2 电商评论分析
-
核心目标:分析用户对商品的评价态度,辅助商家优化产品、平台提升用户体验。
-
案例:电商平台中用户对商品的评价(“还不错哦”“能用么!”),通过分类判断评价为 “正面”“中性”“负面”,统计商品的好评率、差评核心原因(如 “质量差”“尺码不准”)。
-
价值:商家可针对性改进产品(如某款鞋子差评多为 “磨脚”,则优化鞋型),平台为用户优先推荐高好评商品。
2.3 违规内容检测
-
核心目标:识别文本中的违规信息,维护内容生态安全。
-
检测范围:涉黄、涉暴、涉恐、辱骂、虚假宣传等违规内容。
-
应用场景:
-
客服 / 销售对话质检:检测客服是否使用辱骂话术、销售是否夸大产品功效。
-
网站内容审查:过滤论坛、评论区的违规言论,避免不良信息传播。
-
三、自定义类别任务特性
自定义类别任务是文本分类的灵活延伸,核心特点的 “类别定义自由度高”,只要人能通过文本判断的类别,均可作为分类目标。常见场景包括:
-
垃圾邮件分类:类别为 “垃圾邮件”“正常邮件”,通过识别 “广告推销”“诈骗链接” 等特征判断。
-
汽车交易相关性判断:类别为 “与汽车交易相关”“与汽车交易无关”,用于筛选汽车电商平台的有效对话 / 文章。
-
作者风格识别:类别为 “某作者风格”“非某作者风格”,如判断一篇散文是否为鲁迅所作(基于用词、句式特征)。
-
机器生成文本检测:类别为 “机器生成”“人工撰写”,用于识别 AI 生成的新闻、论文等内容。
-
合同文本合规性判断:类别为 “符合规范”“不符合规范”,检测合同中是否存在无效条款、法律风险表述。
-
阅读人群适配性分类:类别为 “未成年适宜”“中年适宜”“老年适宜”“孕妇适宜” 等,用于儿童读物、老年健康文章的精准推送。
四、文本分类的机器学习流程
机器学习实现文本分类需遵循 “定义 - 数据 - 训练 - 预测” 的标准化流程,具体步骤如下:
4.1 流程拆解
-
定义类别:明确分类目标与类别集合(如情感分析定义 “正类”“负类”,领域分类定义 “经济”“体育”“科技” 等)。
-
收集数据:获取带类别标签的文本数据(标注数据),如情感分析需收集大量标注 “正 / 负” 的评论,领域分类需收集标注 “经济 / 体育 / 科技” 的新闻。数据质量直接影响模型效果,需保证标注准确性、类别分布合理性。
-
模型训练:将标注数据输入分类模型,模型学习文本特征与类别间的映射关系。核心逻辑是:模型通过计算 “文本特征属于某类别的概率”,调整参数使预测结果与真实标签尽可能一致。
-
预测应用:将训练好的模型用于未标注文本,输出该文本所属的类别(如输入 “今日油价上涨”,模型预测其类别为 “经济”)。
4.2 流程示意图

五、贝叶斯算法在文本分类中的应用
贝叶斯算法是基于 “贝叶斯公式” 的概率模型,核心思想是通过 “先验概率” 与 “条件概率” 计算 “后验概率”,实现类别预测。
5.1 预备知识:全概率公式
-
公式定义:若事件组{Bi}是样本空间Ω的一个划分(即Bi互斥且⋃Bi=Ω),且P(Bi)>0,则对任意事件A,有:
-
案例解释:扔正常骰子,计算 “结果为 5(事件 A)” 的概率。
-
划分事件:
(结果为奇数,
)、B2(结果为偶数,
)。
-
条件概率:
(奇数包含 1、3、5,共 3 种,5 占 1 种),
(偶数不含 5)。
-
计算结果:
与实际常识一致。
-
5.2 核心公式:贝叶斯公式
- 公式推导:由联合概率
,变形得:
- 符号含义:
-
(P(A)):事件 A 的先验概率(已知的初始概率,如 “感染新冠的概率”)。
-
(P(B)):事件 B 的先验概率(如 “核酸检测呈阳性的概率”)。
-
(P(B|A)):事件 A 发生时事件 B 发生的条件概率(如 “感染新冠后检测呈阳性的概率”)。
-
(P(A|B)):事件 B 发生时事件 A 发生的后验概率(如 “检测呈阳性时实际感染新冠的概率”,即我们需要求解的目标)。
-
5.3 贝叶斯公式的实际应用:核酸检测概率计算
5.3.1 已知条件
- 新冠人群感染率(先验概率(P(A))):0.1%(即 0.001)。
- 检测准确率:
实际情况 \ 检测结果 检测呈阳性 检测呈阴性 感染新冠(A) 99% 1% 未感染新冠 5% 95%
5.3.2 计算过程
-
计算\(P(B)\)(检测呈阳性的总概率,用全概率公式):
-
计算\(P(A|B)\)(检测呈阳性时实际感染的概率):
(即 1.9%)
5.3.3 结论
即使核酸检测呈阳性,实际感染新冠的概率仅约 1.9%,原因是人群感染率(先验概率)极低,误报率(5%)导致大量未感染者被误判为阳性,需结合临床症状进一步判断。
5.4 贝叶斯算法在文本分类中的应用
5.4.1 核心假设
文本属于某类别的概率,仅与文本中包含的词相关(即 “词的独立性假设”,简化计算)。
5.4.2 分类逻辑
假设存在 3 个类别,文本S由
(n 个词)组成,目标是计算
(文本S属于类别
的概率),选择概率最大的类别作为预测结果。
-
应用贝叶斯公式:
其中,P(S)是所有类别共有的分母,比较不同\
的概率时可忽略,只需计算
。
-
词的独立性假设:文本S在类别
下的概率
,等于每个词在\(A_i\)下概率的乘积:\(P(S|A_i)=P(W_1|A_i)×P(W_2|A_i)×...×P(W_n|A_i)\)
-
概率计算示例:若判断文本 “今日 A 股上涨” 是否属于 “经济” 类(\(A_1\)):
:“经济” 类文本在所有文本中的占比(先验概率)。
:“今日” 在 “经济” 类文本中出现的概率,
:“A 股” 在 “经济” 类文本中出现的概率,
:“上涨” 在 “经济” 类文本中出现的概率。
- 计算
,并与 “体育”“科技” 等类别的对应值比较,最大者即为预测类别。
5.5 贝叶斯算法的优缺点
5.5.1 优点
-
简单高效:计算逻辑清晰,无需复杂迭代训练,适合小规模数据场景。
-
可解释性强:概率计算过程可追溯,能明确知道每个词对类别预测的贡献。
-
样本覆盖好时效果优:若训练数据能充分覆盖各类别特征,预测精度较高。
-
支持分批训练:可将训练数据分批输入,无需一次性加载所有数据,降低内存压力。
5.5.2 缺点
-
样本不均衡敏感:若某类样本数量极少,其先验概率\(P(A_i)\)被低估,导致预测偏向多样本类别。
-
未见过特征处理困难:若文本中出现训练数据未见过的词,该词的条件概率\(P(W|A_i)=0\),导致整体概率为 0,需通过 “平滑技术”(如拉普拉斯平滑)解决。
-
特征独立假设不成立:实际文本中词与词存在关联(如 “A 股” 与 “上涨” 常同时出现),独立假设会损失关联信息,影响精度。
-
忽略语序与词义:仅考虑词的出现频率,不考虑词的顺序(如 “我喜欢他” 与 “他喜欢我” 语义不同但词相同),也无法区分多义词(如 “苹果” 指水果或公司)。
六、支持向量机(SVM)在文本分类中的应用
支持向量机(SVM)是一种有监督学习模型,核心思想是寻找 “最大边距超平面”,实现对数据的最优分类,适用于线性可分与线性不可分场景。
6.1 核心概念:最大边距超平面
6.1.1 超平面定义
在二维空间中,超平面是直线(如);在三维空间中是平面;在高维空间中是维度为 “空间维度 - 1” 的线性结构,用于划分不同类别的数据。
6.1.2 边距与最大边距
-
边距:超平面到两类数据中 “最近样本点” 的距离之和。
-
最大边距:SVM 的目标是找到边距最大的超平面,原因是:边距越大,模型对噪声数据的容错能力越强,泛化性能(对新数据的预测能力)越好。
6.1.3 支持向量
两类数据中距离超平面最近的样本点,称为 “支持向量”。SVM 的超平面仅由支持向量决定,其他样本点对超平面位置无影响 —— 这是 SVM 的核心特性,也是其对异常值不敏感的原因。
6.2 线性不可分问题的解决:核函数
6.2.1 问题本质
当数据在原始输入空间中无法用直线(或低维超平面)划分时(如一维数据[-1,0,1]为正样本,([-3,-2,2,3])为负样本),需将数据映射到更高维度的特征空间,使其在高维空间中线性可分。
6.2.2 映射示例
-
原始一维数据x,映射到二维特征空间
:正样本[-1,0,1]映射后为[(-1,1),(0,0),(1,1)],负样本[-3,-2,2,3]映射后为[(-3,9),(-2,4),(2,4),(3,9)]。此时在二维空间中,可通过直线(如y=2)轻松划分两类数据。
6.2.3 核函数的作用
直接将数据映射到高维空间会导致 “维度灾难”—— 计算量随维度增加呈指数级增长(如 3 维数据映射到 9 维,10 维数据映射到 1024 维)。核函数的核心价值是:无需显式进行高维映射,直接在原始空间中计算 “高维空间中两个向量的内积”,大幅降低计算成本。
6.2.4 常见核函数
| 核函数类型 | 公式 | 适用场景 |
|---|---|---|
| 线性核函数 | 数据本身线性可分,或特征维度高(如文本分类的词袋特征) | |
| 多项式核函数 | 数据呈多项式分布,需捕捉非线性关系 | |
| 高斯核函数(RBF) | 数据分布复杂,无法确定非线性关系类型,应用最广泛 | |
| 双曲正切核函数 | 模拟神经网络,适用于需非线性映射且希望输出范围有限的场景 |
6.3 SVM 的多分类解决方案
SVM 本质是二分类模型,处理多分类(K 类)需通过 “拆解策略” 实现,核心有两种方式:
6.3.1 One vs One(一对一)
-
原理:为每对类别构建一个 SVM 分类器,共需构建
个分类器(如 3 类需 3 个分类器:A-B、A-C、B-C)。
-
预测逻辑:将待预测样本输入所有分类器,统计每个类别被预测的次数,次数最多的类别即为最终结果。
-
示例:类别 [A,B,C],样本 X 输入 3 个分类器:
-
SVM (A,B)→A,SVM (A,C)→A,SVM (B,C)→B → A 被预测 2 次,B 被预测 1 次 → 最终类别为 A。
-
6.3.2 One vs Rest(一对多)
-
原理:为每个类别构建一个 SVM 分类器,该分类器将 “当前类别” 视为正类,“其他所有类别” 视为负类,共需构建 K 个分类器(如 3 类需 3 个分类器:A - 其他、B - 其他、C - 其他)。
-
预测逻辑:将待预测样本输入所有分类器,每个分类器输出样本属于 “当前正类” 的概率(或得分),选择得分最高的类别作为最终结果。
-
示例:类别 [A,B,C],样本 X 输入 3 个分类器:
-
SVM (A - 其他)→0.1,SVM (B - 其他)→0.2,SVM (C - 其他)→0.5 → C 得分最高 → 最终类别为 C。
-
6.4 SVM 的优缺点
6.4.1 优点
-
对异常值不敏感:超平面仅由支持向量决定,少量异常值(远离支持向量的样本)不影响超平面位置。
-
样本需求量低:无需大量数据即可学习到稳定的超平面,适合小样本场景。
-
高维数据处理能力强:即使特征维度(如文本的词袋特征维度)远大于样本数量,仍能有效分类,避免过拟合。
6.4.2 缺点
-
大规模数据计算负担重:样本数量过多时,支持向量数量增加,模型训练与预测的时间、空间复杂度显著上升。
-
多分类处理复杂:需通过 One vs One 或 One vs Rest 拆解任务,增加了模型构建与调参的复杂度。
-
核函数与参数选择困难:不同核函数适用于不同数据分布,需通过大量实验尝试(如高斯核的\(\sigma\)参数),无统一选择标准。
七、深度学习在文本分类中的应用
深度学习通过构建多层神经网络,自动学习文本的深层语义特征,解决了传统机器学习(如贝叶斯、SVM)依赖人工特征工程的问题,是当前文本分类的主流方法。
7.1 深度学习文本分类的通用 Pipeline
所有深度学习模型的实现均遵循以下工程流程,通过不同模块的协作完成训练与预测:
| 模块文件 | 核心功能 | 关键操作 |
|---|---|---|
| config.py | 配置模型参数 | 设定学习率、 batch size、训练轮数(epoch)、隐藏层维度等 |
| loader.py | 数据加载与预处理 | 读取文本数据→文本清洗(去特殊符号、停用词)→文本编码(如词嵌入)→划分训练集 / 验证集 / 测试集 |
| model.py | 定义神经网络结构 | 构建模型的层结构(如 Embedding 层、LSTM 层、卷积层),定义前向传播逻辑 |
| evaluate.py | 模型评价 | 定义评价指标(准确率 acc、精确率 precision、召回率 recall、F1-score),每轮训练后计算验证集指标,判断模型性能 |
| main.py | 训练主流程 | 初始化模型、优化器、损失函数→迭代训练(输入训练数据→计算损失→反向传播更新参数)→保存最优模型→测试集评估 |
7.2 主流深度学习模型详解
7.2.1 fastText
-
核心特点:基于 “n-gram 特征” 捕捉文本的局部序列信息,模型结构简单、训练速度快。
-
n-gram 特征:将单个词拆分为多个子序列(如 “apple” 的 3-gram 为 “app”“ppl”“ple”),弥补了传统词袋模型忽略词内字符关系的缺陷。
-
模型结构:文本→n-gram 特征提取→Embedding 层(将每个 n-gram 映射为向量)→Mean Pooling(对所有 n-gram 向量取平均,得到文本向量)→全连接层→类别预测。
-
适用场景:大规模文本分类(如新闻分类、垃圾邮件检测),对速度要求高的场景。

模型参数
| 参数名 | 含义 | 调优建议 |
|---|---|---|
vocab_size |
n-gram 词汇表大小(需提前统计数据中所有 ngram 的数量) | 数据量小:1 万~5 万- 数据量大:10 万~50 万(配合minCount过滤后的数据) |
embed_dim |
Embedding 维度(同官方库dim) |
同官方库建议:50~300,复杂任务取上限 |
num_classes |
分类类别数(必须与数据一致) | 按实际任务指定(如二分类设 2,10 分类设 10) |
ngram_size |
n-gram 的 n 值(同官方库wordNgrams) |
同官方库建议:英文 2~3,中文 1~ |
代码实现
import torch
import torch.nn as nn
import torch.nn.functional as F
class FastText(nn.Module):
def __init__(self, vocab_size, embed_dim, num_classes, ngram_size=2):
super(FastText, self).__init__()
self.embedding = nn.Embedding(vocab_size, embed_dim) # n-gram的Embedding层
self.fc = nn.Linear(embed_dim, num_classes) # 分类层
def forward(self, x):
# x: [batch_size, seq_len] (seq_len是n-gram的数量)
embed = self.embedding(x) # [batch_size, seq_len, embed_dim]
mean_pool = torch.mean(embed, dim=1) # 均值池化:[batch_size, embed_dim]
out = self.fc(mean_pool) # [batch_size, num_classes]
return out
# 示例:假设n-gram词汇表大小为1000,Embedding维度100,分类数2
model = FastText(vocab_size=1000, embed_dim=100, num_classes=2)
x = torch.randint(0, 1000, (32, 10)) # 32个样本,每个样本10个n-gram特征
output = model(x)
print(output.shape) # 输出 torch.Size([32, 2])(32个样本的2分类结果)
7.2.2 TextRNN(基于 RNN/LSTM/GRU)
-
核心思想:利用循环神经网络(RNN)的 “时序记忆能力”,按词的顺序处理文本,捕捉文本的上下文语义关系。
-
改进:LSTM/GRU:传统 RNN 存在 “梯度消失” 问题(无法捕捉长文本的远距离依赖),LSTM(长短期记忆网络)通过 “门控机制”(输入门、遗忘门、输出门)控制信息的遗忘与更新,GRU 是 LSTM 的简化版,性能相近但计算更快。
-
双向 LSTM(BiLSTM):同时从 “左到右” 和 “右到左” 处理文本,捕捉双向上下文(如 “我不喜欢这部电影”,正向处理到 “不” 时,结合反向的 “喜欢”“电影” 更易判断情感为负)。
-
模型结构:文本→Embedding 层→BiLSTM 层(提取双向时序特征)→Dropout 层(防止过拟合)→LSTM 层(进一步提取深层时序特征)→全连接层→Softmax→类别预测。
-
适用场景:需捕捉上下文依赖的文本分类(如情感分析、对话意图识别)。

代码示例
import torch
import torch.nn as nn
import torch.nn.functional as F
class TextRNN(nn.Module):
def __init__(self, vocab_size, embed_dim, hidden_dim, num_classes, num_layers=2, dropout=0.5):
super(TextRNN, self).__init__()
# 1. Embedding层:将词索引转成词向量
self.embedding = nn.Embedding(vocab_size, embed_dim)
# 2. BiLSTM层:双向LSTM捕获上下文
self.bilstm = nn.LSTM(
input_size=embed_dim,
hidden_size=hidden_dim,
num_layers=num_layers,
bidirectional=True, # 双向
batch_first=True # 输入格式:[batch, seq_len, embed_dim]
)
# 3. Dropout层
self.dropout = nn.Dropout(dropout)
# 4. 单向LSTM层(截图中BiLSTM后接LSTM,也可简化为全连接)
self.lstm = nn.LSTM(
input_size=hidden_dim*2, # 双向LSTM输出是2*hidden_dim
hidden_size=hidden_dim,
num_layers=1,
batch_first=True
)
# 5. 全连接层(Linear):映射到分类数
self.fc = nn.Linear(hidden_dim, num_classes)
# 6. Softmax(PyTorch中交叉熵损失会自带Softmax,可省略)
def forward(self, x):
# x形状:[batch_size, seq_len](词索引序列)
# 1. Embedding:[batch_size, seq_len] → [batch_size, seq_len, embed_dim]
embed = self.embedding(x)
# 2. BiLSTM:输出[batch_size, seq_len, 2*hidden_dim]
bilstm_out, _ = self.bilstm(embed)
# 3. Dropout
drop_out = self.dropout(bilstm_out)
# 4. LSTM:输出[batch_size, seq_len, hidden_dim]
lstm_out, _ = self.lstm(drop_out)
# 取最后一个时间步的输出(截图中用“最后一个位置输出向量分类”)
last_out = lstm_out[:, -1, :] # [batch_size, hidden_dim]
# 5. 全连接层:[batch_size, hidden_dim] → [batch_size, num_classes]
logits = self.fc(last_out)
# 6. 输出(若用交叉熵损失,不需要Softmax;若单独预测,加F.softmax(logits, dim=1))
return logits
7.2.3 TextCNN(基于卷积神经网络)
-
核心思想:借鉴图像 CNN 的 “局部特征提取” 能力,用一维卷积核(如 3-gram、4-gram 卷积核)提取文本的局部语义块(如短语特征 “非常好”“太难吃”)。
-
模型结构:文本→Embedding 层(得到文本矩阵:n 个词 ×embedding 维度)→多尺寸卷积层(用不同尺寸的卷积核提取不同长度的局部特征)→Max-over-time Pooling(对每个卷积核的输出取最大值,保留最显著的局部特征)→全连接层(Dropout 防止过拟合)→Softmax→类别预测。
-
优势:并行计算能力强(卷积操作可并行),训练速度快,能有效捕捉短语级特征。
-
适用场景:短文本分类(如电商评论、微博情感分析),需突出局部关键特征的场景。

代码
import torch
import torch.nn as nn
import torch.nn.functional as F
class TextCNN(nn.Module):
def __init__(self, vocab_size, embed_dim, num_classes, filter_sizes=[2,3,4], num_filters=100, dropout=0.5):
super(TextCNN, self).__init__()
# 1. 词嵌入层
self.embedding = nn.Embedding(vocab_size, embed_dim)
# 2. 卷积层(不同卷积核宽度)
self.convs = nn.ModuleList([
nn.Conv2d(1, num_filters, (fs, embed_dim)) # 输入通道1,输出通道num_filters,卷积核(fs, embed_dim)
for fs in filter_sizes
])
# 3. Dropout层
self.dropout = nn.Dropout(dropout)
# 4. 全连接层(分类)
self.fc = nn.Linear(num_filters * len(filter_sizes), num_classes)
def forward(self, x):
# x形状:[batch_size, seq_len]
x = self.embedding(x) # [batch_size, seq_len, embed_dim]
x = x.unsqueeze(1) # 增加通道维度:[batch_size, 1, seq_len, embed_dim]
# 卷积+ReLU+池化
conv_outs = []
for conv in self.convs:
out = F.relu(conv(x)) # [batch_size, num_filters, seq_len - fs + 1, 1]
out = out.squeeze(3) # 去掉最后一维:[batch_size, num_filters, seq_len - fs + 1]
out = F.max_pool1d(out, out.size(2)) # Max-over-time Pooling:[batch_size, num_filters, 1]
out = out.squeeze(2) # 去掉最后一维:[batch_size, num_filters]
conv_outs.append(out)
# 拼接所有卷积核的输出
x = torch.cat(conv_outs, 1) # [batch_size, num_filters * len(filter_sizes)]
x = self.dropout(x)
x = self.fc(x) # [batch_size, num_classes]
return x
7.2.4 Gated CNN(门控卷积神经网络)
-
核心改进:在传统 CNN 的基础上增加 “门控机制”,动态控制卷积特征的权重,增强模型对重要特征的关注。
-
门控逻辑:文本矩阵→两个并行卷积层(A 层、B 层)→B 层输出经过 Sigmoid 激活函数(得到 0-1 的权重)→A 层输出与 B 层权重逐元素相乘(重要特征权重接近 1,不重要特征权重接近 0)→得到门控后的特征。
-
优势:相比传统 CNN,能更灵活地筛选有效特征,提升复杂文本的分类精度。

代码
import torch
import torch.nn as nn
import torch.nn.functional as F
class GatedCNNLayer(nn.Module):
def __init__(self, embed_dim, out_channels, kernel_size):
super().__init__()
# 卷积A(特征提取)
self.conv_A = nn.Conv1d(
in_channels=embed_dim, # 输入通道数=词向量维度
out_channels=out_channels,
kernel_size=kernel_size,
padding=kernel_size//2 # 保持序列长度不变
)
# 卷积B(门控权重)
self.conv_B = nn.Conv1d(
in_channels=embed_dim,
out_channels=out_channels,
kernel_size=kernel_size,
padding=kernel_size//2
)
# 偏置项(对应公式中的b、c)
self.b = nn.Parameter(torch.zeros(out_channels))
self.c = nn.Parameter(torch.zeros(out_channels))
def forward(self, x):
# x形状:[batch_size, seq_len, embed_dim] → 转置为Conv1d要求的[batch, embed_dim, seq_len]
x = x.transpose(1, 2)
# 卷积A + 偏置
A = self.conv_A(x) + self.b.unsqueeze(1) # [batch, out_channels, seq_len]
# 卷积B + 偏置 → Sigmoid激活(门控)
B = torch.sigmoid(self.conv_B(x) + self.c.unsqueeze(1))
# 门控操作:逐元素相乘
gated_out = A * B
# 转置回[batch, seq_len, out_channels]
return gated_out.transpose(1, 2)
# 文本分类完整模型示例
class GatedCNNTextClassifier(nn.Module):
def __init__(self, vocab_size, embed_dim, num_classes, out_channels=128, kernel_size=3):
super().__init__()
self.embedding = nn.Embedding(vocab_size, embed_dim) # 词嵌入层
self.gated_cnn = GatedCNNLayer(embed_dim, out_channels, kernel_size) # 门控卷积层
self.fc = nn.Linear(out_channels, num_classes) # 分类层
def forward(self, x):
# x: [batch_size, seq_len](输入文本的词索引)
embed = self.embedding(x) # [batch, seq_len, embed_dim]
gated_feat = self.gated_cnn(embed) # [batch, seq_len, out_channels]
# 全局平均池化(或最大池化)
pooled = gated_feat.mean(dim=1) # [batch, out_channels]
logits = self.fc(pooled) # [batch, num_classes]
return logits
7.2.5 TextRCNN(RCNN 融合模型)
-
核心思想:融合 RNN 的 “时序特征” 与 CNN 的 “局部特征”,同时捕捉文本的上下文依赖与短语级特征。
-
模型结构:文本→Embedding 层→BiLSTM 层(提取双向上下文特征,得到每个词的上下文向量)→将 “词嵌入向量 + 双向 LSTM 输出向量” 拼接→卷积层(提取拼接后的局部特征)→Max Pooling→全连接层→类别预测。
-
优势:结合了 RNN 和 CNN 的优点,在长文本与复杂语义场景中表现更优。

代码
import torch
import torch.nn as nn
import torch.nn.functional as F
class TextRCNN(nn.Module):
def __init__(self, vocab_size, embed_dim, hidden_size, num_classes, num_layers=1, dropout=0.5):
super(TextRCNN, self).__init__()
# 1. Embedding层
self.embedding = nn.Embedding(vocab_size, embed_dim)
# 2. 双向LSTM
self.bilstm = nn.LSTM(
input_size=embed_dim,
hidden_size=hidden_size,
num_layers=num_layers,
bidirectional=True, # 双向
batch_first=True
)
# 3. TextCNN(卷积+池化)
self.cnn = nn.Conv1d(
in_channels=embed_dim + 2*hidden_size, # 输入维度:词向量+双向LSTM输出
out_channels=hidden_size,
kernel_size=3, # 卷积核大小
padding=1
)
self.pool = nn.AdaptiveMaxPool1d(1) # 全局最大池化
# 4. 分类层
self.fc = nn.Linear(hidden_size, num_classes)
self.dropout = nn.Dropout(dropout)
def forward(self, x):
# x: [batch_size, seq_len]
embed = self.embedding(x) # [batch_size, seq_len, embed_dim]
# 双向LSTM:output [batch_size, seq_len, 2*hidden_size]
lstm_out, _ = self.bilstm(embed)
# 拼接:词向量 + LSTM输出 → [batch_size, seq_len, embed_dim + 2*hidden_size]
concat = torch.cat([embed, lstm_out], dim=-1)
# 转置适配CNN输入:[batch_size, embed_dim+2*hidden_size, seq_len]
concat = concat.permute(0, 2, 1)
# CNN+池化
cnn_out = F.relu(self.cnn(concat)) # [batch_size, hidden_size, seq_len]
pool_out = self.pool(cnn_out).squeeze(-1) # [batch_size, hidden_size]
# 分类
out = self.fc(self.dropout(pool_out)) # [batch_size, num_classes]
return out
7.2.6 BERT(预训练语言模型)
-
核心突破:基于 “预训练 - 微调” 模式,在大规模无标注文本上预训练(学习通用语言知识),再在具体分类任务上微调(适配任务特征),大幅提升文本分类的精度。
-
核心机制:
-
双向 Transformer encoder:同时关注文本中所有词的相互关系(如 “他喜欢苹果,因为它很甜”,“它” 指代 “苹果”),解决了传统模型单向处理的局限。
-
[CLS] token:在文本开头添加特殊 token [CLS],其对应的输出向量作为整个文本的语义表示,用于分类。
-
-
BERT 特征使用方式:
-
直接使用 [CLS] token 的向量作为文本特征,输入全连接层分类。
-
对所有 token 的向量取 Max/Average Pooling,得到文本向量。
-
将 BERT 输出的向量输入 LSTM/CNN,进一步提取特征。
-
融合 BERT 中间层的输出(如取第 3、5、7 层的向量拼接),增强特征的丰富性。
-
-
适用场景:对精度要求高的复杂任务(如法律文本分类、医疗文本诊断、细粒度情感分析)。
八、文本分类的常见问题与解决方案
在文本分类实践中,常面临数据稀疏、标签不均衡、多标签分类等问题,需针对性解决:
8.1 数据稀疏问题
8.1.1 问题定义
训练数据量过少(如某类别仅几十条样本),模型在训练集上可收敛(拟合训练数据),但在测试集上预测准确率极低(泛化能力差)。
8.1.2 解决方案
-
标注更多数据:最直接有效的方法,通过人工标注或半监督标注(如主动学习选择高价值样本标注)增加数据量。
-
数据增强:构造相似样本,如文本同义替换(“很好”→“非常好”)、随机插入 / 删除停用词、句子语序调整(不改变语义)、翻译回译(中文→英文→中文)。
-
使用预训练模型:预训练模型(如 BERT)已在大规模数据上学习了通用语言知识,微调时需少量任务数据即可达到较好效果,减少对标注数据的依赖。
-
增加规则弥补:对数据稀疏的类别,制定人工规则(如 “包含‘涨停’‘跌停’的文本属于金融类”),辅助模型预测。
-
调整阈值:在二分类中,降低少数类的预测阈值(如将 “正类” 预测概率阈值从 0.5 调整为 0.3),用召回率(尽可能识别少数类)换取准确率。
-
重新定义类别:合并相似类别(如将 “篮球”“足球” 合并为 “球类”),减少类别数量,提升每个类别的样本量。
8.2 标签不均衡问题
8.2.1 问题定义
不同类别的样本数量差异极大(如某类别样本数 10000,另一类别仅 50),模型会偏向预测多样本类别,导致少样本类别预测精度极低。
8.2.2 解决方案
-
基础方案:数据稀疏的所有解决方案均适用(如标注更多少样本类数据、数据增强、使用预训练模型)。
-
针对性方案:
-
过采样:复制少样本类的样本(或通过 SMOTE 等算法生成相似样本),使少样本类的样本量与多样本类接近。需注意避免过拟合(如复制次数过多导致模型记住重复样本)。
-
降采样:随机删除多样本类的部分样本(保留核心样本),平衡类别分布。需注意避免删除关键样本(可通过聚类选择代表性样本保留)。
-
调整样本权重:在损失函数中为少样本类分配更高的权重(如多样本类权重为 1,少样本类权重为 100),使模型在训练时更关注少样本类的预测误差。
-
8.3 多标签分类问题
8.3.1 问题定义
多标签分类与多分类的核心区别:多分类中每个样本仅属于一个类别(如文本属于 “经济” 或 “体育”),多标签中每个样本可属于多个类别(如电影描述 “战斗中负伤的前海军战士操纵阿凡达”,标签为 “动作”“科幻”)。
8.3.2 解决方案
-
分解为多个二分类问题:为每个标签构建一个二分类模型,判断样本是否属于该标签。示例:标签为 [动作,科幻,爱情],构建 3 个二分类模型(动作 - 非动作、科幻 - 非科幻、爱情 - 非爱情),样本 “动作 + 科幻” 在 3 个模型中分别预测为 “1、1、0”。优势:实现简单,可针对每个标签单独优化。
-
转化为多分类问题:将所有可能的标签组合视为一个新类别,如标签 [动作,科幻,爱情] 的组合类别包括 “动作”“科幻”“爱情”“动作 + 科幻”“动作 + 爱情”“科幻 + 爱情”“动作 + 科幻 + 爱情”“无标签”,共 8 个类别。优势:能捕捉标签间的关联(如 “动作 + 科幻” 常同时出现)。缺点:标签数量多时,组合类别呈指数级增长(如 10 个标签有 1024 个组合),样本稀疏问题严重。
-
直接使用多标签损失函数:无需拆解任务,通过修改损失函数使模型直接输出多个标签的概率。常用损失函数:BCELoss(二分类交叉熵损失),每个标签独立计算损失,总和为总损失。模型输出:每个标签的概率(0-1),设定阈值(如 0.5),概率大于阈值的标签即为预测结果。优势:最直接的多标签解决方案,无需修改数据结构,适合标签数量较多的场景。
Pipline
config.py
import torch
# 数据配置
DATA_PATH = "文本分类练习.csv" # 你的CSV数据路径
LABEL_COL = 0 # 标签列索引(第一列)
TEXT_COL = 1 # 文本列索引(第二列)
MAX_SEQ_LEN = 50 # 最大序列长度(后续根据数据长度统计可调整)
TEST_RATIO = 0.1 # 测试集占比
VAL_RATIO = 0.1 # 验证集占比(从训练集中拆分)
CHECK_DATA_HEAD = True # 新增:检查CSV是否有表头,自动处理
MIN_FREQ = 2 # 词汇表最小词频(过滤低频词)
# 训练配置
BATCH_SIZE = 32
EPOCHS = 2
LEARNING_RATE = {
"fasttext": 0.01,
"textcnn": 0.001,
"bilstm": 0.001,
"bert": 2e-5
}
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") # 自动选择GPU/CPU
DROPOUT = 0.5
# 模型配置
EMBEDDING_SIZE = 128 # fasttext/TextCNN/BiLSTM的嵌入维度
HIDDEN_SIZE = 128 # BiLSTM隐藏层维度
CNN_PARAMS = {
"kernel_sizes": [3, 4, 5], # 多尺寸卷积核
"num_kernels": 128 # 每个尺寸的卷积核数量
}
BERT_MODEL_NAME = "bert-base-chinese" # 中文BERT预训练模型
loader.py
import pandas as pd
import jieba
import torch
import numpy as np
from collections import defaultdict
from sklearn.model_selection import train_test_split
from torch.utils.data import Dataset, DataLoader
from transformers import AutoTokenizer
import config
# 加载停用词(可选,提升模型效果)
def load_stopwords():
try:
with open("stopwords.txt", "r", encoding="utf-8") as f:
stopwords = set(f.read().splitlines())
except FileNotFoundError:
stopwords = set() # 无停用词文件则不使用
return stopwords
STOPWORDS = load_stopwords()
# 文本清洗与分词
def clean_and_tokenize(text):
# 去除特殊符号、数字,保留中文
text = "".join([c for c in text if c.isalpha() or c.isspace()])
# 分词
tokens = jieba.lcut(text.strip())
# 过滤停用词和长度<=1的词
tokens = [token for token in tokens if token not in STOPWORDS and len(token) > 1]
return tokens
# 统计文本长度分布(重点关注数据长度)
def analyze_text_length(texts):
tokenized_texts = [clean_and_tokenize(text) for text in texts]
lengths = [len(tokens) for tokens in tokenized_texts]
print(f"文本长度统计:")
print(f" 最小长度:{min(lengths)}")
print(f" 最大长度:{max(lengths)}")
print(f" 平均长度:{np.mean(lengths):.2f}")
print(f" 中位数长度:{np.median(lengths):.2f}")
print(f" 95%分位数长度:{np.percentile(lengths, 95):.2f}")
return tokenized_texts, lengths
# 手动构建词汇表(替代torchtext的build_vocab_from_iterator)
def build_vocab_manual(tokenized_texts):
# 统计词频
word_freq = defaultdict(int)
for tokens in tokenized_texts:
for token in tokens:
word_freq[token] += 1
# 筛选高频词(保留出现次数>=MIN_FREQ的词)
filtered_words = [word for word, freq in word_freq.items() if freq >= config.MIN_FREQ]
# 构建词汇表:<pad>→0(填充符),<unk>→1(未登录词)
word_to_idx = {
"<pad>": 0,
"<unk>": 1
}
for word in filtered_words:
word_to_idx[word] = len(word_to_idx)
print(f"词汇表大小:{len(word_to_idx)}")
return word_to_idx
# 自定义数据集类(适配AutoTokenizer与手动词汇表)
class CommentDataset(Dataset):
def __init__(self, texts, labels, vocab=None, is_bert=False):
self.texts = texts
self.labels = labels
self.vocab = vocab # 仅非BERT模型使用
self.is_bert = is_bert
# 初始化AutoTokenizer(BERT专用)
self.bert_tokenizer = AutoTokenizer.from_pretrained(config.BERT_MODEL_NAME) if is_bert else None
def __len__(self):
return len(self.labels)
def __getitem__(self, idx):
text = self.texts[idx]
label = self.labels[idx]
if self.is_bert:
# BERT编码:使用AutoTokenizer(替代原BertTokenizer)
encoding = self.bert_tokenizer(
text,
max_length=config.MAX_SEQ_LEN,
padding="max_length",
truncation=True,
return_tensors="pt"
)
return {
"input_ids": encoding["input_ids"].flatten(),
"attention_mask": encoding["attention_mask"].flatten(),
"label": torch.tensor(label, dtype=torch.long)
}
else:
# 非BERT编码:手动词汇表映射+截断/填充(替代torchtext的vocab)
tokens = clean_and_tokenize(text)
# 词汇表映射(未登录词→<unk>的索引1)
indices = [self.vocab.get(token, self.vocab["<unk>"]) for token in tokens]
# 截断(超过MAX_SEQ_LEN)或填充(不足MAX_SEQ_LEN)
if len(indices) > config.MAX_SEQ_LEN:
indices = indices[:config.MAX_SEQ_LEN]
else:
indices += [self.vocab["<pad>"]] * (config.MAX_SEQ_LEN - len(indices))
return {
"text": torch.tensor(indices, dtype=torch.long),
"label": torch.tensor(label, dtype=torch.long)
}
# 数据加载主函数(核心修改:表头检测)
def load_data(is_bert=False):
# 1. 读取CSV数据(新增表头检测)
print(f"正在读取数据文件:{config.DATA_PATH}")
try:
# 先尝试读取前5行,判断是否有表头(标签列是否为数字0/1)
sample_df = pd.read_csv(config.DATA_PATH, nrows=5, header=None)
# 检查第一列(标签列)是否为数字0/1
is_label_numeric = pd.to_numeric(sample_df.iloc[:, config.LABEL_COL], errors="coerce").notna().all()
has_header = not is_label_numeric
if has_header and config.CHECK_DATA_HEAD:
print("检测到CSV存在表头,将跳过表头读取数据")
df = pd.read_csv(config.DATA_PATH, header=0) # 跳过表头
# 重新指定列索引(表头行不参与数据)
labels = pd.to_numeric(df.iloc[:, config.LABEL_COL], errors="coerce").values.astype(int)
texts = df.iloc[:, config.TEXT_COL].astype(str).values
else:
print("CSV无表头或表头为有效标签,直接读取数据")
df = pd.read_csv(config.DATA_PATH, header=None)
labels = pd.to_numeric(df.iloc[:, config.LABEL_COL], errors="coerce").values.astype(int)
texts = df.iloc[:, config.TEXT_COL].astype(str).values
# 数据清洗:移除标签为NaN的行
valid_mask = ~np.isnan(labels)
labels = labels[valid_mask]
texts = texts[valid_mask]
print(f"数据读取完成,有效数据量:{len(texts)}条")
print(f"标签分布:正类(1){sum(labels)}条,负类(0){len(labels)-sum(labels)}条")
except Exception as e:
print(f"数据读取错误:{str(e)}")
raise ValueError("请检查CSV文件格式(第一列0/1标签,第二列文本)")
# 2. 统计文本长度(无修改)
print("="*50)
print("开始统计文本长度...")
tokenized_texts, lengths = analyze_text_length(texts)
print("="*50)
# 3. 划分数据集(无修改)
train_texts, test_texts, train_labels, test_labels = train_test_split(
texts, labels, test_size=config.TEST_RATIO, random_state=42, stratify=labels
)
train_texts, val_texts, train_labels, val_labels = train_test_split(
train_texts, train_labels, test_size=config.VAL_RATIO/(1-config.TEST_RATIO),
random_state=42, stratify=train_labels
)
print(f"数据划分结果:")
print(f" 训练集:{len(train_texts)}条")
print(f" 验证集:{len(val_texts)}条")
print(f" 测试集:{len(test_texts)}条")
# 4. 构建词汇表(无修改)
vocab = None
if not is_bert:
print("构建词汇表...")
train_tokenized = [clean_and_tokenize(text) for text in train_texts]
vocab = build_vocab_manual(train_tokenized)
# 5. 创建DataLoader(无修改)
train_dataset = CommentDataset(train_texts, train_labels, vocab, is_bert)
val_dataset = CommentDataset(val_texts, val_labels, vocab, is_bert)
test_dataset = CommentDataset(test_texts, test_labels, vocab, is_bert)
train_loader = DataLoader(train_dataset, batch_size=config.BATCH_SIZE, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=config.BATCH_SIZE, shuffle=False)
test_loader = DataLoader(test_dataset, batch_size=config.BATCH_SIZE, shuffle=False)
return train_loader, val_loader, test_loader, vocab
evaluate.py
import torch
import torch.nn.functional as F
import time
import config
def evaluate_model(model, data_loader, model_name):
"""
评估模型性能:准确率 + 预测速度
"""
model.eval() # 切换到评估模式
total_correct = 0
total_samples = 0
total_time = 0
with torch.no_grad(): # 禁用梯度计算,加快速度
for batch in data_loader:
# 提取batch数据(适配BERT与非BERT模型)
if model_name == "bert":
input_ids = batch["input_ids"].to(config.DEVICE)
attention_mask = batch["attention_mask"].to(config.DEVICE)
labels = batch["label"].to(config.DEVICE)
# 计时开始
start_time = time.time()
outputs = model(input_ids=input_ids, attention_mask=attention_mask)
# 计时结束
total_time += time.time() - start_time
else:
texts = batch["text"].to(config.DEVICE)
labels = batch["label"].to(config.DEVICE)
# 计时开始
start_time = time.time()
outputs = model(texts)
# 计时结束
total_time += time.time() - start_time
# 计算预测结果
preds = torch.argmax(F.softmax(outputs, dim=1), dim=1)
total_correct += (preds == labels).sum().item()
total_samples += labels.size(0)
# 计算指标
accuracy = total_correct / total_samples
# 计算预测100条数据的平均耗时
avg_time_per_100 = (total_time / total_samples) * 100
return {
"accuracy": accuracy,
"predict_time_100": avg_time_per_100 # 预测100条的耗时(秒)
}
model.py
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import BertForSequenceClassification
import config
# 1. fastText模型
class FastText(nn.Module):
def __init__(self, vocab_size):
super(FastText, self).__init__()
self.embedding = nn.Embedding(
num_embeddings=vocab_size,
embedding_dim=config.EMBEDDING_SIZE,
padding_idx=0 # <pad>对应的索引为0
)
self.fc = nn.Linear(config.EMBEDDING_SIZE, 2) # 二分类(0/1)
self.dropout = nn.Dropout(config.DROPOUT)
def forward(self, x):
# x: [batch_size, max_seq_len]
embed = self.embedding(x) # [batch_size, max_seq_len, embedding_size]
embed = self.dropout(embed)
avg_pool = torch.mean(embed, dim=1) # 平均池化:[batch_size, embedding_size]
output = self.fc(avg_pool) # [batch_size, 2]
return output
# 2. TextCNN模型
class TextCNN(nn.Module):
def __init__(self, vocab_size):
super(TextCNN, self).__init__()
self.embedding = nn.Embedding(
num_embeddings=vocab_size,
embedding_dim=config.EMBEDDING_SIZE,
padding_idx=0
)
# 多尺寸卷积层
self.convs = nn.ModuleList([
nn.Conv1d(
in_channels=config.EMBEDDING_SIZE,
out_channels=config.CNN_PARAMS["num_kernels"],
kernel_size=ks
) for ks in config.CNN_PARAMS["kernel_sizes"]
])
self.fc = nn.Linear(
len(config.CNN_PARAMS["kernel_sizes"]) * config.CNN_PARAMS["num_kernels"],
2
)
self.dropout = nn.Dropout(config.DROPOUT)
def forward(self, x):
# x: [batch_size, max_seq_len]
embed = self.embedding(x).permute(0, 2, 1) # 转换维度:[batch_size, embedding_size, max_seq_len]
embed = self.dropout(embed)
# 卷积+最大池化
conv_outputs = []
for conv in self.convs:
conv_out = conv(embed) # [batch_size, num_kernels, max_seq_len - kernel_size + 1]
pool_out = F.max_pool1d(conv_out, kernel_size=conv_out.size(2)).squeeze(2) # [batch_size, num_kernels]
conv_outputs.append(pool_out)
# 拼接所有卷积核的输出
concat_out = torch.cat(conv_outputs, dim=1) # [batch_size, num_kernels * kernel_sizes_num]
output = self.fc(concat_out) # [batch_size, 2]
return output
# 3. BiLSTM模型
class BiLSTM(nn.Module):
def __init__(self, vocab_size):
super(BiLSTM, self).__init__()
self.embedding = nn.Embedding(
num_embeddings=vocab_size,
embedding_dim=config.EMBEDDING_SIZE,
padding_idx=0
)
self.lstm = nn.LSTM(
input_size=config.EMBEDDING_SIZE,
hidden_size=config.HIDDEN_SIZE,
bidirectional=True, # 双向LSTM
batch_first=True,
dropout=config.DROPOUT if config.HIDDEN_SIZE > 1 else 0
)
self.fc = nn.Linear(config.HIDDEN_SIZE * 2, 2) # 双向→2*hidden_size
self.dropout = nn.Dropout(config.DROPOUT)
def forward(self, x):
# x: [batch_size, max_seq_len]
embed = self.embedding(x) # [batch_size, max_seq_len, embedding_size]
embed = self.dropout(embed)
# LSTM前向传播
lstm_out, _ = self.lstm(embed) # [batch_size, max_seq_len, 2*hidden_size]
# 取最后一个时间步的输出
last_hidden = lstm_out[:, -1, :] # [batch_size, 2*hidden_size]
output = self.fc(last_hidden) # [batch_size, 2]
return output
# 4. BERT模型
class BertClassifier(nn.Module):
def __init__(self):
super(BertClassifier, self).__init__()
self.bert = BertForSequenceClassification.from_pretrained(
config.BERT_MODEL_NAME,
num_labels=2 # 二分类
)
def forward(self, input_ids, attention_mask):
# input_ids: [batch_size, max_seq_len]
# attention_mask: [batch_size, max_seq_len](标记是否为有效词)
output = self.bert(input_ids=input_ids, attention_mask=attention_mask)
return output[0] # [batch_size, 2]
# 模型初始化函数
def init_model(model_name, vocab_size=None):
if model_name == "fasttext":
assert vocab_size is not None, "fasttext需要词汇表大小"
model = FastText(vocab_size).to(config.DEVICE)
elif model_name == "textcnn":
assert vocab_size is not None, "textcnn需要词汇表大小"
model = TextCNN(vocab_size).to(config.DEVICE)
elif model_name == "bilstm":
assert vocab_size is not None, "bilstm需要词汇表大小"
model = BiLSTM(vocab_size).to(config.DEVICE)
elif model_name == "bert":
model = BertClassifier().to(config.DEVICE)
else:
raise ValueError(f"不支持的模型:{model_name}")
return model
main.py
import torch
import torch.nn as nn
import torch.optim as optim
import pandas as pd
from tqdm import tqdm # 进度条
import config
import loader
import model
import evaluate
import os
# 设置环境变量允许重复OpenMP库(仅临时测试用)
os.environ["KMP_DUPLICATE_LIB_OK"] = "TRUE"
def train_model(model_name):
"""
单模型训练流程
"""
print(f"\n{'='*60}")
print(f"开始训练模型:{model_name}")
print(f"{'='*60}")
# 1. 加载数据
is_bert = (model_name == "bert")
train_loader, val_loader, test_loader, vocab = loader.load_data(is_bert=is_bert)
vocab_size = len(vocab) if not is_bert else None
# 2. 初始化模型、损失函数、优化器
net = model.init_model(model_name, vocab_size)
criterion = nn.CrossEntropyLoss() # 二分类交叉熵损失
optimizer = optim.Adam(net.parameters(), lr=config.LEARNING_RATE[model_name])
# 3. 训练循环
best_val_acc = 0.0
best_model_path = f"best_{model_name}.pth"
for epoch in range(config.EPOCHS):
net.train() # 切换到训练模式
train_loss = 0.0
train_bar = tqdm(train_loader, desc=f"Epoch {epoch+1}/{config.EPOCHS}")
for batch in train_bar:
optimizer.zero_grad() # 清空梯度
# 前向传播
if model_name == "bert":
input_ids = batch["input_ids"].to(config.DEVICE)
attention_mask = batch["attention_mask"].to(config.DEVICE)
labels = batch["label"].to(config.DEVICE)
outputs = net(input_ids=input_ids, attention_mask=attention_mask)
else:
texts = batch["text"].to(config.DEVICE)
labels = batch["label"].to(config.DEVICE)
outputs = net(texts)
# 计算损失
loss = criterion(outputs, labels)
train_loss += loss.item()
# 反向传播+参数更新
loss.backward()
optimizer.step()
# 更新进度条
train_bar.set_postfix(loss=loss.item())
# 每轮训练后验证
val_metrics = evaluate.evaluate_model(net, val_loader, model_name)
val_acc = val_metrics["accuracy"]
print(f"Epoch {epoch+1} - 训练损失:{train_loss/len(train_loader):.4f},验证准确率:{val_acc:.4f}")
# 保存最优模型(基于验证集准确率)
if val_acc > best_val_acc:
best_val_acc = val_acc
torch.save(net.state_dict(), best_model_path)
print(f"保存最优模型(验证准确率:{best_val_acc:.4f})")
# 4. 测试集评估(加载最优模型)
print(f"\n开始测试模型:{model_name}")
net.load_state_dict(torch.load(best_model_path))
test_metrics = evaluate.evaluate_model(net, test_loader, model_name)
print(f"测试集准确率:{test_metrics['accuracy']:.4f}")
print(f"预测100条数据耗时:{test_metrics['predict_time_100']:.4f}秒")
# 返回测试集指标
return {
"model": model_name,
"learning_rate": config.LEARNING_RATE[model_name],
"hidden_size/embedding_size": config.HIDDEN_SIZE if model_name == "bilstm" else config.EMBEDDING_SIZE,
"卷积核参数": f"尺寸{config.CNN_PARAMS['kernel_sizes']},数量{config.CNN_PARAMS['num_kernels']}" if model_name == "textcnn" else "-",
"test_accuracy": test_metrics["accuracy"],
"predict_time_100": test_metrics["predict_time_100"]
}
if __name__ == "__main__":
# 待训练的模型列表
model_list = ["fasttext", "textcnn", "bilstm", "bert"]
# 存储所有模型的结果
results = []
# 逐个训练模型
for model_name in model_list:
result = train_model(model_name)
results.append(result)
# 生成实验结果表格
results_df = pd.DataFrame(results)
results_df.to_csv("模型对比结果.csv", index=False, encoding="utf-8-sig")
print(f"\n{'='*60}")
print("所有模型训练完成!实验结果如下:")
print(f"{'='*60}")
print(results_df.round(4))

根据上面的实验数据,我们得出如下结论:
- 精度排序:BERT > TextCNN > BiLSTM > fastText,预训练模型(BERT)在语义理解上优势明显。
- 速度排序:fastText > TextCNN > BiLSTM > BERT,简单模型(fastText)训练与预测速度更快。
更多推荐

所有评论(0)