
1. 从零搭建框架到对话预测为什么我要自己重写RNN预测网络很多人学深度学习第一步就是import torch或者import tensorflow然后调几个API把模型跑通觉得自己已经“入门”了。但真到了要自己动手写一个能用的预测网络时问题就全冒出来了——张量维度对不上、梯度传着传着就没了、训练loss不降反升、预测出来的对话驴唇不对马嘴。这些坑光靠调库是永远踩不到的。我前面已经带着大家从零搭了一套深度学习框架包括张量、自动求导、基础层、优化器这些核心组件。现在到了第13篇要干一件真正有挑战的事用我们自己建的架构重写一个RNN预测网络实现对话预测功能。说白了就是让模型学会“你问一句它答一句”这种最基础的序列到序列的映射能力。为什么选对话预测作为RNN的实战场景因为对话数据天然就是序列一问一答之间有时序依赖而且长度可变、语义复杂非常适合拿来检验一个RNN网络到底有没有真正跑通。更重要的是这个任务足够小小到你能在一台普通笔记本上跑完但又足够真实真实到你能从中理解RNN的核心机制——隐藏状态怎么传、梯度怎么回传、序列怎么对齐。这篇文章我会把整个程序的每一行关键代码都拆开讲清楚包括数据怎么准备、网络怎么搭、前向传播怎么走、损失怎么算、反向传播怎么实现、预测阶段怎么解码。如果你跟着我前面的文章一路走过来这篇会让你对整个框架的理解上一个台阶如果你是直接跳到这篇也没关系我会把必要的背景补全让你能独立复现。提示本文假设你已经有了一个能跑通基础张量运算和自动求导的框架。如果你还没有建议先回头看看前面的内容否则后面的代码你会看得云里雾里。2. 对话预测任务的数据准备从原始文本到可训练张量2.1 对话数据的特殊性在哪里对话预测和普通的文本分类、情感分析不一样。分类任务是把一整段文本映射到一个标签输入输出都是定长的。但对话是一对一的序列映射输入是一个问句序列输出是一个答句序列两个序列长度都可能不一样而且每个时间步的输出都依赖于前面所有时间步的输入。这就带来几个必须解决的问题。第一变长序列怎么处理。你不能要求所有问句都是5个词、所有答句都是7个词那太死板了。常见的做法是设定一个最大长度短的补零长的截断。第二词表怎么建。对话数据里词汇量可能很大但很多词只出现一两次如果全放进词表模型参数会爆炸。所以需要做词频过滤低频词统一用特殊标记代替。第三输入和输出怎么对齐。在训练时我们通常把答句的最后一个词作为目标前面的词作为解码器的输入这叫“teacher forcing”策略。我这次用的是一份小型的客服对话数据集大概两千多组问答对内容涉及退换货、物流查询、支付问题这些场景。数据量不大但足够验证网络结构是否正确。2.2 词表构建与序列编码的实操细节词表构建这一步很多人会直接调现成的分词器但既然我们是自己搭框架就得自己实现。我的做法是先把所有对话文本按字符切分中文场景下字符级比词级更稳妥不用考虑分词歧义然后统计每个字符的出现频率。def build_vocab(sentences, min_freq2): freq {} for sent in sentences: for ch in sent: freq[ch] freq.get(ch, 0) 1 # 特殊标记 vocab {pad: 0, sos: 1, eos: 2, unk: 3} idx 4 for ch, count in sorted(freq.items(), keylambda x: -x[1]): if count min_freq: vocab[ch] idx idx 1 return vocab这里有几个细节值得说。pad的索引设为0后面做mask的时候直接用tensor ! 0就能筛出有效位置。sos和eos分别标记序列的开始和结束解码时靠它们控制生成过程。unk处理低频字符防止词表过大。编码的时候每个问句前面加sos、后面加eos答句也一样。然后统一padding到最大长度。这里有个坑padding的位置会影响RNN的隐藏状态更新。如果你直接让RNN处理padding的0向量隐藏状态会被“污染”。所以要么用pack_padded_sequence这种机制跳过padding要么在计算损失时把padding位置的损失mask掉。我选择后者因为自己实现pack机制太复杂而mask损失更直观。def encode(sentence, vocab, max_len): ids [vocab.get(sos)] [vocab.get(ch, vocab[unk]) for ch in sentence] [vocab.get(eos)] if len(ids) max_len: ids ids[:max_len] else: ids ids [vocab[pad]] * (max_len - len(ids)) return ids注意padding的长度建议取所有样本长度的95分位数而不是最大值。最大值可能是个别超长样本会导致大量padding浪费计算资源。2.3 批处理与数据加载器的实现自己搭框架数据加载器也得自己写。核心逻辑就是每次取一个batch把问句和答句分别堆成矩阵。问句矩阵形状是(batch_size, max_q_len)答句矩阵形状是(batch_size, max_a_len)。class DialogDataset: def __init__(self, pairs, vocab, max_q_len, max_a_len): self.q_data [encode(q, vocab, max_q_len) for q, a in pairs] self.a_data [encode(a, vocab, max_a_len) for q, a in pairs] self.batch_size 32 def __len__(self): return len(self.q_data) // self.batch_size def __getitem__(self, idx): start idx * self.batch_size end start self.batch_size q_batch self.q_data[start:end] a_batch self.a_data[start:end] return q_batch, a_batch这里我故意没有用随机打乱因为对话数据里有些样本是有关联的打乱反而可能破坏上下文。当然如果你的数据是独立同分布的打乱没问题。3. 用自建框架重写RNN单元从公式到代码的完整映射3.1 RNN的核心公式与隐藏状态传递机制RNN的本质就一个公式$$h_t \tanh(W_{ih} x_t b_{ih} W_{hh} h_{t-1} b_{hh})$$其中$x_t$是当前时间步的输入$h_{t-1}$是上一时间步的隐藏状态$W_{ih}$和$W_{hh}$是两组权重矩阵$b$是偏置。输出层再做一个线性变换$$y_t W_{ho} h_t b_o$$这个公式看起来简单但自己实现的时候维度对齐是最容易出错的地方。假设词嵌入维度是128隐藏层维度是256batch_size是32。那么$x_t$的形状是(32, 128)$h_{t-1}$的形状是(32, 256)。$W_{ih}$必须是(128, 256)$W_{hh}$必须是(256, 256)这样两个矩阵乘法的结果才能相加。我在第一次实现的时候把$W_{hh}$写成了(256, 128)结果矩阵乘法直接报维度不匹配。排查了半天才发现是转置搞反了。所以你在写的时候一定要把每个矩阵的形状在纸上画出来确认无误再写代码。3.2 在自建框架中实现RNNCell类基于我们自己的张量类RNNCell的实现大概长这样class RNNCell: def __init__(self, input_size, hidden_size): self.W_ih Tensor.randn(input_size, hidden_size) * 0.01 self.W_hh Tensor.randn(hidden_size, hidden_size) * 0.01 self.b_ih Tensor.zeros(1, hidden_size) self.b_hh Tensor.zeros(1, hidden_size) self.params [self.W_ih, self.W_hh, self.b_ih, self.b_hh] def forward(self, x, h_prev): # x: (batch, input_size), h_prev: (batch, hidden_size) linear1 x.matmul(self.W_ih) self.b_ih linear2 h_prev.matmul(self.W_hh) self.b_hh h_next (linear1 linear2).tanh() return h_next这里的关键是matmul和tanh都必须是我们框架里实现了自动求导的操作。如果tanh没有实现反向传播那梯度传到这里就断了训练根本没法进行。我在框架里给tanh写的反向传播是grad * (1 - tanh(x)^2)这是标准公式但要注意tanh(x)的值在前向传播时已经算出来了反向传播时直接复用不要重新算一遍否则计算图会出问题。权重初始化用0.01倍的标准正态分布这是防止梯度爆炸的常用手段。如果你用1.0倍初始隐藏状态会很大tanh直接饱和梯度接近零模型学不动。3.3 序列展开从单步到多步的循环逻辑有了RNNCell接下来要把它按时间步展开。假设输入序列长度是20就要循环20次每次把当前时间步的输入和上一步的隐藏状态喂进去。class RNN: def __init__(self, input_size, hidden_size, output_size): self.cell RNNCell(input_size, hidden_size) self.W_ho Tensor.randn(hidden_size, output_size) * 0.01 self.b_o Tensor.zeros(1, output_size) self.params self.cell.params [self.W_ho, self.b_o] def forward(self, x_seq, h0None): # x_seq: (batch, seq_len, input_size) batch_size, seq_len, _ x_seq.shape if h0 is None: h0 Tensor.zeros(batch_size, self.cell.W_hh.shape[0]) h h0 outputs [] for t in range(seq_len): h self.cell.forward(x_seq[:, t, :], h) y h.matmul(self.W_ho) self.b_o outputs.append(y) return outputs, h这里有个性能问题每次循环都创建一个新的计算图节点序列长了之后计算图会非常大反向传播时内存占用很高。我在实测中发现序列长度超过50之后训练速度明显下降。解决办法是限制最大序列长度或者用截断反向传播只回传最近N步的梯度。对于对话预测这种任务20到30的长度基本够用。提示如果你发现训练时内存暴涨先检查是不是序列太长导致计算图过大。可以打印一下计算图的节点数量超过一万个节点就要考虑优化了。4. 编码器-解码器架构的搭建与训练流程4.1 为什么对话预测需要编码器-解码器结构最简单的RNN只能做一对一映射但对话是一对多的序列映射。你输入一个问句模型要输出一个完整的答句而且答句的每个词都依赖于问句的语义和前面已经生成的词。这就需要一个编码器先把问句压缩成一个上下文向量然后解码器基于这个向量逐步生成答句。编码器就是一个普通的RNN把问句的每个词依次喂进去最后一步的隐藏状态就是整个问句的语义表示。解码器也是RNN但它的初始隐藏状态来自编码器每一步的输入是上一个时间步生成的词训练时是真实答句的上一个词输出是当前时间步的词概率分布。这种结构有个名字叫“序列到序列”Seq2Seq是对话预测、机器翻译这些任务的经典架构。虽然现在Transformer大行其道但RNN版本的Seq2Seq依然是理解序列建模的最佳入口。4.2 编码器与解码器的具体实现编码器直接复用上面的RNN类但只需要最后一步的隐藏状态class Encoder: def __init__(self, vocab_size, embed_size, hidden_size): self.embed Tensor.randn(vocab_size, embed_size) * 0.01 self.rnn RNN(embed_size, hidden_size, hidden_size) self.params [self.embed] self.rnn.params def forward(self, x): # x: (batch, seq_len) 整数索引 embedded self.embed[x] # (batch, seq_len, embed_size) _, h_final self.rnn.forward(embedded) return h_final解码器稍微复杂一点因为每一步都要输出词表上的概率分布class Decoder: def __init__(self, vocab_size, embed_size, hidden_size): self.embed Tensor.randn(vocab_size, embed_size) * 0.01 self.rnn RNN(embed_size, hidden_size, vocab_size) self.params [self.embed] self.rnn.params def forward(self, x, h0): embedded self.embed[x] outputs, h_final self.rnn.forward(embedded, h0) return outputs, h_final注意解码器的输出维度是vocab_size因为每个时间步都要预测下一个词是词表里的哪一个。损失函数用交叉熵把每个时间步的输出和真实标签做对比。4.3 训练循环中的损失计算与梯度更新训练循环的伪代码大概是这样for epoch in range(num_epochs): for q_batch, a_batch in dataset: # 前向传播 h_enc encoder.forward(q_batch) # 解码器输入答句去掉最后一个词前面加sos dec_input a_batch[:, :-1] dec_target a_batch[:, 1:] outputs, _ decoder.forward(dec_input, h_enc) # 计算损失 loss 0 for t in range(len(outputs)): loss cross_entropy(outputs[t], dec_target[:, t]) # 反向传播 loss.backward() # 更新参数 optimizer.step() optimizer.zero_grad()这里有几个关键点。第一解码器的输入和目标是错开一位的这叫“shifted target”。第二损失要对每个时间步求和或平均我一般用平均这样不同长度的序列对损失的贡献更均衡。第三zero_grad必须在step之后调用否则梯度会累积。我在第一次跑的时候忘了zero_grad结果loss直接飞到NaN。排查了好久才发现是梯度累积导致的。这个坑很隐蔽因为前几个batch看起来正常到后面才爆炸。注意如果你的loss出现NaN优先检查三件事梯度有没有清零、学习率是不是太大、有没有除零操作。这三个原因占了NaN问题的九成以上。5. 对话预测的解码策略与效果验证5.1 贪心解码与束搜索的取舍训练完之后怎么用模型生成答句最简单的是贪心解码每一步都选概率最大的词直到生成eos或者达到最大长度。def greedy_decode(encoder, decoder, question, vocab, max_len30): h encoder.forward(question) input_id vocab[sos] result [] for _ in range(max_len): x Tensor([[input_id]]) output, h decoder.forward(x, h) probs softmax(output[0]) input_id argmax(probs) if input_id vocab[eos]: break result.append(input_id) return result贪心解码的问题是容易陷入局部最优生成的句子可能不通顺。束搜索beam search保留top-k个候选每一步扩展后再选最优的k个效果通常更好。但束搜索的计算量是贪心的k倍对于实时对话场景k一般取3到5。我实测下来在这个小数据集上贪心解码和束搜索的差距不明显因为数据量小、句子短。但如果你的数据复杂束搜索值得一试。5.2 预测效果的评估与常见问题分析评估对话预测的质量不能只看loss。loss低不代表生成的句子合理。我一般从三个维度看语法正确性生成的句子是不是人话、语义相关性答句和问句有没有关系、多样性是不是所有问句都回同一句话。常见问题有几个。第一模型倾向于生成高频词比如“好的”“谢谢”因为这样loss最低。解决办法是在损失里加一个频率惩罚项降低高频词的权重。第二模型可能重复生成同一个词比如“我我我我”这是因为隐藏状态陷入了循环。可以在解码时加一个重复惩罚对已经生成的词降低其概率。第三长问句的答句质量明显下降因为编码器把长序列压缩成一个固定维度的向量信息损失严重。这个问题在RNN架构下没有完美解法只能靠增加隐藏层维度或者用注意力机制缓解。5.3 从训练曲线判断模型是否真正学到了东西训练过程中我会同时记录训练loss和验证loss。如果训练loss持续下降但验证loss开始上升说明过拟合了需要加正则化或者减少参数量。如果两个loss都不降说明模型没学到东西可能是学习率太小、梯度消失、或者数据有问题。梯度消失是RNN的经典问题。当序列较长时反向传播的梯度会指数衰减前面的时间步几乎收不到梯度。判断方法很简单打印每一层的梯度范数如果前面几层的梯度接近零就是梯度消失了。解决办法包括用LSTM/GRU替代普通RNN、用梯度裁剪、或者缩短序列长度。我在这个项目里用的是普通RNN序列长度控制在20以内梯度消失问题不严重。但如果你要处理更长的序列强烈建议换成LSTM。LSTM的门控机制能有效缓解梯度消失代价是参数量增加、计算变慢。6. 自己搭框架重写RNN的几点实战体会走完这一整套流程我最大的感受是自己实现一遍比看十遍公式都管用。以前调库的时候nn.RNN一行代码就搞定了但隐藏状态怎么传、梯度怎么回传、padding怎么处理全是黑盒。自己写完之后这些细节全都清清楚楚。另一个体会是维度对齐是最大的坑。矩阵乘法、广播、拼接每一步都要确认形状。我的习惯是在每个关键操作后面打印形状虽然看起来笨但能省下大量调试时间。还有一点小数据集上不要追求复杂模型。我这个对话预测网络只有两层RNN参数量不到一百万在两千条数据上跑几十个epoch就能收敛。如果你一上来就堆多层LSTM、加注意力、加残差连接反而容易过拟合而且调试难度成倍增加。最后分享一个实用技巧在训练前先用一个极小的数据集比如10条样本跑通整个流程。确认前向传播、反向传播、参数更新、预测解码都能跑通之后再换全量数据。这样能把数据问题和代码问题分开排查效率高很多。我每次写新网络都是这个流程屡试不爽。