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_T Linear(1) 打分 Sigmoid 0~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 章里那个"浓缩了全文"的隐藏状态。


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

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

python 复制代码
import torch.nn as nn

class SentimentGRU(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)                       # 打分

    def forward(self, x):                 # x: (B, S) 整数矩阵
        e = self.emb(x)                   # (B, S, E) 查表成向量
        _, h = self.gru(e)                # h: (1, B, H) 最后一步记忆
        return self.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 内部到底怎么"通读",用一段最直白的循环讲,比看张量拼起来更清楚:

python 复制代码
h = torch.zeros(hidden)          # 记忆清零,开始读
for word_vec in review_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)\\,\\BigL=−N1i=1∑Nyilogy\^i+(1−yi)log(1−y\^i)

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

python 复制代码
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()

for ep in range(20):
    for st in range(0, 200, 32):
        x, y = pack_batch(idx[st:st+32])            # 取一批,填充+掩码
        lo = lossf(model(x), y)
        opt.zero_grad(); lo.backward(); opt.step()

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

前几轮的真实数字:

训练轮 0 2 4 6 8 10
损失 0.6996 0.5121 0.2426 0.1666 0.1359 0.0830

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


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

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

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

python 复制代码
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.6944 0.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.054 25 轮下降 0.640 验证准确率 50% → 73.85% 第 21 轮峰值后走平 随机基线 0.693 / 50% 没学会的分界线

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

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

相关推荐
阳明山水1 小时前
销量预测的“隐形杀手”:概念漂移与自适应学习实战
人工智能·深度学习·算法·机器学习·架构
ZHOU_WUYI10 小时前
7. fastwam 模型 _predict_action_noise_with_cache部分
人工智能·pytorch·深度学习
华清远见成都中心11 小时前
卷积神经网络(CNN)为什么能够识别图像?
人工智能·深度学习·cnn
G311354227312 小时前
大模型不可用时,业务还能不能继续:企业需要设计降级方案
大数据·服务器·数据库·人工智能·深度学习
TechEdu20260612 小时前
[人工智能]TensorFlow深度学习框架工程实践概览
人工智能·深度学习·ai·tensorflow
一碗白开水一15 小时前
入门实践工程九:基于 BERT 的中文情感分类微调~附:安装依赖库及工程源码
人工智能·深度学习·机器学习·自然语言处理·分类·bert
xiaoxiaoxiaolll15 小时前
AI赋能复合材料力学:神经网络与” 多尺度仿真
人工智能·深度学习·神经网络
卡梅德生物科技小能手18 小时前
卡梅德生物科普 TPBG(滋养层糖蛋白)
经验分享·深度学习·生活
淼澄研学20 小时前
PyTorch模型训练5大避坑指南:解决显存泄漏与设备不匹配报错
人工智能·pytorch·python