一、文本分类基础认知

1.1 核心定义

文本分类是预先设定文本类别集合,对单篇文本预测其所属类别的 NLP 任务。核心逻辑是建立 “文本内容 - 类别标签” 的映射关系,本质是通过算法学习文本特征与类别间的关联模式。

1.2 典型案例

分类类型 文本示例 对应类别 应用场景
情感分析 “这家饭店太难吃了” 负类 电商评论情感判断、舆情情绪分析
情感分析 “这家菜很好吃” 正类 同上
领域分类 “今日 A 股行情大好” 经济 新闻频道划分、资讯推荐领域定位
领域分类 “今日湖人击败勇士” 体育 同上

二、文本分类核心应用场景

文本分类广泛渗透于信息处理各环节,核心场景可分为以下四类:

2.1 资讯内容打标签

  • 核心目标:为新闻、文章等资讯分配多级类别标签,实现内容结构化。

  • 案例:新浪网的分类体系(一级分类:新闻、体育、财经、娱乐等;二级分类:体育下的 NBA、英超、中超,财经下的股票、基金、外汇),用户可通过分类快速定位感兴趣内容。

  • 价值:提升资讯检索效率,支撑个性化推荐(如向体育爱好者推送 NBA 新闻)。

2.2 电商评论分析

  • 核心目标:分析用户对商品的评价态度,辅助商家优化产品、平台提升用户体验。

  • 案例:电商平台中用户对商品的评价(“还不错哦”“能用么!”),通过分类判断评价为 “正面”“中性”“负面”,统计商品的好评率、差评核心原因(如 “质量差”“尺码不准”)。

  • 价值:商家可针对性改进产品(如某款鞋子差评多为 “磨脚”,则优化鞋型),平台为用户优先推荐高好评商品。

2.3 违规内容检测

  • 核心目标:识别文本中的违规信息,维护内容生态安全。

  • 检测范围:涉黄、涉暴、涉恐、辱骂、虚假宣传等违规内容。

  • 应用场景

    • 客服 / 销售对话质检:检测客服是否使用辱骂话术、销售是否夸大产品功效。

    • 网站内容审查:过滤论坛、评论区的违规言论,避免不良信息传播。

三、自定义类别任务特性

自定义类别任务是文本分类的灵活延伸,核心特点的 “类别定义自由度高”,只要人能通过文本判断的类别,均可作为分类目标。常见场景包括:

  • 垃圾邮件分类:类别为 “垃圾邮件”“正常邮件”,通过识别 “广告推销”“诈骗链接” 等特征判断。

  • 汽车交易相关性判断:类别为 “与汽车交易相关”“与汽车交易无关”,用于筛选汽车电商平台的有效对话 / 文章。

  • 作者风格识别:类别为 “某作者风格”“非某作者风格”,如判断一篇散文是否为鲁迅所作(基于用词、句式特征)。

  • 机器生成文本检测:类别为 “机器生成”“人工撰写”,用于识别 AI 生成的新闻、论文等内容。

  • 合同文本合规性判断:类别为 “符合规范”“不符合规范”,检测合同中是否存在无效条款、法律风险表述。

  • 阅读人群适配性分类:类别为 “未成年适宜”“中年适宜”“老年适宜”“孕妇适宜” 等,用于儿童读物、老年健康文章的精准推送。

四、文本分类的机器学习流程

机器学习实现文本分类需遵循 “定义 - 数据 - 训练 - 预测” 的标准化流程,具体步骤如下:

4.1 流程拆解

  1. 定义类别:明确分类目标与类别集合(如情感分析定义 “正类”“负类”,领域分类定义 “经济”“体育”“科技” 等)。

  2. 收集数据:获取带类别标签的文本数据(标注数据),如情感分析需收集大量标注 “正 / 负” 的评论,领域分类需收集标注 “经济 / 体育 / 科技” 的新闻。数据质量直接影响模型效果,需保证标注准确性、类别分布合理性。

  3. 模型训练:将标注数据输入分类模型,模型学习文本特征与类别间的映射关系。核心逻辑是:模型通过计算 “文本特征属于某类别的概率”,调整参数使预测结果与真实标签尽可能一致。

  4. 预测应用:将训练好的模型用于未标注文本,输出该文本所属的类别(如输入 “今日油价上涨”,模型预测其类别为 “经济”)。

4.2 流程示意图

五、贝叶斯算法在文本分类中的应用

贝叶斯算法是基于 “贝叶斯公式” 的概率模型,核心思想是通过 “先验概率” 与 “条件概率” 计算 “后验概率”,实现类别预测。

5.1 预备知识:全概率公式

  • 公式定义:若事件组{Bi​}是样本空间Ω的一个划分(即Bi​互斥且⋃Bi​=Ω),且P(Bi​)>0,则对任意事件A,有:P(A) = \sum_{i}P(B_i)P(A|B_i)

  • 案例解释:扔正常骰子,计算 “结果为 5(事件 A)” 的概率。

    • 划分事件:B_1​(结果为奇数, P(B_1)=\frac{1}{2})、B2​(结果为偶数,P(B_2)=\frac{1}{2})。

    • 条件概率:P(A|B_1)=\frac{1}{3}(奇数包含 1、3、5,共 3 种,5 占 1 种),P(A|B_2)=0(偶数不含 5)。

    • 计算结果:P(A)=P(B_1)P(A|B_1) + P(B_2)P(A|B_2)=\frac{1}{2}*\frac{1}{3}+\frac{1}{2}*0=\frac{1}{6}与实际常识一致。

5.2 核心公式:贝叶斯公式

  • 公式推导:由联合概率P(AB)=P(A|B)P(B)=P(B|A)P(A),变形得:P(A|B)=\frac{P(B|A)P(A)}{P(B)}
  • 符号含义
    • (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%
    未感染新冠\overline{A} 5% 95%
5.3.2 计算过程
  1. 计算\(P(B)\)(检测呈阳性的总概率,用全概率公式):P(B)=P(B|A)P(A)+P(B|\overline{A})P(\overline{A}))=0.99×0.001 + 0.05×(1-0.001))=0.00099 + 0.04995 = 0.05094

  2. 计算\(P(A|B)\)(检测呈阳性时实际感染的概率):P(A|B)=\frac{P(B|A)P(A)}{P(B)}=\frac{0.99*0.001}{0.05094}\approx 0.019(即 1.9%)

5.3.3 结论

即使核酸检测呈阳性,实际感染新冠的概率仅约 1.9%,原因是人群感染率(先验概率)极低,误报率(5%)导致大量未感染者被误判为阳性,需结合临床症状进一步判断。

5.4 贝叶斯算法在文本分类中的应用

5.4.1 核心假设

文本属于某类别的概率,仅与文本中包含的词相关(即 “词的独立性假设”,简化计算)。

5.4.2 分类逻辑

假设存在 3 个类别A_1,A_2,A_3,文本S由W_1,W_2,...,W_n(n 个词)组成,目标是计算P(A_i|S)(文本S属于类别A_i的概率),选择概率最大的类别作为预测结果。

  1. 应用贝叶斯公式:P(A_i|S)=\frac{P(S|A_i)P(A_i)}{P(S)}其中,P(S)是所有类别共有的分母,比较不同\A_i的概率时可忽略,只需计算P(S|A_i)P(A_i)

  2. 词的独立性假设:文本S在类别A_i下的概率P(S|A_i),等于每个词在\(A_i\)下概率的乘积:\(P(S|A_i)=P(W_1|A_i)×P(W_2|A_i)×...×P(W_n|A_i)\)

  3. 概率计算示例:若判断文本 “今日 A 股上涨” 是否属于 “经济” 类(\(A_1\)):

    • P(A_1):“经济” 类文本在所有文本中的占比(先验概率)。
    • P(W_1|A_1):“今日” 在 “经济” 类文本中出现的概率,P(W_2|A_1):“A 股” 在 “经济” 类文本中出现的概率,P(W_3|A_1):“上涨” 在 “经济” 类文本中出现的概率。
    • 计算P(S|A_1)P(A_1),并与 “体育”“科技” 等类别的对应值比较,最大者即为预测类别。

5.5 贝叶斯算法的优缺点

5.5.1 优点
  1. 简单高效:计算逻辑清晰,无需复杂迭代训练,适合小规模数据场景。

  2. 可解释性强:概率计算过程可追溯,能明确知道每个词对类别预测的贡献。

  3. 样本覆盖好时效果优:若训练数据能充分覆盖各类别特征,预测精度较高。

  4. 支持分批训练:可将训练数据分批输入,无需一次性加载所有数据,降低内存压力。

5.5.2 缺点

  1. 样本不均衡敏感:若某类样本数量极少,其先验概率\(P(A_i)\)被低估,导致预测偏向多样本类别。

  2. 未见过特征处理困难:若文本中出现训练数据未见过的词,该词的条件概率\(P(W|A_i)=0\),导致整体概率为 0,需通过 “平滑技术”(如拉普拉斯平滑)解决。

  3. 特征独立假设不成立:实际文本中词与词存在关联(如 “A 股” 与 “上涨” 常同时出现),独立假设会损失关联信息,影响精度。

  4. 忽略语序与词义:仅考虑词的出现频率,不考虑词的顺序(如 “我喜欢他” 与 “他喜欢我” 语义不同但词相同),也无法区分多义词(如 “苹果” 指水果或公司)。

六、支持向量机(SVM)在文本分类中的应用

支持向量机(SVM)是一种有监督学习模型,核心思想是寻找 “最大边距超平面”,实现对数据的最优分类,适用于线性可分与线性不可分场景。

6.1 核心概念:最大边距超平面

6.1.1 超平面定义

在二维空间中,超平面是直线(如wx+b=0);在三维空间中是平面;在高维空间中是维度为 “空间维度 - 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,映射到二维特征空间\phi(x)=[x, x^2]:正样本[-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 常见核函数
核函数类型 公式 适用场景
线性核函数 K(x,x')=(x\cdot x') 数据本身线性可分,或特征维度高(如文本分类的词袋特征)
多项式核函数 K(x,x')=(1+(x\cdot x'))^d 数据呈多项式分布,需捕捉非线性关系
高斯核函数(RBF) K(x,x')=exp(-\frac{|x-x'|^2}{2\sigma^2}) 数据分布复杂,无法确定非线性关系类型,应用最广泛
双曲正切核函数 K(x,x')=tanh(1+(x\cdot x')) 模拟神经网络,适用于需非线性映射且希望输出范围有限的场景

6.3 SVM 的多分类解决方案

SVM 本质是二分类模型,处理多分类(K 类)需通过 “拆解策略” 实现,核心有两种方式:

6.3.1 One vs One(一对一)
  • 原理:为每对类别构建一个 SVM 分类器,共需构建K(K-1)/2个分类器(如 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 优点
  1. 对异常值不敏感:超平面仅由支持向量决定,少量异常值(远离支持向量的样本)不影响超平面位置。

  2. 样本需求量低:无需大量数据即可学习到稳定的超平面,适合小样本场景。

  3. 高维数据处理能力强:即使特征维度(如文本的词袋特征维度)远大于样本数量,仍能有效分类,避免过拟合。

6.4.2 缺点
  1. 大规模数据计算负担重:样本数量过多时,支持向量数量增加,模型训练与预测的时间、空间复杂度显著上升。

  2. 多分类处理复杂:需通过 One vs One 或 One vs Rest 拆解任务,增加了模型构建与调参的复杂度。

  3. 核函数与参数选择困难:不同核函数适用于不同数据分布,需通过大量实验尝试(如高斯核的\(\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 特征使用方式

    1. 直接使用 [CLS] token 的向量作为文本特征,输入全连接层分类。

    2. 对所有 token 的向量取 Max/Average Pooling,得到文本向量。

    3. 将 BERT 输出的向量输入 LSTM/CNN,进一步提取特征。

    4. 融合 BERT 中间层的输出(如取第 3、5、7 层的向量拼接),增强特征的丰富性。

  • 适用场景:对精度要求高的复杂任务(如法律文本分类、医疗文本诊断、细粒度情感分析)。

八、文本分类的常见问题与解决方案

在文本分类实践中,常面临数据稀疏、标签不均衡、多标签分类等问题,需针对性解决:

8.1 数据稀疏问题

8.1.1 问题定义

训练数据量过少(如某类别仅几十条样本),模型在训练集上可收敛(拟合训练数据),但在测试集上预测准确率极低(泛化能力差)。

8.1.2 解决方案
  1. 标注更多数据:最直接有效的方法,通过人工标注或半监督标注(如主动学习选择高价值样本标注)增加数据量。

  2. 数据增强:构造相似样本,如文本同义替换(“很好”→“非常好”)、随机插入 / 删除停用词、句子语序调整(不改变语义)、翻译回译(中文→英文→中文)。

  3. 使用预训练模型:预训练模型(如 BERT)已在大规模数据上学习了通用语言知识,微调时需少量任务数据即可达到较好效果,减少对标注数据的依赖。

  4. 增加规则弥补:对数据稀疏的类别,制定人工规则(如 “包含‘涨停’‘跌停’的文本属于金融类”),辅助模型预测。

  5. 调整阈值:在二分类中,降低少数类的预测阈值(如将 “正类” 预测概率阈值从 0.5 调整为 0.3),用召回率(尽可能识别少数类)换取准确率。

  6. 重新定义类别:合并相似类别(如将 “篮球”“足球” 合并为 “球类”),减少类别数量,提升每个类别的样本量。

8.2 标签不均衡问题

8.2.1 问题定义

不同类别的样本数量差异极大(如某类别样本数 10000,另一类别仅 50),模型会偏向预测多样本类别,导致少样本类别预测精度极低。

8.2.2 解决方案
  • 基础方案:数据稀疏的所有解决方案均适用(如标注更多少样本类数据、数据增强、使用预训练模型)。

  • 针对性方案

    1. 过采样:复制少样本类的样本(或通过 SMOTE 等算法生成相似样本),使少样本类的样本量与多样本类接近。需注意避免过拟合(如复制次数过多导致模型记住重复样本)。

    2. 降采样:随机删除多样本类的部分样本(保留核心样本),平衡类别分布。需注意避免删除关键样本(可通过聚类选择代表性样本保留)。

    3. 调整样本权重:在损失函数中为少样本类分配更高的权重(如多样本类权重为 1,少样本类权重为 100),使模型在训练时更关注少样本类的预测误差。

8.3 多标签分类问题

8.3.1 问题定义

多标签分类与多分类的核心区别:多分类中每个样本仅属于一个类别(如文本属于 “经济” 或 “体育”),多标签中每个样本可属于多个类别(如电影描述 “战斗中负伤的前海军战士操纵阿凡达”,标签为 “动作”“科幻”)。

8.3.2 解决方案
  1. 分解为多个二分类问题:为每个标签构建一个二分类模型,判断样本是否属于该标签。示例:标签为 [动作,科幻,爱情],构建 3 个二分类模型(动作 - 非动作、科幻 - 非科幻、爱情 - 非爱情),样本 “动作 + 科幻” 在 3 个模型中分别预测为 “1、1、0”。优势:实现简单,可针对每个标签单独优化。

  2. 转化为多分类问题:将所有可能的标签组合视为一个新类别,如标签 [动作,科幻,爱情] 的组合类别包括 “动作”“科幻”“爱情”“动作 + 科幻”“动作 + 爱情”“科幻 + 爱情”“动作 + 科幻 + 爱情”“无标签”,共 8 个类别。优势:能捕捉标签间的关联(如 “动作 + 科幻” 常同时出现)。缺点:标签数量多时,组合类别呈指数级增长(如 10 个标签有 1024 个组合),样本稀疏问题严重。

  3. 直接使用多标签损失函数:无需拆解任务,通过修改损失函数使模型直接输出多个标签的概率。常用损失函数: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)训练与预测速度更快。
Logo

有“AI”的1024 = 2048,欢迎大家加入2048 AI社区

更多推荐