GRU 这个词在 PyTorch 圈子里出现的频率极高但真正把torch.nn.GRU的输入输出形状、batch_first的语义、h_n和output的区别一次性讲透的资料并不多。我前后在几个文本分类和时序预测项目里用过它也帮同事排查过不下几十次的维度报错发现绝大多数问题都集中在同一个地方输入输出的三维张量到底哪一维是什么以及在多层、双向的情况下这些维度怎么组合。这篇就把torch.nn.GRU的输入输出彻底拆开讲一遍从构造参数到前向传播从最小可运行示例到变长序列的 padding 陷阱再到实际项目里该怎么接下游网络。如果你刚开始接触循环神经网络或者已经用过但每次写代码都要翻文档确认形状这篇应该能让你以后少查几次 API。1. torch.nn.GRU 到底在算什么从门控逻辑到构造参数很多人写nn.GRU的时候是把它当黑盒用的输入丢进去输出拿出来形状对得上就行。但只要涉及调试黑盒就会变成折磨。所以先把内部在算什么捋清楚后面所有形状规则都能从这套逻辑里推导出来不用死记。1.1 门控结构决定了你要关心哪些形状GRU 的核心是两个门加一个候选状态。重置门r_t决定上一时刻的隐藏状态有多少要参与候选状态的计算更新门z_t决定新旧状态各占多少比例候选状态n_t则是当前输入和历史信息的融合结果。用公式写出来是这样r_t sigmoid(W_ir x_t b_ir W_hr h_{t-1} b_hr)z_t sigmoid(W_iz x_t b_iz W_hz h_{t-1} b_hz)n_t tanh(W_in x_t b_in r_t * (W_hn h_{t-1} b_hn))h_t (1 - z_t) * n_t z_t * h_{t-1}从形状角度读这些公式能得到几个关键结论。第一输入x_t的最后一维必须等于input_size隐藏状态h_t的最后一维必须等于hidden_size这两者相互独立你完全可以设计一个input_size100、hidden_size32的模型把高维输入压到低维隐空间。第二三个门的计算都要把x_t和h_{t-1}投影到hidden_size维度上再相加所以 PyTorch 把三组权重在输出维度上拼接存储这也是为什么你去打印weight_ih_l0会看到形状是(3 * hidden_size, input_size)而不是三个独立的矩阵。那3这个因子对应的是哪三个门顺序是什么这点官方文档写得比较隐蔽实际是r、z、n的顺序也就是重置门、更新门、候选状态。你在做手写对齐或者加载预训练权重时如果搞错了顺序模型不会报错但效果会莫名其妙地崩掉。我见过有人把权重导出来在 NumPy 里重写推理结果精度差了一大截排查半天就是这里的顺序反了。注意如果你打算把 PyTorch 训练好的权重导出到其他框架或者手写推理务必确认门控顺序是重置门、更新门、候选状态而不是按直觉猜的更新门在前。再看参数量。单层单向的 GRU 一共有3 * hidden_size * (input_size hidden_size) 6 * hidden_size个参数前面的3 * hidden_size * (input_size hidden_size)是两组权重矩阵后面的6 * hidden_size来自两个偏置项b_ih和b_hh每个都是3 * hidden_size。这个公式在估算显存和模型大小时很有用。举个实际数字input_size300、hidden_size128的单层 GRU参数量是3 * 128 * (300 128) 6 * 128 164352 768 165120大概 16 万参数比一层 512 维的全连接层约 26 万还小。所以 GRU 本身并不吃参数真正吃资源的是它按时间步展开的计算量。1.2 nn.GRU 构造参数逐项翻译nn.GRU的构造函数签名里参数不多但每一个都会影响输入输出的形状或者训练行为逐个过一遍。input_size就是每个时间步输入向量的维度。做词向量输入时它等于词向量维度做传感器时序时它等于每个时刻的特征数量。注意它不等于词表大小也不等于 batch size这是新手最容易混淆的一点。hidden_size是隐状态维度也是输出的最后一维的基础值调大能提升表达能力但会线性增加计算量。num_layers是堆叠层数默认 1。层数大于 1 时第 i 层的输出会作为第 i1 层的输入所以中间层不需要你手动指定维度框架会自动处理。但要注意多层时h_0的第一维会从1变成num_layers这是形状报错的高发区。bias默认True关掉的话上面参数量公式里的6 * hidden_size就没了。实践中很少有人关除非你在做极致的量化压缩。batch_first默认是False这是 PyTorch 循环神经网络家族的历史遗留设计也是被吐槽最多的一个参数。它默认要求输入形状是(seq_len, batch, input_size)也就是序列长度在前。现代的数据加载器习惯把 batch 放第一维所以实际项目里基本上都会显式写batch_firstTrue。这个参数只影响输入x和输出output的前两维顺序不影响h_0和h_n后者永远是 batch 在第二维。dropout只在num_layers 1时生效作用于层与层之间。如果你设置了num_layers1却传了dropout0.5框架会给出警告并且忽略这个参数。我见过有人以为单层 GRU 也能靠这个参数做正则化训练了半天没效果就是这个原因。bidirectional打开后hidden_size的实际输出维度会翻倍output的最后一维变成2 * hidden_sizeh_n的第一维也会翻倍。这个后面单独展开说。2. 输入张量的形状规则为什么总是三维GRU 的输入必须是三维张量这个约束经常让从全连接网络转过来的人不适应。全连接层可以接受(batch, features)的二维输入GRU 为什么不行因为多出来的一维就是时间。2.1 输入 x 的三个维度分别代表什么batch_firstFalse默认时x的形状是(seq_len, batch, input_size)。三个维度依次是时间步数、批大小、单步特征维度。batch_firstTrue时变成(batch, seq_len, input_size)只是把前两维换了个位置语义不变。为什么默认把序列长度放前面因为早期 PyTorch 的设计参考了 Torch7 和 cuDNN 的接口约定cuDNN 底层更倾向于时间维在前这样在按时间步循环时内存访问更连续。虽然现在的实现已经做了优化但这个默认值一直保留下来了。实际写代码时我的习惯是统一用batch_firstTrue理由是 DataLoader 出来的 batch 天然是第一维如果再用默认值就得在 forward 里做两次transpose代码里到处是.transpose(0, 1)很难维护而且 transpose 返回的是视图后续如果接view或者reshape很容易踩到内存不连续的坑。举个具体例子。假设你在做一个基于传感器数据的动作识别任务采集频率 50Hz窗口长度 2 秒每条样本就是 100 个时间步每个时间步有 6 个特征三轴加速度加三轴角速度。batch size 取 32那么batch_firstTrue时输入形状就是(32, 100, 6)input_size6。这三个数字分别对应什么写代码时一定要在心里默念一遍因为(100, 32, 6)和(32, 100, 6)在很多情况下都能跑通不会报错但语义完全错了模型学不到任何东西。关于输入维度还有一个容易忽略的点input_size在模型定义时就固定了运行时如果喂进来的特征维度对不上会直接抛RuntimeError。这个错误信息通常长这样Expected input of size (..., 6), got (..., 8)。看到这类报错先检查特征工程那一步是不是多加了几列比如把时间戳或者 ID 也一起塞进去了。2.2 batch_first 参数到底改了什么这个参数的作用范围经常被误解。它只改两件事输入x的前两维顺序以及输出output的前两维顺序。它不改h_0的形状也不改h_n的形状。这两个张量永远保持(num_layers * num_directions, batch, hidden_size)。所以你会遇到这种组合batch_firstTrue输入是(32, 100, 6)h_0却是(1, 32, 32)而不是(32, 1, 32)。新手看到这个组合会本能地觉得不一致然后把h_0写成(32, 1, 32)接着就会收到这样的报错RuntimeError: Expected hidden size (1, 32, 32), got (32, 1, 32)这个报错信息其实很友好直接把期望值和实际值都打出来了。遇到就按提示改不用猜。还有一种情况是h_0干脆不传。这时候框架会默认用全零初始化等价于h_0 torch.zeros(num_layers * num_directions, batch, hidden_size, devicex.device)。不传h_0在大部分任务里是合理的因为零初始状态经过几个时间步就会被输入信息覆盖掉。但在一些对初始状态敏感的任务里比如很短的序列显式传一个可学习的初始状态可能会有帮助这时候把h_0定义成nn.Parameter就可以了。2.3 初始隐藏状态 h_0 的维度怎么定h_0的形状公式是(num_layers * num_directions, batch, hidden_size)。这里的乘法是直接相乘不是拼接。单向时num_directions1双向时是2。三层双向的话第一维就是6。这个第一维的内部排列顺序也值得记一下按层排列每层内部先正向后反向。也就是说num_layers2、双向时索引0是第 0 层正向1是第 0 层反向2是第 1 层正向3是第 1 层反向。如果你想让不同层用不同的初始状态就按这个顺序去填充。如果所有层都用零直接传None或者创建全零张量都行。我一般会在 forward 里写if h0 is None: h0 torch.zeros(self.num_layers * self.num_directions, x.size(0), self.hidden_size, devicex.device, dtypex.dtype)注意这里的dtypex.dtype混合精度训练时这个细节很关键。如果忘了写默认创建的h_0是float32而x可能是float16两者相加就会报 dtype 不匹配。这类错误在纯float32环境下不会出现一上 AMP 就暴露了。3. 输出到底是什么output 与 h_n 的区别与联系nn.GRU的前向传播返回一个元组第一个是output第二个是h_n。很多人搞不清这两个到底有什么区别觉得反正都是隐藏状态随便取一个用。实际上它们的语义差别很大用错了会直接影响模型效果。3.1 output 的每一帧从哪来output包含了每个时间步的隐藏状态。batch_firstTrue时形状是(batch, seq_len, num_directions * hidden_size)。对于单向、单层的 GRUoutput[:, t, :]就是第 t 个时间步的隐状态h_t。这里有个细节output只包含最后一层的输出。如果你堆了 3 层output是第 3 层每个时间步的输出第 0 层和第 1 层的中间结果你是拿不到的。这一点在做特征提取或者可视化时很重要想要中间层的表示只能自己手动逐层调用或者拆成多个单层 GRU 手动串联。再说长度。output的seq_len维度和输入的seq_len完全一致即使你传了h_0也一样。这意味着一件事如果你做了 paddingoutput在 padding 位置也会有值。这些值是基于 padding 内容算出来的语义上是无意义的后面接下游网络时必须做 mask否则 padding 的噪声会污染结果。这是我在实际项目里踩过的最深的坑之一下面第 5 节会展开讲。3.2 h_n 与 output[-1] 的等价性验证这是最常被问到的一个问题output[:, -1, :]和h_n[-1]是一回事吗在单向、单层的情况下答案是完全相同数值上一模一样。原因是h_n存的就是最后一个时间步的隐状态而output[:, -1, :]也是最后一个时间步的输出两者指向同一个东西。在单向、多层的情况下output[:, -1, :]仍然等于h_n[-1]。因为output本来就是最后一层的输出序列h_n[-1]也是最后一层的最后时刻状态。在多方向或者双向的时候就不同了。双向时output[:, -1, :]是[正向的最后一步, 反向的最后一步]而反向分支的最后一步对应的其实是序列的第一个位置。同时h_n[-2:]里存的是[正向最后一步, 反向最后一步]这里的反向最后一步对应的是输入序列的开头。所以output[:, -1, :]和h_n[-2:]虽然都是两个hidden_size的拼接但反向那半段不是同一个东西。可以用一段几行的代码验证这个等价关系这也是我自己写完 GRU 之后的固定检查动作import torch import torch.nn as nn torch.manual_seed(42) gru nn.GRU(input_size5, hidden_size6, num_layers2, batch_firstTrue) x torch.randn(4, 7, 5) out, hn gru(x) # 单向多层output 最后一帧等于 h_n 的最后一层 print(torch.allclose(out[:, -1, :], hn[-1], atol1e-6)) # True如果你把这个检查写进单元测试以后重构网络结构时能第一时间发现行为变化。3.3 单向/双向/多层组合下的形状对照这张表是我自己整理的贴在工位上一段时间后来记住了才撕掉。基准设定是batch4、seq_len7、input_size5、hidden_size6、batch_firstTrue。配置output 形状h_n 形状output[-1] 与 h_n 的关系单层单向(4, 7, 6)(1, 4, 6)相等两层单向(4, 7, 6)(2, 4, 6)out[:, -1, :] 等于 hn[-1]单层双向(4, 7, 12)(2, 4, 6)不等仅前 6 维对应正向末尾两层双向(4, 7, 12)(4, 4, 6)不等需按层按方向拆如果batch_firstFalse把上面表格里 output 的前两维对调即可h_n不变。这个规律记住之后任何维度报错都能在脑子里快速定位。还有一个常见的取用方式做分类任务时从双向模型里提取句子表示。标准写法是把h_n拆成正向和反向两部分然后拼接num_layers, batch, hidden hn.shape[0] // 2, hn.shape[1], hn.shape[2] hn hn.view(num_layers, 2, batch, hidden) forward_last hn[-1, 0] # 正向最后一步 backward_last hn[-1, 1] # 反向最后一步 sentence_vec torch.cat([forward_last, backward_last], dim1) # (batch, 2*hidden)这段代码几乎是双向文本分类的标配值得直接记住。注意hn.view的第一个参数是层数不是num_layers * 2这里写错的话张量元素总数对不上会立刻报错所以还算安全。4. 可直接复现的完整示例前面讲的都是规则这一节直接把代码贴出来从构造到前向传播到结果验证照着跑一遍就能把形状彻底搞明白。4.1 最小可运行示例import torch import torch.nn as nn torch.manual_seed(0) batch_size 4 seq_len 7 input_size 5 hidden_size 6 num_layers 2 gru nn.GRU( input_sizeinput_size, hidden_sizehidden_size, num_layersnum_layers, batch_firstTrue, ) # 输入batch 在第一维 x torch.randn(batch_size, seq_len, input_size) # 初始状态层数在前batch 在第二维 h0 torch.zeros(num_layers, batch_size, hidden_size) out, hn gru(x, h0) print(input :, x.shape) # torch.Size([4, 7, 5]) print(output:, out.shape) # torch.Size([4, 7, 6]) print(h_n :, hn.shape) # torch.Size([2, 4, 6]) # 验证最后一层的最后一帧 print(torch.allclose(out[:, -1, :], hn[-1], atol1e-6)) # True跑通之后可以试着把batch_first改成False同时把x改成(7, 4, 5)观察output的形状变化。这个对照实验做一次以后就不会再混淆了。4.2 手写单步 GRU 对齐官方实现想要真正确定自己理解了门控顺序和公式最好的办法是手写一个单步版本跟官方实现对一遍数值。import torch import torch.nn as nn torch.manual_seed(1) hidden_size, input_size 6, 5 gru nn.GRU(input_size, hidden_size, batch_firstTrue) w_ih gru.weight_ih_l0 # (3*hidden, input) w_hh gru.weight_hh_l0 # (3*hidden, hidden) b_ih gru.bias_ih_l0 # (3*hidden,) b_hh gru.bias_hh_l0 # (3*hidden,) def manual_step(x_t, h_prev): gi x_t w_ih.T b_ih gh h_prev w_hh.T b_hh i_r, i_z, i_n gi.chunk(3, dim-1) h_r, h_z, h_n gh.chunk(3, dim-1) r torch.sigmoid(i_r h_r) z torch.sigmoid(i_z h_z) n torch.tanh(i_n r * h_n) return (1 - z) * n z * h_prev x torch.randn(1, 3, input_size) out, _ gru(x) h torch.zeros(1, hidden_size) for t in range(x.size(1)): h manual_step(x[:, t, :], h) print(torch.allclose(h, out[:, -1, :], atol1e-6)) # True这段代码有两个价值点。一是确认了chunk(3)出来的顺序是重置门、更新门、候选状态二是确认了b_ih和b_hh是两个独立的偏置不是共享的。如果你在做模型移植或者量化这段代码可以直接当参考实现。4.3 文本分类里的实际接法光看形状示例还不够实际项目里怎么把 GRU 接进一个完整的分类网络是更实际的问题。下面是一个可以直接改改就用的文本分类模型。import torch import torch.nn as nn from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence class GRUClassifier(nn.Module): def __init__(self, vocab_size, emb_dim, hidden_size, num_classes, num_layers1, dropout0.3): super().__init__() self.emb nn.Embedding(vocab_size, emb_dim, padding_idx0) self.gru nn.GRU( input_sizeemb_dim, hidden_sizehidden_size, num_layersnum_layers, batch_firstTrue, bidirectionalTrue, dropoutdropout if num_layers 1 else 0.0, ) self.dropout nn.Dropout(dropout) self.fc nn.Linear(hidden_size * 2, num_classes) def forward(self, token_ids, lengths): emb self.emb(token_ids) # (B, L, E) packed pack_padded_sequence( emb, lengths.cpu(), batch_firstTrue, enforce_sortedFalse ) out, hn self.gru(packed) out, _ pad_packed_sequence(out, batch_firstTrue) # (B, L, 2H) # 用 mask 做平均池化避开 padding 位置 mask (token_ids ! 0).unsqueeze(-1).float() # (B, L, 1) summed (out * mask).sum(dim1) pooled summed / mask.sum(dim1).clamp(min1e-9) return self.fc(self.dropout(pooled))这里选择平均池化而不是取h_n原因是平均池化对短文本更稳尤其是当句子长度差异很大时h_n会被最后一个有效词的信息过度主导。当然如果你的任务是判断整句的语义倾向取h_n拼接也完全可行把pooled那几行换成从hn里拆方向拼接即可。padding_idx0这个参数别漏。它让 embedding 层在反向传播时跳过 id 为 0 的位置避免 padding 污染词向量。配合pack_padded_sequence两个一起用才能把 padding 的影响压到最低。5. 变长序列、padding 与那些容易被忽略的坑真实数据里序列长度几乎不可能整齐划一要么 padding要么截断。这一节讲清楚 padding 到底会带来什么后果以及怎么处理才干净。5.1 右侧 padding 对 h_n 的污染假设一个 batch 里有两条序列长度分别是 3 和 5padding 到 5。右边补两个零向量。如果不做任何处理直接喂给单向 GRU会发生什么第一个样本的h_n是经过 5 个时间步算出来的其中后两个时间步的输入是零向量。注意输入是零不代表隐状态不变化。因为z_t和r_t都是sigmoid(W 0 W h b)只要偏置不为零门控值就不是固定值h_t会继续演化。所以最终得到的h_n和真正读到第 3 个词就停下的结果是不一样的。这个差异有多大如果序列很长、padding 比例很小影响可以忽略。但如果序列普遍很短而 padding 很多或者偏置初始化得比较激进h_n会明显偏移。我做过一个对比实验在一个长度为 10 到 15 的短文本任务上不做 pack 和做 pack 的准确率差了将近 2 个百分点这对分类任务来说不算小。5.2 pack_padded_sequence 的正确用法pack_padded_sequence的作用就是把 padding 位置直接从计算中剔除让 GRU 在每个 batch 内只处理有效的时间步。用法上有几个必须注意的点。第一lengths必须是 CPU 上的 int64 张量。如果你从 GPU 上的张量直接切出来会报错。所以代码里要写.cpu()。第二enforce_sortedFalse建议显式写上。默认为True时要求一个 batch 内的序列按长度降序排列很多人的数据加载器不做这个排序就会收到Expected sorted相关的报错。设成False后框架内部会自动处理排序和还原用起来省心。第三pack 之后output的类型是PackedSequence不能直接做索引或者切片必须先pad_packed_sequence还原。还原之后output的序列长度等于这个 batch 内的最大长度不是全局最大长度如果后面要拼接或者做定长操作可能还需要再 pad 一次。第四pad_packed_sequence返回的第二项是lengths很多人用不上就忽略了。但如果你的下游逻辑依赖有效长度这个值是有用的。还有一个小细节pack 之后h_n是准确的不受 padding 影响。所以在只需要最后状态的任务里其实可以 pack 之后直接取h_n跳过还原那一步能省一点显存和时间。5.3 双向模型里 output 最后一帧不等于序列结尾这是前面提过但值得单独强调的一点。双向 GRU 的反向分支是从序列末尾往开头读的所以在原始时间顺序下反向分支的最后一步对应的是序列的第一个位置。这意味着output[:, -1, :]的后半段反向部分实际上是反向分支读完整条序列之后的状态它对应的语义位置是序列的开头。如果你把它当成句尾特征来用逻辑就错了。对于定长序列这个问题不明显因为所有位置都有有效内容。但对于 padding 过的变长序列output[:, -1, :]的反向部分对应的是 padding 区域还是第一个有效词取决于你有没有做大调整。一旦搞混模型可能仍然能训练但收敛会变慢效果不稳定。我的建议是双向模型要取全局表示一律走h_n拆分拼接那条路或者用带 mask 的池化。不要图省事直接切output[:, -1, :]。6. 常见报错与排查技巧实录前面讲的是原理和正确用法这一节讲实战里最容易撞到的问题。我把这几年遇到的典型报错整理成了一张速查表遇到问题先对照能省不少时间。6.1 形状类报错速查表报错信息关键词常见原因处理方式input must have 3 dimensions, got 2把单条序列(L, F)直接喂进去加unsqueeze(0)补 batch 维或改用nn.GRUCellExpected hidden size (1, 4, 6), got (4, 1, 6)h_0维度顺序写反改成(num_layers*num_directions, batch, hidden)Expected input size (..., 6), got (..., 8)特征维度与input_size不匹配检查特征工程列数或改模型定义For batched 3-D input, hx should also be 3-Dh_0传了二维张量补上第一维的层数dropout相关 warning 但训练正常num_layers1却设了 dropout把 dropout 设为 0 或加大层数lengths相关报错lengths在 GPU 上或未排序.cpu()加enforce_sortedFalse这张表里最值得说的是第一条。nn.GRU和nn.GRUCell的区别就在这里。GRUCell只处理一个时间步输入是二维的(batch, input_size)适合你想自己写循环、需要在中途插入自定义逻辑的场景。GRU处理整个序列输入必须三维。两者名字很像但输入要求完全不同选错了就会一直报维度错误。判断标准很简单如果你需要在每个时间步之间做额外操作比如注意力、条件判断、动态停止就用GRUCell如果只是标准的序列编码用GRU更快更省事。6.2 dtype、device 与 dropout 的隐性坑形状问题解决之后剩下的坑大多是隐性的不报错但影响结果。一个是 dtype。混合精度训练时如果h_0手动创建的时候没指定dtype默认是float32与float16的输入不匹配会报一个expected scalar type Half but found Float的错误。解决方式就是在创建张量时统一用x.dtype和x.device。更省事的写法是直接不传h_0让框架按输入自动创建这样 dtype 和 device 都不会错。另一个是 device。模型在 GPU 上但h_0在 CPU 上时报错信息有时候不会直说是设备问题而是给一个比较绕的维度错误。所以养成习惯凡是在 forward 里新建张量一律带上devicex.device。第三个是 dropout 的警告。前面提过num_layers1时传非零 dropout 会有一条UserWarning。这条警告有些人直接忽略了但在做对比实验时会产生误解以为自己加了正则化其实没有。看到这条警告就老老实实把 dropout 设成 0或者把层数加到 2 以上。还有一个小坑跟batch_first的切换有关。有些开源代码在内部用了默认的batch_firstFalse然后在你传数据的地方做了transpose。如果你接手这类代码又改了batch_first很可能只改一处导致输入和输出一边转置了另一边没转模型能跑但学不对。接手别人代码时先全局搜一下batch_first和transpose(0, 1)出现的位置心里有数再动手。6.3 参数初始化与性能调优经验默认的初始化对 GRU 来说能用但不算最优。循环权重用正交初始化通常能改善梯度传播输入权重用 Xavier 也比我见过的默认均匀分布更稳。def init_gru_weights(module): for name, param in module.named_parameters(): if weight_ih in name: nn.init.xavier_uniform_(param) elif weight_hh in name: nn.init.orthogonal_(param) elif bias in name: nn.init.zeros_(param) model.apply(init_gru_weights)如果希望初始时更新门偏向保留旧状态可以把更新门对应的偏置设成正值这样sigmoid的输出大于 0.5更倾向于记住历史。不过这个技巧对长序列更明显短序列任务里差别不大属于可选优化。性能方面有几个实测有效的做法。第一保证在 GPU 上跑GRU 的 GPU 加速比 CPU 快一个数量级以上而且批次越大优势越明显因为按时间步展开的循环可以在 batch 维度上并行。第二hidden_size对齐到 8 的倍数比如 128、256、512在一些硬件上能更快这个优势不像卷积那么显著但试试不亏。第三num_layers增加的收益通常小于把hidden_size加大带来的收益但层数多了会让梯度传得更深、训练更难。我一般先固定层数为 1 或 2把hidden_size调到一个合理值再考虑加层。速度上 GRU 比同配置的 LSTM 快一些因为门少了一个参数量和计算量都小。如果你的任务对精度要求不是极致GRU 通常是性价比更高的选择。我做过一个中文短文本分类的对比同样的训练轮数下GRU 的准确率只比 LSTM 低零点几个百分点但单轮训练时间少了大约四分之一调试迭代的体验明显更好。最后说一个我自己的习惯。每次新写一个 GRU 相关的模块我会先写三四行断言把输入输出的形状全部打出来确认一遍包括output、h_n、以及下游输入的形状。这几行代码跑一次只要几百毫秒但能省掉后面几十分钟的调试。特别是涉及多层加双向的组合时纸面上推算很容易出错跑一下最踏实。上面这些内容基本覆盖了torch.nn.GRU从输入到输出的全部关键点。我个人在实际操作中的体会是形状这件事并不复杂难的是双向和多层叠加之后的组合以及 padding 这种不报错但会悄悄影响结果的问题。把第 3 节的对照表和第 5 节的 pack 用法记住日常开发里遇到的大部分情况都能应付。