云计算百科
云计算领域专业知识百科平台

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

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

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

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

🎯 本章目标

  • 拼出 Embedding → GRU → Linear(1) → Sigmoid 的完整网络;
  • 用 200 条小样本验证代码正确(过拟合 = 实现没写错);
  • 在 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 章里那个"浓缩了全文"的隐藏状态。


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

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

    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 内部到底怎么"通读",用一段最直白的循环讲,比看张量拼起来更清楚:

    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)\\,\\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()

    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 轮,损失曲线长这样:

    在这里插入图片描述

    前几轮的真实数字:

    训练轮0246810
    损失 0.6996 0.5121 0.2426 0.1666 0.1359 0.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.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 三模型横向对比实验

    赞(0)
    未经允许不得转载:网硕互联帮助中心 » python神经网络编程入门(二十七)——RNN IMBD搭建情感分类器与基础训练
    分享到: 更多 (0)

    评论 抢沙发

    评论前必须登录!