当前位置: 首页 > news >正文

python神经网络编程入门(二十七)——RNN IMBD搭建情感分类器与基础训练

引言:菜都切好了,开火烧菜

前两章一直在"备菜":第 12 章把文字变成整数,第 13 章把长短不一的影评装进(50000,500)(50000, 500)(50000,500)的统一模具。数据洗得干干净净,词表、批次、掩码都备好了。可光有食材摆着,成不了菜——得生火、下锅、翻炒。

这一章就是把"食材"真正倒进锅里:搭起一张最简单的情感分类器,让模型把一条影评读一遍,最后吐出一个 0 到 1 之间的打分,越接近 1 越像好评。先证明这套代码能学会,再在真正的数据上跑起来。

🎯本章目标

  1. 拼出Embedding → GRU → Linear(1) → Sigmoid的完整网络;
  2. 用 200 条小样本验证代码正确(过拟合 = 实现没写错);
  3. 在 8000 条影评上正式训练 25 轮,看懂损失下降、准确率爬升,以及后期的过拟合信号。

一、三件套:查字典、通读全文、拍板打分

先看整个网络长什么样。一条影评从整数序列进来,要经过三层:

Embedding查表(B,S,E)词ID→向量GRU沿时间读(B,S,H)最后一步 h_TLinear(1)打分Sigmoid0~1好感度一条影评 → 一个 0~1 的打分

这三层各有各的活,都能用生活里的动作对上号:

  • Embedding(查字典):整数 ID 只是"词在词表里的编号",编号本身没有含义。查表那一步把编号换成一段稠密向量——就像碰到生词去翻字典,翻到的是这个词的意思。第 11 章讲过,这一段向量是能学习的,训练后语义相近的词会靠得近。
  • GRU(通读全文记重点):第 9 章的主角。它一个词一个词地读,手里捏着一份"记忆",每读一个词就更新一次记忆。读到最后一个词,记忆里就浓缩了整篇影评的要点。
  • Linear(1) + Sigmoid(拍板打分):把最后那份记忆压缩成一个数字,再 sigmoid 压到 0~1 之间,当作好评概率。

整条路用一行公式说清楚:

y^=σ(W hT+b)\hat{y} = \sigma\big(W\, h_T + b\big)y^=σ(WhT+b)

其中hTh_ThT是 GRU 读完最后一个词后的记忆,W,bW, bW,b是最后一层线性变换的参数,σ\sigmaσ是 Sigmoid。hTh_ThT就是第 9 章里那个"浓缩了全文"的隐藏状态。


二、骨架代码:三层拼起来

把上面这张图翻译成代码,寥寥十几行:

importtorch.nnasnnclassSentimentGRU(nn.Module):def__init__(self,vocab=5002,embed=64,hidden=128):super().__init__()self.emb=nn.Embedding(vocab,embed,padding_idx=0)# 查字典self.gru=nn.GRU(embed,hidden,batch_first=True)# 通读self.fc=nn.Linear(hidden,1)# 打分defforward(self,x):# x: (B, S) 整数矩阵e=self.emb(x)# (B, S, E) 查表成向量_,h=self.gru(e)# h: (1, B, H) 最后一步记忆returnself.fc(h[-1]).squeeze(-1)# (B,) 未压缩的打分

有几个细节值得停下来看:

  • vocab=5002:词表大小。第 12 章留了 5000 个高频词,加上<PAD>=0<UNK>=1,一共 5002 个。
  • padding_idx=0:告诉 Embedding,编号 0 是填充位。这样 PAD 会被查成全零向量,等于什么都没读,也就不会污染 GRU 的记忆——这是第 13 章掩码思想在"读序列"这里的落地。第 13 章掩码主要拦的是"逐词预测"的损失;分类任务只取最后一步记忆打分,PAD 用零向量挡住即可,不用再单独算掩码。
  • h[-1]:GRU 返回的hhh形状是(1,B,H)(1, B, H)(1,B,H),第 0 维是层数(这里只有 1 层),h[-1]取的是最后一层、最后一个时间步的隐藏状态,也就是"通读完的记忆"。

GRU 内部到底怎么"通读",用一段最直白的循环讲,比看张量拼起来更清楚:

h=torch.zeros(hidden)# 记忆清零,开始读forword_vecinreview_vecs:# 一个词一个词地读h=gru_step(h,word_vec)# 读一个词,更新一次记忆score=sigmoid(linear(h))# 读完,用最后的记忆拍板

gru_step就是第 9 章那一整套更新门、重置门、候选记忆的公式——代码里写成一行,但心里要装着它是在"逐词翻新记忆"。


三、先拿 200 条试刀:小样本过拟合

代码写完了,怎么知道没写错?先拿一小撮数据试。这是最划算的验错法:挑 200 条影评,让模型反复背。如果代码是对的,200 条很快就能背下来——损失一路掉到接近 0,这就是"过拟合",反而说明实现正确。反过来,如果 200 条都学不动,说明前向或反向有 bug,再大的数据也白搭。

训练用的损失函数是二分类交叉熵(BCE):

L=−1N∑i=1N[ yilog⁡y^i+(1−yi)log⁡(1−y^i) ]\mathcal{L} = -\tfrac{1}{N}\sum_{i=1}^{N}\Big[\,y_i\log\hat{y}_i + (1-y_i)\log(1-\hat{y}_i)\,\Big]L=N1i=1N[yilogy^i+(1yi)log(1y^i)]

模型还没学会、瞎猜时,损失会停在−ln⁡12=ln⁡2≈0.693-\ln\tfrac12 = \ln 2 \approx 0.693ln21=ln20.693——这是"随机猜测"的天然底线,后面对比有没有进步就看它。

torch.manual_seed(42);np.random.seed(42)idx=np.random.choice(25000,200,replace=False)# 随机抽 200 条model=SentimentGRU(vocab=5002,embed=32,hidden=32)opt=torch.optim.Adam(model.parameters(),lr=1e-2)lossf=nn.BCEWithLogitsLoss()forepinrange(20):forstinrange(0,200,32):x,y=pack_batch(idx[st:st+32])# 取一批,填充+掩码lo=lossf(model(x),y)opt.zero_grad();lo.backward();opt.step()

跑 20 轮,损失曲线长这样:

前几轮的真实数字:

训练轮0246810
损失0.69960.51210.24260.16660.13590.0830

0.6996一路掉到0.0830——200 条班子基本被背下来了。这证明前向、反向、损失、更新这一整套链路是通的。可以放心上大菜了。


四、全量开火:洗牌、切分、训练

小样本只是验刀,接下来在真正的数据上训练。这里藏着一个极其容易踩的坑:这份数据是按标签排好序的——前一半是好评(标签 1)、后一半是差评(标签 0)。如果直接拿前 8000 条训练、紧挨着的 2000 条当验证,验证集里就会全是同一类标签,测出来的准确率毫无意义(要么虚高、要么虚低)。

所以第一步必须打乱顺序,再随机切分

all_idx=np.arange(25000)np.random.shuffle(all_idx)# 先洗牌,把好评差评打散train_idx=all_idx[:8000]# 训练:8000 条val_idx=all_idx[8000:10000]# 验证:2000 条

然后挂上 Adam 优化器,用 1e-3 的学习率正式训练 25 轮。每轮结束后在验证集上测一次准确率。真实运行输出(挑几轮展示):

epoch 0: train_loss ≈ 0.6944 | val_acc ≈ 0.5020 epoch 5: train_loss ≈ 0.5955 | val_acc ≈ 0.5200 epoch 10: train_loss ≈ 0.3958 | val_acc ≈ 0.5990 epoch 15: train_loss ≈ 0.1851 | val_acc ≈ 0.6855 epoch 20: train_loss ≈ 0.0931 | val_acc ≈ 0.7350 epoch 24: train_loss ≈ 0.0542 | val_acc ≈ 0.7305

把起止的关键数字收在一张表里,一眼看清幅度:

指标第 0 轮第 24 轮变化
训练损失0.69440.0542↓ 0.640
验证准确率50.2%73.05%↑ 22.9%
  • 训练损失:从0.6944降到0.0542,稳稳离开了0.693的随机基线,模型确实在"读懂"好评差评;
  • 验证准确率:从50.2%一路爬到峰值73.85%(第 21 轮),到第 24 轮微回落到73.05%

练完 25 轮,再拉那两条真实影评看打分,一开头的模型和现在的模型判若两人:

idx=13 label=1 score=0.991 | i enjoyed the night ... one of the better movies of the summer idx=12529 label=0 score=0.003 | i had some ... for the movie since it had a nice star cast ...

great的好评拿到0.991,含terrible的差评只有0.003——这次不仅对,而且非常自信

不过曲线里藏着一个要留意的信号:第 21 轮之后,训练损失还在往下掉(0.085 → 0.054),验证准确率却不涨反微降(73.85% → 73.05%)。这就是标准的过拟合:模型开始把训练集一字不差地背下来,对没见过的数据却帮不上忙。训练损失越低并不代表越好——这正是第 16 章引入 Dropout 等正则化手段的理由。


五、常见坑与自查

  • 不洗牌直接切分:数据前半全是好评、后半全是差评,懒得洗牌会让验证集变成"单一种类",准确率失真。shuffle再切分,这一步不能省。
  • 忘记padding_idx=0:不告诉 Embedding 谁是填充位,PAD 也会被当成普通词参与计算,GRU 的"空气"也读了,记忆被污染。
  • 取错隐藏状态GRU返回的第二个值是(1,B,H)(1, B, H)(1,B,H),不取h[-1]而直接拿去喂线性层,维度对不上会报错,或取到非最后一步的状态。
  • 评估时忘了切eval()模式:训练循环里顺手加的 Dropout 在评估时也必须关掉;这里没有 Dropout,但养成model.eval()的习惯,第 16 章会用上。
  • nn.BCELoss而不是BCEWithLogitsLoss:前者要先把 logit 过 Sigmoid,数值上更易不稳定;后者把 Sigmoid 融进损失里,更稳妥,代码里用的就是它。

小结与预告

这一章把前面备好的数据真正喂进了网络,走通了"读一遍 → 打个分"的完整链路:

  • 三层骨架Embedding(查字典)→GRU(通读记忆)→Linear(1)+Sigmoid(打分),一条影评变成一个 0~1 的打分;
  • 小样本验刀:200 条上损失0.70 → 0.08,证明代码没写错;
  • 洗牌教训:数据按标签排序,必须先shuffle再切分,否则验证集失真;
  • 全量 25 轮:训练损失0.694 → 0.054,验证准确率50% → 73.85%(第 21 轮峰值);此后损失继续降、验证走平,真实验到了过拟合。

本章的核心数据,一张小看板收尾:

小样本 200 条0.70 → 0.08损失背下全部训练损失0.694 → 0.05425 轮下降 0.640验证准确率50% → 73.85%第 21 轮峰值后走平随机基线0.693 / 50%没学会的分界线

路已经通了,接下来就是怎么让模型变聪明——第 15 章把 RNN、GRU、LSTM 三大模型拉到同一张桌子上比个高下,看谁收敛最快、谁最终精度最高。

下一篇(二十八):RNN vs LSTM vs GRU 三模型横向对比实验

http://www.jsqmd.com/news/1355169/

相关文章:

  • 终极指南:如何在macOS上使用BlackHole实现零延迟音频环回
  • js-stellar-sdk错误处理完全手册:解决90%的Stellar开发问题
  • word转图片怎么转?7款PDF格式转换工具实测盘点,免费与官方方法一次说清
  • MongoKit索引优化指南:提升MongoDB查询性能的完整方案
  • 如何在5分钟内为Tailwind项目添加Apple式平滑圆角?Corner Smoothing插件快速上手
  • 如何构建企业级语义层:Cube Core实战架构指南与性能优化策略
  • Ookii.Dialogs.WinForms高级技巧:如何实现Vista风格文件对话框
  • Grapple.nvim项目作用域详解:Git仓库、LSP与自定义作用域配置教程
  • Klock实战教程:如何在Android与iOS项目中集成日期时间功能
  • 解决mechabar常见问题:从依赖安装到主题适配的完整解决方案
  • Fast-DDS:重新定义分布式实时通信的3大技术架构突破
  • 移民中介哪家靠谱?先查这3个资质再签字 - 北极星移民
  • Rocky Linux 9.0 完整安装 containerd(K8s )教程
  • web前端基础到入门——15day
  • 字节跳动算法面试通关指南:2年高频LeetCode题目深度解析与实战策略
  • 为什么选择OpenAlphaDiffract?材料科学研究者不可错过的AI辅助工具
  • 被子植物DNA甲基化预测新突破:a2zchromatin-methylation模型原理解析与代码实现
  • GeoJSON终极指南:解锁地理数据处理的完整工具生态
  • Cap开源录屏工具完整指南:3步快速上手与专业录制方案
  • 深度学习音频降噪实战:DeepFilterNet让声音瞬间清晰的完整指南
  • 如何快速集成Django-photologue?5分钟上手的安装与配置攻略
  • 如何使用Open GPX Tracker:从安装到导出GPX文件的终极教程
  • 碳硅共轭:人类与AI深度协同的哲学与科学架构
  • hackernews-TUI常见问题解决:从安装失败到性能优化的完整方案
  • 终极指南:如何在Pixi.js中轻松集成Live2D模型实现二次元角色动画
  • web前端基础到入门——14day
  • 如何在10分钟内用RVC WebUI实现专业级AI语音克隆?完整免费指南
  • 探索全球29k+机场数据:Airports项目终极指南与实用价值解析
  • 提升用户体验的细节:Corner Smoothing插件在移动端应用的最佳实践
  • Rosetta 国际化库技术解密:298字节背后的架构哲学