为何不同RNN架构在字符级姓名分类任务上性能差异巨大?
自定义RNN性能较差的核心原因
- 梯度被人为截断,丧失时序建模能力
你的实现中self.prev_hidden = hidden.detach()会将隐藏状态的梯度完全截断,反向传播时梯度只能传递到当前时间步,完全无法利用历史字符的时序信息,相当于退化成了单字符分类模型,根本发挥不了RNN的序列建模优势。而教程实现中没有手动detach隐藏状态,梯度可以沿着时间步反向传递,能够学习到姓名中前后字符的关联规律。 - 输出层缺少分类必要的激活函数
字符级姓名分类是多分类任务,教程实现最后加了nn.LogSoftmax(dim=1),和NLLLoss损失函数适配,计算的是标准的多分类对数似然损失。你的实现中输出层直接返回全连接层的原始输出,没有做归一化,损失计算逻辑不匹配,会导致模型收敛困难。 - 特征融合方式表达能力不足
你用torch.add合并输入变换结果和隐藏层变换结果,相当于强制要求输入特征和隐藏状态特征维度对齐后做元素级相加,模型能学习到的特征组合方式非常受限。而教程实现是将输入和隐藏状态做torch.cat拼接后再传入线性层,模型可以自主学习输入和历史隐藏状态的融合规则,特征表达能力远强于元素相加的方式。 - 隐藏状态管理逻辑错误
你把初始隐藏状态作为类属性初始化,且没有在每个样本推理前重置隐藏状态,会导致前一个样本的隐藏状态残留到下一个样本的计算中,不同姓名的序列信息互相干扰。而教程实现提供了initHidden方法,每次处理新姓名前都会重置初始隐藏状态,避免了不同样本间的信息污染。 - 权重设计的表达效率差异
教程的实现中输入到隐藏层、输入到输出层共享了拼接后的输入特征,能够同时利用当前字符和历史隐藏状态的信息生成输出,而你的实现中输出层仅从当前隐藏状态获取信息,丢失了直接从当前输入获取关键特征的通路,拟合能力更弱。
内容的提问来源于stack exchange,提问作者MLNOOB
相关产品推荐
相关产品推荐

