首页/新闻资讯/正文详情

从零复现LSTM天池新闻文本分类:一份能跑通的Python源码

发布时间:2026/9/23 7:53:11 来源:云帆数科 栏目:资讯中心
从零复现LSTM天池新闻文本分类:一份能跑通的Python源码
简介这份Python源码包围绕天池新闻文本分类比赛展开采用LSTM作为核心模型适合人工智能、计算机及相关专业学生、教师与企业员工用于课程设计、毕业设计或赛题复现。包内共25个文件以14个py脚本为主体辅以9个pyc编译文件、1个txt词表与1个json配置压缩包约58KB结构紧凑便于快速阅读与二次开发。代码涵盖LSTM编码器、TextCNN编码器、注意力机制、BERT编码器、数据预处理、模型训练与对抗训练等模块并配有训练入口脚本可帮助读者理解从文本向量化到模型评估的完整流程。已有161人学习说明该方案具备一定参考价值。对于希望掌握新闻文本分类实战、对比不同编码器效果或在此基础上迁移到其他NLP任务的读者这份源码能提供可运行的基线实现与清晰的模块划分降低从零搭建的门槛。1. 从零复现 LSTM 天池新闻文本分类一份能跑通的 Python 源码该长什么样天池新闻文本分类这个赛题很多人第一次跑 LSTM 都会卡在同一个地方代码能跑但分数上不去或者干脆在词表构建那一步就报内存错误。我见过太多人把python环境配好、pycharm配置完拿到一份源码直接python train.py结果要么是KeyError要么是 loss 不降最后不了了之。这份基于 LSTM 的新闻文本分类源码核心要解决的就是把 14 类新闻标题和正文映射成固定长度序列用 Embedding LSTM 全连接做多分类。它适合已经会python基础语法、想拿一个完整 NLP 项目练手的人也适合想搞懂「为什么我的 LSTM 比别人的低 5 个点」的熟手。下面我按自己复现时的顺序把数据、模型、训练、调参和踩坑一条条拆开。2. 数据读取与词表构建别让内存和 OOV 拖垮你的 LSTM2.1 天池新闻数据的真实结构和读取方式天池新闻文本分类的原始数据通常是train_set.csv和test_a.csv每行包含label和text两列text是新闻标题加正文的拼接用空格分词后的形式。很多人直接用pd.read_csv读然后text.split()做词表这在数据量小的时候没问题但天池这个赛题训练集有 20 万条每条平均几百个词全量加载后内存占用很容易超过 8G。我一般会先看数据分布再决定要不要做截断。import pandas as pd from collections import Counter # 读取训练集指定分隔符和列名 train_df pd.read_csv(train_set.csv, sep\t) # 查看类别分布确认是否均衡 print(train_df[label].value_counts().sort_index()) # 统计每条文本的词数分布决定 max_len text_len train_df[text].apply(lambda x: len(x.split())) print(text_len.describe(percentiles[0.5, 0.9, 0.95, 0.99]))这段代码先确认标签是否从 0 到 13 连续再通过分位数看 95% 的文本长度落在哪里。如果 95% 分位数是 800那max_len设 1000 就够设 2000 只会让 LSTM 的序列过长梯度回传变慢显存也吃紧。参数上sep\t是天池数据常见的制表符分隔如果实际是逗号改成sep,即可。percentiles列表里我习惯加 0.99防止极端长文本影响判断。2.2 词表构建的两种策略和 OOV 处理词表构建直接决定 Embedding 层的输入质量。常见做法是统计所有训练文本的词频取 top N 个词剩下的映射为UNK。但这里有个坑如果只对训练集建词表测试集里出现的新词全变UNK模型在验证集上会掉点。我一般会把训练集和测试集合并后一起统计词频再切分。另外词表大小不是越大越好天池这个赛题词表控制在 5 万到 10 万之间比较稳太大 Embedding 参数量暴涨太小 OOV 太多。from collections import Counter # 合并训练和测试的文本统一建词表 test_df pd.read_csv(test_a.csv, sep\t) all_text pd.concat([train_df[text], test_df[text]], ignore_indexTrue) # 统计词频只保留出现次数 2 的词 word_counter Counter() for text in all_text: word_counter.update(text.split()) # 构建词表0 留给 padding1 留给 UNK vocab {PAD: 0, UNK: 1} for word, count in word_counter.most_common(): if count 2: break if len(vocab) 100000: break vocab[word] len(vocab) print(f词表大小: {len(vocab)})这里count 2过滤掉只出现一次的词能显著减小词表且对精度影响很小。len(vocab) 100000是硬上限防止个别高频噪声词撑爆词表。建完词表后把文本转成 id 序列时遇到不在词表里的词就填 1。注意PAD和UNK的 id 必须固定后面 Embedding 层要对应。2.3 把文本转成定长 id 序列的完整函数有了词表下一步是把每条文本转成固定长度的 id 列表。短了补 0长了截断。截断策略有从头部截、从尾部截、头尾各截一半新闻文本的关键信息通常在开头我一般保留前max_len个词。def text_to_ids(text, vocab, max_len1000): 将文本转为定长 id 序列不足补 0超出截断 words text.split() ids [vocab.get(w, 1) for w in words] # 1 是 UNK if len(ids) max_len: ids ids [0] * (max_len - len(ids)) else: ids ids[:max_len] return ids # 应用到训练集和测试集 train_df[ids] train_df[text].apply(lambda x: text_to_ids(x, vocab)) test_df[ids] test_df[text].apply(lambda x: text_to_ids(x, vocab))vocab.get(w, 1)保证 OOV 词映射到UNK。补 0 操作放在后面因为 LSTM 对 padding 位置可以通过pack_padded_sequence忽略但为了简单很多源码直接补 0 后接Embedding再送 LSTM效果也能接受。如果显存够max_len可以设 1200 左右再大收益递减。3. LSTM 模型搭建从 Embedding 到分类头的每一层参数3.1 模型整体结构和前向传播逻辑这份源码的模型部分通常是一个nn.Module包含 Embedding 层、LSTM 层、全连接层。Embedding 把 id 序列映射成稠密向量LSTM 提取序列特征取最后一个时间步的输出或做平均池化再经过全连接映射到 14 类。我一般会在 LSTM 后加一个 Dropout防止过拟合。下面是一个可直接用的模型定义。import torch import torch.nn as nn class LSTMClassifier(nn.Module): def __init__(self, vocab_size, embed_dim300, hidden_dim256, num_layers2, num_classes14, dropout0.3): super().__init__() # padding_idx0 表示 0 对应的 embedding 不参与梯度更新 self.embedding nn.Embedding(vocab_size, embed_dim, padding_idx0) self.lstm nn.LSTM(embed_dim, hidden_dim, num_layers, batch_firstTrue, bidirectionalTrue, dropoutdropout if num_layers 1 else 0) self.dropout nn.Dropout(dropout) self.fc nn.Linear(hidden_dim * 2, num_classes) # 双向所以乘 2 def forward(self, x): # x: (batch, seq_len) emb self.embedding(x) # (batch, seq_len, embed_dim) out, (h, c) self.lstm(emb) # out: (batch, seq_len, hidden*2) # 取最后一个时间步的输出 out out[:, -1, :] out self.dropout(out) logits self.fc(out) return logitspadding_idx0让 padding 的 embedding 保持为 0 且不更新减少噪声。bidirectionalTrue让 LSTM 同时看前后文对新闻分类这种任务通常比单向高 1 到 2 个点。num_layers2时 LSTM 内部会加 dropout但只有层数大于 1 才生效。hidden_dim * 2是因为双向拼接。取out[:, -1, :]是取最后一个时间步如果序列补了很多 0最后一个时间步可能是 padding这时改用平均池化更稳。3.2 关键参数怎么设embed_dim、hidden_dim、num_layers这三个参数直接决定模型容量和训练速度。embed_dim常见 128、256、300天池这个赛题用 300 预训练词向量初始化效果更好但源码里如果没提供预训练向量从随机初始化开始300 和 256 差别不大。hidden_dim我一般设 256双向后输出 512再大显存吃紧且容易过拟合。num_layers设 2 足够3 层以上训练慢且提升有限。参数常用值影响embed_dim128 / 256 / 300太小欠拟合太大过拟合hidden_dim128 / 256 / 512256 是精度和速度的平衡点num_layers1 / 2 / 32 层性价比最高dropout0.2 / 0.3 / 0.5过拟合严重时调到 0.5max_len800 / 1000 / 1200覆盖 95% 文本长度即可如果训练集 loss 降得很快但验证集不降先把dropout调到 0.5再把hidden_dim降到 128。如果训练集 loss 都降不下去检查词表是不是太小或者max_len截断太狠。3.3 用 pack_padded_sequence 处理变长序列的正确姿势上面模型对 padding 位置也做了 LSTM 计算虽然padding_idx0减少了 embedding 噪声但 LSTM 仍会在 padding 上消耗计算。更规范的做法是用pack_padded_sequence只对真实长度做 LSTM。这需要先按长度排序再 pack最后 unpack。from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence def forward_with_pack(self, x, lengths): emb self.embedding(x) # 按长度降序排序pack 要求 lengths, sort_idx lengths.sort(descendingTrue) emb emb[sort_idx] packed pack_padded_sequence(emb, lengths.cpu(), batch_firstTrue) packed_out, (h, c) self.lstm(packed) out, _ pad_packed_sequence(packed_out, batch_firstTrue) # 还原原始顺序 _, unsort_idx sort_idx.sort() out out[unsort_idx] out out[:, -1, :] return self.fc(self.dropout(out))lengths必须是 CPU 上的 int64 张量。排序后要记住sort_idx最后用unsort_idx还原否则 batch 内顺序错乱loss 会异常。如果嫌麻烦直接补 0 送 LSTM 也能跑但显存占用会高 20% 左右。4. 训练循环与验证让 loss 真正降下来的几个开关4.1 数据加载和 batch 划分训练循环第一步是把 id 序列转成 Tensor用DataLoader分 batch。天池数据量大batch_size设 64 或 128 比较稳。太小训练慢太大显存不够且梯度更新次数少。from torch.utils.data import Dataset, DataLoader class NewsDataset(Dataset): def __init__(self, ids, labelsNone): self.ids ids self.labels labels def __len__(self): return len(self.ids) def __getitem__(self, idx): ids torch.tensor(self.ids[idx], dtypetorch.long) if self.labels is not None: label torch.tensor(self.labels[idx], dtypetorch.long) return ids, label return ids train_dataset NewsDataset(train_df[ids].tolist(), train_df[label].tolist()) train_loader DataLoader(train_dataset, batch_size128, shuffleTrue, num_workers2)num_workers2在 Linux 上能加速数据读取Windows 上如果报错就设 0。shuffleTrue只在训练集开验证集和测试集不要开。4.2 优化器、学习率和损失函数的选择优化器我一般用 Adam学习率 1e-3配合ReduceLROnPlateau在验证集 loss 不降时减半。损失函数用CrossEntropyLoss如果类别不均衡可以加weight参数。天池这个赛题类别基本均衡不加权重也行。import torch.optim as optim device torch.device(cuda if torch.cuda.is_available() else cpu) model LSTMClassifier(vocab_sizelen(vocab)).to(device) criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr1e-3) scheduler optim.lr_scheduler.ReduceLROnPlateau(optimizer, modemin, factor0.5, patience2) for epoch in range(10): model.train() total_loss 0 for ids, labels in train_loader: ids, labels ids.to(device), labels.to(device) optimizer.zero_grad() logits model(ids) loss criterion(logits, labels) loss.backward() # 梯度裁剪防止 LSTM 梯度爆炸 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0) optimizer.step() total_loss loss.item() print(fEpoch {epoch}, Loss: {total_loss / len(train_loader):.4f})clip_grad_norm_的max_norm5.0是 LSTM 训练的后悔药不加的话 loss 可能突然变 NaN。patience2表示验证集 loss 连续 2 个 epoch 不降就减学习率。如果显存够batch_size可以加到 256学习率相应调到 2e-3。4.3 验证集划分和早停策略训练集不能全用来训练要切 10% 做验证。早停是防止过拟合最直接的手段验证集 loss 连续 3 个 epoch 不降就停。from sklearn.model_selection import train_test_split train_ids, val_ids, train_labels, val_labels train_test_split( train_df[ids].tolist(), train_df[label].tolist(), test_size0.1, random_state42) val_dataset NewsDataset(val_ids, val_labels) val_loader DataLoader(val_dataset, batch_size128, shuffleFalse) best_val_loss float(inf) patience_counter 0 for epoch in range(20): # 训练部分省略同上 model.eval() val_loss 0 with torch.no_grad(): for ids, labels in val_loader: ids, labels ids.to(device), labels.to(device) logits model(ids) val_loss criterion(logits, labels).item() val_loss / len(val_loader) scheduler.step(val_loss) if val_loss best_val_loss: best_val_loss val_loss torch.save(model.state_dict(), best_model.pt) patience_counter 0 else: patience_counter 1 if patience_counter 3: print(早停触发) breakrandom_state42保证每次切分一致方便对比实验。保存best_model.pt而不是最后一个 epoch 的模型因为最后一个可能已经过拟合。早停的patience设 3 比较稳设 1 容易误停。5. 避坑与排查LSTM 新闻分类里最容易翻车的 5 个地方5.1 现象loss 一直是 NaN训练无法继续原因通常是学习率太大或梯度爆炸。LSTM 的梯度回传路径长不加裁剪很容易爆。解决把学习率降到 1e-4加上clip_grad_norm_(model.parameters(), max_norm5.0)。如果还是 NaN检查输入 id 里有没有超出词表大小的值nn.Embedding遇到越界 id 会直接报错或产生 NaN。5.2 现象训练集准确率 99%验证集只有 60%这是典型过拟合。原因可能是模型太大、训练轮数太多、dropout 太小。解决把hidden_dim从 256 降到 128dropout从 0.3 提到 0.5加早停。另外检查验证集是不是和训练集同分布天池的测试集和训练集分布基本一致如果验证集掉点严重多半是过拟合而不是分布问题。5.3 现象报错RuntimeError: Expected hidden[0] size (2, 128, 256), got (2, 128, 128)这是双向 LSTM 的 hidden 维度没对齐。hidden_dim设 256双向后输出 512但初始化 hidden 时如果按 256 写就会报这个错。解决不手动初始化 hidden让 LSTM 自己初始化或者把 hidden 的维度写成(num_layers * 2, batch, hidden_dim)。我一般直接不传 hidden省事。5.4 现象词表建完发现UNK占比超过 20%说明词表太小或者过滤太狠。原因可能是count 2把太多词过滤了或者max_len截断导致长尾词没统计到。解决把词表上限提到 20 万count 2改成count 1即保留所有词。但词表太大会让 Embedding 参数量暴涨需要权衡。另一个办法是用字符级 token但新闻分类里词级通常更好。5.5 现象测试集提交后分数比验证集低很多原因可能是测试集里有些词没在词表里全变UNK。解决建词表时合并测试集文本确保测试集里的词也在词表中。另外检查测试集的 id 转换函数和训练集是否完全一致max_len和截断策略要相同。如果验证集用了pack_padded_sequence而测试集没用也会导致不一致。6. 进阶技巧用预训练词向量和 FGM 对抗训练再提 2 个点6.1 加载预训练词向量初始化 Embedding随机初始化的 Embedding 在 20 万条数据上能学到不错的表示但如果有预训练词向量比如腾讯词向量或 Word2Vec加载后冻结或微调通常能再提 1 到 2 个点。加载逻辑是遍历词表找到对应词向量填进 Embedding 矩阵。import numpy as np def load_pretrained_embedding(vocab, embed_dim300): # 假设预训练词向量存为 dict: word - np.array pretrained {} with open(word_vectors.txt, r, encodingutf-8) as f: for line in f: parts line.strip().split() word parts[0] vec np.array([float(x) for x in parts[1:]]) pretrained[word] vec embedding_matrix np.random.normal(0, 0.1, (len(vocab), embed_dim)) hit 0 for word, idx in vocab.items(): if word in pretrained: embedding_matrix[idx] pretrained[word] hit 1 print(f命中预训练词: {hit}/{len(vocab)}) return embedding_matrix # 赋值给 Embedding 层 embedding_matrix load_pretrained_embedding(vocab) model.embedding.weight.data.copy_(torch.tensor(embedding_matrix, dtypetorch.float)) # 可以选择冻结 embedding # model.embedding.weight.requires_grad False命中率低于 50% 说明预训练词向量和当前语料领域差异大效果可能不明显。冻结 Embedding 适合数据量小的情况数据量大时微调更好。6.2 FGM 对抗训练在 LSTM 上的实现FGM 是一种对抗训练方法通过在 Embedding 层加扰动让模型对噪声更鲁棒。实现上是在每次loss.backward()后对 Embedding 权重加一个梯度方向的扰动再算一次 loss 并更新。class FGM: def __init__(self, model, epsilon1.0): self.model model self.epsilon epsilon self.backup {} def attack(self): for name, param in self.model.named_parameters(): if param.requires_grad and embedding in name: self.backup[name] param.data.clone() norm torch.norm(param.grad) if norm ! 0: r_at self.epsilon * param.grad / norm param.data.add_(r_at) def restore(self): for name, param in self.model.named_parameters(): if name in self.backup: param.data self.backup[name] self.backup {} # 训练循环里使用 fgm FGM(model, epsilon1.0) for ids, labels in train_loader: optimizer.zero_grad() logits model(ids) loss criterion(logits, labels) loss.backward() fgm.attack() # 加扰动 logits_adv model(ids) loss_adv criterion(logits_adv, labels) loss_adv.backward() fgm.restore() # 恢复 optimizer.step()epsilon1.0是扰动幅度太大训练不稳定太小没效果。FGM 只对 Embedding 层加扰动因为 LSTM 层加扰动计算量大且收益低。加了 FGM 后训练时间增加约 30%但验证集通常能提 0.5 到 1 个点。6.3 验证提分是否真实的交叉验证方法单次验证集划分有随机性可能这次提了下次没提。我一般用 5 折交叉验证每折都跑一遍看平均分和方差。如果 FGM 在 5 折里平均提了 0.8 个点方差不大那才是真提分。from sklearn.model_selection import KFold kf KFold(n_splits5, shuffleTrue, random_state42) scores [] for fold, (tr_idx, va_idx) in enumerate(kf.split(train_df)): tr_ids train_df[ids].iloc[tr_idx].tolist() tr_labels train_df[label].iloc[tr_idx].tolist() va_ids train_df[ids].iloc[va_idx].tolist() va_labels train_df[label].iloc[va_idx].tolist() # 训练模型并评估记录验证集准确率 # scores.append(acc) print(f5 折平均准确率: {np.mean(scores):.4f}, 标准差: {np.std(scores):.4f})标准差超过 0.5 个点说明模型不稳定需要检查数据划分或模型初始化。我自己的习惯是每次改完模型先跑 5 折确认稳定后再上全量数据训练提交。这套流程跑下来LSTM 在天池新闻文本分类上做到 0.92 到 0.94 的准确率是正常的再往上就要换 BERT 了。希望帮到你。本文还有配套的精品资源点击获取

相关推荐

agent-skills 实战指南:为 AI 编程助手构建可复用技能模块
agent-skills 实战指南:为 AI 编程助手构建可复用技能模块

1. 从零认识 agent-skills:它到底解决了什么问题第一次看到agent-skills这个词,很多人会以为是某个新出的 AI 模型或者插件市场。其实不是。它更像是一套给 AI coding agent 准备的“技能包规范”——你可以把它理解成给 AI 编程助手写的“操作手册 工具… · 2026/9/23 7:53:11

从普通Prompt到思维链CoT:数学推理提示词实战指南
从普通Prompt到思维链CoT:数学推理提示词实战指南

如果你只把提示词工程当成“把问题描述得更清楚”,那你大概率会在数学推理上碰一鼻子灰。我最近做了一组很简单的对比测试:让同一个大模型计算一道包含加权平均和混合运算的数学题,用普通的 Prompt 直接问,它一本正经地给出了一个… · 2026/9/23 7:53:11

从AI Coding到AI Engineering:16万行代码的工程化实践
从AI Coding到AI Engineering:16万行代码的工程化实践

1. 项目背景:16 万行代码,从“手写”到“AI 协同”的转折点先交代一下背景。这个项目是一个中大型业务系统,涵盖管理后台、用户端 API、定时任务、消息推送、数据对账等多个模块。按传统开发方式估算,16 万行代码大概是一个 6 到 … · 2026/9/23 7:53:04

英语偏旁部首入门到精通:揭秘代码里的字符拆解逻辑
英语偏旁部首入门到精通:揭秘代码里的字符拆解逻辑

英语偏旁部首入门到精通:揭秘代码里的字符拆解逻辑 复制来的代码跑不通,报错信息满屏红字,你盯着屏幕抓耳挠腮,根本不知道从哪下手调。这种“黑盒”体验,是每个开发者从新手迈向 入门到精通… · 2026/9/23 8:37:19

vray渲染器踩坑实录
vray渲染器踩坑实录

V-Ray渲染器性能优化避坑:3个让出图慢10倍的致命错误 复制来的V-Ray渲染参数跑不通,或者跑出来的图黑乎乎一片、噪点满天飞,是不是让你抓狂?别急,这通常是场景设置和硬件配置的冲突,不是你的错。很多新手卡在第一步,因为直接套用网上通用… · 2026/9/23 8:37:19

无线运动耳机性能优化实战:告别堆栈报错
无线运动耳机性能优化实战:告别堆栈报错

无线运动耳机性能优化实战:告别堆栈报错 盯着满屏红色的StackTrace,眼睛都花了还是找不到Bug在哪?别急,这行代码没报错,但你的无线运动耳机在剧烈运动时音频断连、延迟高企,这才是真正的“性能优化”噩梦。很多开发者一上来就调参数,结果… · 2026/9/23 8:36:54

FPGA进位链实现高精度TDC的原理与工程实践
FPGA进位链实现高精度TDC的原理与工程实践

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views … · 2026/9/23 8:36:47

yfd 入门到精通:3 步搞定 StackTrace 报错与底层原理
yfd 入门到精通:3 步搞定 StackTrace 报错与底层原理

yfd 入门到精通:3 步搞定 StackTrace 报错与底层原理 面对满屏红色的 StackTrace,你是不是只想把电脑摔了?别急,这不仅是你的噩梦,也是所有开发者从入门到精通必须跨越的坎。yfd… · 2026/9/23 8:36:47

Link Park避坑指南:从报错到精通的保姆级教程
Link Park避坑指南:从报错到精通的保姆级教程

Link Park避坑指南:从报错到精通的保姆级教程 刚接完一个 Link Park 相关的后端需求,测试环境跑起来,日志直接吐了满屏的 java.lang.NullPointerException 和… · 2026/9/23 8:36:47

3招搞定手机怎么下载微信面试难题实战项目解析
3招搞定手机怎么下载微信面试难题实战项目解析

3招搞定手机怎么下载微信面试难题实战项目解析 面试被问“手机怎么下载微信”背后的原理,90%的人答不上来。别笑,这看似弱智的问题,实则是考察你对移动应用分发机制、安全校验及网络协议理解的试金石。我带过不少校招新人,他们背了八股文,却连一个A… · 2026/9/23 0:00:03

你有新短消息请注意查收:3个新手避坑指南搞定消息系统选型
你有新短消息请注意查收:3个新手避坑指南搞定消息系统选型

你有新短消息请注意查收:3个新手避坑指南搞定消息系统选型 面试被问“高并发下如何保证消息不丢失”,你张口就是“用Redis”,结果面试官追问“如果Redis宕机了怎么办”,你瞬间卡壳。这种场景太常见了,很多新手在背八股文时,只记住了技术名词… · 2026/9/23 0:00:29

Win7无线热点配置工具源码解析:解决API失效的3个实战技巧
Win7无线热点配置工具源码解析:解决API失效的3个实战技巧

Win7无线热点配置工具源码解析:解决API失效的3个实战技巧 Win7无线热点配置工具在Win10/11上跑不动?不是你的问题,是版本升级后 API 全变了。很多老项目里的 netsh wlan… · 2026/9/23 0:00:36

了解更多?预约专属演示

我们的顾问将为您一对一讲解产品与方案

企业微信二维码