资讯动态

中文文本分类实战:TextCNN与RNN的TensorFlow实现

发布时间:2026/9/16 19:05:44 来源:尧图企业网站定制
简介面向中文文本分类入门者与NLP初学者基于TensorFlow完成中文文本分类同时使用卷积神经网络CNN和循环神经网络RNN。项目参考了经典论文《Convolutional Neural Networks for Sentence Classification》与《Character-level Convolutional Neural Networks for Text Classification》并在中文数据集上做简化落地覆盖数据预处理、模型构建、训练与预测等关键环节。压缩包共16个文件以9个Python脚本为主体承担CNN/RNN模型定义、数据加载、训练与预测功能另含4张模型结构及训练过程示意图、1个Shell数据准备脚本、1份Markdown说明和1个txt依赖列表整体体积仅409KB轻量易读。当前已有472人学习/下载适合希望快速通过可运行代码理解深度学习文本分类原理的读者也可作为课程设计或工程实践的参考基线。1. 中文文本分类的两种路线TextCNN 与 RNN 怎么选做中文文本分类很多人第一反应是上 BERT但在工业界和竞赛里TextCNN 和 RNN 依然有不可替代的位置。原因很直接显存占用小、训练速度快、调参空间直观而且在小数据集上表现往往不输预训练模型。这个项目的价值在于它用 TensorFlow 实现了字符级 CNN 和 RNN 两条完整链路从数据预处理到训练、评估、预测全部打通适合用来理解「卷积如何捕捉 n-gram 特征」和「循环网络如何建模序列依赖」这两类核心思路。项目基于 TensorFlow 1.x使用了字符级输入而非词级输入这对中文特别友好——中文分词本身就有误差传播问题而字符级 CNN 直接对单个汉字建模用卷积核宽度去隐式学习短语和局部搭配。RNN 则选择了 LSTM 单元处理长距离依赖时比普通 RNN 稳定得多。如果你正在做新闻分类、情感分析、意图识别这类任务这个项目的网络结构、数据管道和训练脚本可以直接改造成自己的基线模型也可以作为对比实验中的 strong baseline。整份代码读下来你会发现深度学习的文本分类并没有那么玄核心就是「把文本变成张量然后让网络自己找特征」。2. 数据管道与字符级预处理从原始文本到 batch 张量2.1 cnews 数据集的结构与加载逻辑项目中使用的是 THUCNews 的子集cnews原始文本按类别分目录存放。cnews_loader.py是整个项目的入口它承担了读文件、建词典、转索引、padding、生成 batch 的全部工作。先看数据处理的核心函数read_file和build_vocab的调用流程# 读取原始文件每行是一条样本标签\t内容 def read_file(filename): contents, labels [], [] with open(filename, r, encodingutf-8) as f: for line in f: line line.strip() if not line: continue label, content line.split(\t, 1) contents.append(content) labels.append(label) return contents, labels # 构建字符级词典vocab_size 控制保留的最高频字数 def build_vocab(train_dir, vocab_dir, vocab_size5000): contents, labels read_file(train_dir) word_count {} for content in contents: for char in content: word_count[char] word_count.get(char, 0) 1 sorted_words sorted(word_count.items(), keylambda x: x[1], reverseTrue) words [w for w, c in sorted_words[:vocab_size - 2]] with open(vocab_dir, w, encodingutf-8) as f: f.write(UNK\n PAD\n \n.join(words))这段代码的关键在于词典是按字符出现频率截断的而不是全量收录。vocab_size5000意味着只保留训练集中出现次数最多的约 5000 个字符低频字符统一映射为UNKpadding 位用PAD填充。这里要注意因为字符表是按频率排序的所以UNK和PAD一定要放在最前面保证它们的索引分别是 0 和 1否则 embedding 层初始化时容易埋坑。2.2 文本转序列的 padding 与 truncatingRNN 和 CNN 都要求输入是定长张量所以需要对变长文本做统一长度处理。cnews_loader.py中通过seq_length参数控制最大长度处理逻辑如下def process_file(filename, word_to_id, cat_to_id, seq_length600): contents, labels read_file(filename) x_data, y_data [], [] for content, label in zip(contents, labels): ids [word_to_id.get(char, word_to_id[UNK]) for char in content] if len(ids) seq_length: # 超过最大长度直接截断尾部 ids ids[:seq_length] else: # 不足最大长度在右侧补 PAD ids ids [word_to_id[PAD]] * (seq_length - len(ids)) x_data.append(ids) y_data.append(cat_to_id[label]) return np.array(x_data), np.array(y_data)截断策略是「留头去尾」因为新闻文本的主题信息通常集中在前半部分。seq_length600是项目作者基于 cnews 数据集的长度分布选的默认值。如果你换数据集我一般的做法是先统计 95 分位的文本长度再取整到 50 的倍数而不是拍脑袋定 600。process_file返回的x_data形状是[样本数, 600]后续送入网络时在 embedding 层会被扩展成[batch_size, 600, embedding_size]。2.3 batch 迭代器与类别映射batch_iter是训练时的数据供给器它把全量数据随机打乱后按 batch_size 切分def batch_iter(x, y, batch_size64, num_epochs10, shuffleTrue): data_len len(x) num_batch_per_epoch int((data_len - 1) / batch_size) 1 for epoch in range(num_epochs): if shuffle: indices np.random.permutation(np.arange(data_len)) x_shuffle, y_shuffle x[indices], y[indices] else: x_shuffle, y_shuffle x, y for i in range(num_batch_per_epoch): start_idx i * batch_size end_idx min(start_idx batch_size, data_len) yield x_shuffle[start_idx:end_idx], y_shuffle[start_idx:end_idx]num_epochs在迭代器层面控制整体训练轮数run_cnn.py里传入 10 表示全部数据过 10 遍。注意int((data_len - 1) / batch_size) 1这个写法是为了确保最后不足一个 batch 的数据也能被取出来避免丢样本。标签部分cat_to_id在read_category中从类别文件加载比如体育、财经、房产、家居、教育、科技、社会、时尚、游戏、娱乐对应 10 个分类的 one-hot 向量。3. TextCNN 的实现拆解卷积核宽度、池化与全连接3.1 为什么 CNN 能处理文本序列文本分类里的 CNN 通常指 Kim (2014) 提出的 TextCNN核心思想是把文本矩阵当作「图像」的行卷积核在词/字符的 embedding 维度上做一维滑动。每个卷积核相当于一个 n-gram 检测器例如宽度为 3 的卷积核捕捉连续 3 个字符的局部搭配模式。多个不同宽度的卷积核并行就能同时捕获 bigram、trigram、4-gram 等多种粒度的特征再通过最大池化提取每个特征图最显著的部分。3.2 cnn_model.py 的完整结构看cnn_model.py中的核心类TCNNConfig和TextCNNclass TCNNConfig: embedding_dim 64 # 字符向量的维度 seq_length 600 # 输入文本长度 num_classes 10 # 分类数 num_filters 256 # 每种卷积核的数量 filter_sizes [2, 3, 4] # 卷积核宽度 lr 1e-3 # 学习率 dropout_keep_prob 0.5 # dropout 保留比例 class TextCNN: def __init__(self, config): self.input_x tf.placeholder(tf.int32, [None, config.seq_length], nameinput_x) self.input_y tf.placeholder(tf.float32, [None, config.num_classes], nameinput_y) self.keep_prob tf.placeholder(tf.float32, namekeep_prob) # 字符 embedding 层随机初始化后参与训练 embedding tf.get_variable(embedding, [config.vocab_size, config.embedding_dim], initializertf.truncated_normal_initializer(stddev0.1)) embedding_inputs tf.nn.embedding_lookup(embedding, self.input_x) # 在 embedding 矩阵上增加通道维度形状变为 [batch, 600, 64, 1] embedding_inputs tf.expand_dims(embedding_inputs, -1) pooled_outputs [] for i, filter_size in enumerate(config.filter_sizes): with tf.name_scope(conv-maxpool-%s % filter_size): # 卷积核形状 [filter_size, embedding_dim, 1, num_filters] conv tf.layers.conv2d(embedding_inputs, filtersconfig.num_filters, kernel_size[filter_size, config.embedding_dim], activationtf.nn.relu) # 输出形状 [batch, 600 - filter_size 1, 1, num_filters] pooled tf.layers.max_pooling2d(conv, pool_size[config.seq_length - filter_size 1, 1], strides1) pooled_outputs.append(pooled) # 拼接所有卷积核的输出到同一个维度 num_filters_total config.num_filters * len(config.filter_sizes) h_pool tf.concat(pooled_outputs, 3) h_pool_flat tf.reshape(h_pool, [-1, num_filters_total]) # dropout 全连接层输出分类概率 h_drop tf.nn.dropout(h_pool_flat, keep_probself.keep_prob) self.logits tf.layers.dense(h_drop, config.num_classes, namefc) self.prob tf.nn.softmax(self.logits, nameprob)这里最容易被忽略的是kernel_size[filter_size, embedding_dim]——第二个维度直接取 embedding 的完整宽度意味着卷积核只在文本长度方向上滑动不会跨 embedding 维度滑动。这对文本任务是正确的每个字符在同一维度的语义空间内比较才有意义。max_pooling2d的pool_size取[seq_length - filter_size 1, 1]把整个特征图压成一个标量相当于「整个文本中这个 n-gram 模式最强烈的信号」。3.3 训练脚本与前向验证run_cnn.py中训练循环的核心代码with tf.Session() as sess: sess.run(tf.global_variables_initializer()) for batch_x, batch_y in batch_iter(x_train, y_train, batch_size64, num_epochs10): feed_dict {model.input_x: batch_x, model.input_y: batch_y, model.keep_prob: 0.5} _, loss, acc sess.run([train_op, model.loss, model.acc], feed_dictfeed_dict) # 每个批次结束后按步数打印训练指标 if step % 100 0: print(step {}: loss {:.4f}, acc {:.4f}.format(step, loss, acc))训练时keep_prob0.5评估和预测时keep_prob1.0这是 dropout 的标准用法。model.loss在类的定义中用的是 softmax 交叉熵model.acc则是将logits的 argmax 与真实标签比较计算平均正确率。TextCNN 在这个数据集上的表现通常训练 3~5 个 epoch 就能达到 90% 以上的准确率收敛速度明显快于 RNN这也是它在工业界常被用作第一版基线的原因。4. RNN 序列建模为什么这里用 LSTM 而不是朴素 RNN4.1 LSTM 的门控机制与文本拟合能力rnn_model.py里的 RNN 结构使用的是 LSTM 单元而不是原始 RNN 单元。理由很好理解基础 RNN 在反向传播时梯度要连乘多个时间步的权重矩阵序列超过 200 步就容易梯度消失或爆炸。LSTM 通过遗忘门、输入门、输出门三个门控结构让梯度可以沿「记忆通道」长距离传播。中文文本平均长度在几百个字符之间seq_length600的设定下LSTM 几乎是必然选择。看rnn_model.py的网络定义部分class TRNNConfig: embedding_dim 64 seq_length 600 num_classes 10 hidden_dim 128 # LSTM 隐藏层单元数 num_layers 2 # 双层 LSTM rnn_cell lstm # 可选 rnn / gru lr 1e-3 class TextRNN: def __init__(self, config): self.input_x tf.placeholder(tf.int32, [None, config.seq_length], nameinput_x) self.input_y tf.placeholder(tf.float32, [None, config.num_classes], nameinput_y) self.keep_prob tf.placeholder(tf.float32, namekeep_prob) # embedding 层结构与 CNN 共用 embedding tf.get_variable(embedding, [config.vocab_size, config.embedding_dim], initializertf.truncated_normal_initializer(stddev0.1)) inputs tf.nn.embedding_lookup(embedding, self.input_x) # 输入形状 [batch, 600, 64] if config.rnn_cell lstm: cell_fn tf.nn.rnn_cell.BasicLSTMCell elif config.rnn_cell rnn: cell_fn tf.nn.rnn_cell.BasicRNNCell else: cell_fn tf.nn.rnn_cell.GRUCell cells [cell_fn(config.hidden_dim) for _ in range(config.num_layers)] # 多层 RNN 之间用 dropout 连接 rnn_cells tf.nn.rnn_cell.MultiRNNCell([tf.nn.rnn_cell.DropoutWrapper(cell, output_keep_probself.keep_prob) for cell in cells]) # 动态展开序列输入形状 [batch, max_time, input_size] outputs, state tf.nn.dynamic_rnn(rnn_cells, inputs, dtypetf.float32) # 取最后一个时间步的输出作为整个序列的表示 self.logits tf.layers.dense(outputs[:, -1, :], config.num_classes, namefc) self.prob tf.nn.softmax(self.logits, nameprob)hidden_dim128控制了 LSTM 记忆容量num_layers2让网络有机会学习更高层的语义抽象。第一层 LSTM 关注字符间的短程依赖第二层在此基础上组合出更长距离的模式。DropoutWrapper包裹在每个 cell 外层只对输出做 dropout——注意这里不要对输入也加 dropout否则会干扰 embedding 的语义空间。4.2 最后一个时间步输出 vs 全局池化RNN 的序列输出有两种常用取法取最后一个时间步的 hidden state或者对所有时间步的输出做平均池化/最大池化。项目里用的是outputs[:, -1, :]这种方法直觉上对应「读完整段文本后的最终记忆」。它对文本结尾信息有偏重如果任务的关键信号分布在中部或开头建议改成tf.reduce_mean(outputs, axis1)做全局平均池化。我实际对比过两种方式在情感分析类任务上平均池化通常更稳在新闻分类上最后一位输出略微占优具体还是要看你数据的类别分布。4.3 对比实验的观察视角acc_loss_rnn.png和acc_loss_cnn.png是两种模型训练过程的 loss 和 acc 曲线图cnn_architecture.png和rnn_architecture.png是结构示意图。运行run_cnn.py和run_rnn.py后你会观察到几个典型差异CNN 在前 500 步内准确率快速攀升而 RNN 需要更长时间热身单个 epoch 的耗时 RNN 明显高于 CNN因为 LSTM 的时序计算难以充分并行但 RNN 在长文本上的最终准确率通常高出 CNN 1~2 个百分点。这两个模型没有绝对的强弱之分在资源允许时把两者都跑一遍取集成或择优部署都是常见做法。5. 预测流程与模型复用从训练权重到单条样本分类5.1 predict.py 的加载与推理细节predict.py展示了如何从磁盘恢复训练好的模型并完成单条文本分类。关键步骤是重建相同的网络结构、加载 checkpoint、然后把原始文本走一遍完整的预处理链def predict_one(sentence, model_dir, word_to_id, cat_to_id, seq_length600): # 重建模型结构必须与训练时完全一致 config TCNNConfig() model TextCNN(config) saver tf.train.Saver() with tf.Session() as sess: # 从 checkpoint 恢复全部变量 saver.restore(sess, tf.train.latest_checkpoint(model_dir)) # 单条文本转 id 序列长度不足补 PAD ids [word_to_id.get(char, word_to_id[UNK]) for char in sentence] if len(ids) seq_length: ids ids[:seq_length] else: ids ids [word_to_id[PAD]] * (seq_length - len(ids)) feed_dict {model.input_x: [ids], model.keep_prob: 1.0} prob sess.run(model.prob, feed_dictfeed_dict) pred_idx np.argmax(prob, axis1)[0] return cat_to_id_reverse[pred_idx], prob[0][pred_idx]推理时keep_prob一定要设为 1.0如果沿用训练时的 0.5输出概率会被随机置零导致预测结果不稳定。另外sentence里的字符如果在训练集词典中不存在会被映射为UNK的 id 0embedding 层对这个 id 的向量是训练过程中学出来的「未知字符平均表示」。5.2 checkpoint 文件管理与服务化改造训练结束后save_dir下会生成.meta、.index、.data-*三类文件和 checkpoint 索引文件。.meta里保存的是网络结构图.data里是变量值。tf.train.latest_checkpoint(save_dir)会自动找到最新的 checkpoint不用手写具体 epoch 编号。如果你要部署成 HTTP 接口建议把predict_one里的tf.Session()提取为全局单例避免每次请求都重建图。可以用tf.graph_util.convert_variables_to_constants把模型冻结成单文件 pb 格式这样就不需要再维护变量列表和 checkpoint 元数据。5.3 超参数调整的优先级建议如果换数据集后效果不理想按以下顺序排查先调seq_length过短会截掉有效信息过长会增加 padding 占比浪费算力再调embedding_dim从 64 起步往 100、128 试文本语义空间偏大时低维向量表达不够然后是卷积核数量和宽度组合——常见配置是[2, 3, 4]或[3, 4, 5]各 128 或 256 个宽度越大越偏短语级特征最后才是学习率。对于 RNNhidden_dim从 128 往 256 调num_layers在数据量不足时保持 1~2 层即可过深反而过拟合。对抗过拟合优先调整dropout_keep_prob从 0.5 降到 0.3或者把num_filters从 256 降到 128而不是盲目加数据增强——文本分类的空间扰动很难定义字符级 dropout 或同义词替换并不是银弹。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价