第一次把nn.GRU塞进模型里的时候,我并没有被什么门控公式难住,反而是被它那套输入输出形状约定绊了一跤。当时我按(batch, seq, feature)的顺序把张量喂进去,程序不报错,loss 也照常往下掉,但验证集指标始终比预期差一截,排查了两天才发现是自己把seq_len和batch两个维度搞反了——偏偏那两个数字刚好都是 32,广播机制把错误完全掩盖了。这篇就把torch.nn.GRU的输入和输出从头拆一遍:input、h_0、output、h_n这四个张量分别是什么形状,多层和双向配置会把形状搅成什么样,变长序列为什么必须做 pack,以及训练循环里那些文档不会明说、但踩一次就够的坑。文中代码都能直接复制运行,适合已经会写 PyTorch 训练循环、但每次用到 GRU 都要回去翻文档的人。
1. 从一次维度错位说起:torch.nn.GRU 输入端的四条硬性约定
1.1 input 的每个维度分别管什么
构造 GRU 时只有两个参数是必填的:input_size和hidden_size。hidden_size好理解,就是隐状态的宽度;input_size最容易被误读,它描述的是"单个时间步上、单个样本的特征维度",既不是序列长度,也不是 batch 大小。
举个具体的例子:你有一批句子,先通过nn.Embedding转成向量,得到形状(N, L, E)的张量,此时喂给 GRU 的input_size应该是E,而不是L,也不是N。我在早期项目里就见过同事把input_size写成词表大小vocab_size,训练能跑但收敛极慢,原因就是参数矩阵的第一维被凭空放大了几万倍。
默认顺序是(seq_len, batch, input_size),也就是"时间维排在最前面"。为什么不是直觉上的 batch 优先?因为 PyTorch 的 RNN 系列最早是对着底层高性能实现对齐的,那里时间维天然就在第 0 维,框架为了少一次转置就沿用了这个约定。想要 batch 优先,得显式打开batch_first=True。
1.2 batch_first 只管两头,不管中间
这里有个反直觉的细节,也是我见过最多的记混点:batch_first只影响input和output,完全不影响h_0和h_n。隐状态的 batch 维永远固定在第 1 位。
| 张量 | batch_first=False(默认) | batch_first=True |
|---|---|---|
| input | (L, N, H_in) | (N, L, H_in) |
| output | (L, N, D*H_out) | (N, L, D*H_out) |
| h_0 | (D*num_layers, N, H_out) | 同左,不受影响 |
| h_n | (D*num_layers, N, H_out) | 同左,不受影响 |
表里的L是序列长度,N是 batch,D是方向数(单向为 1,双向为 2),H_out就是hidden_size。
提示:如果你手写
h_0时用(N, H)两维传进去,会直接报维度错误;写成(num_layers, N, H)在单向情况下能跑,但双向配置下又会对不上。最稳的写法是torch.zeros(num_layers * num_directions, N, hidden_size),把这几个量算清楚再传。
1.3 h_0 可以不给,但给了就得全对
h_0是可选参数,不传的话模块内部会自动填零。但只要传,就必须同时满足三个条件:形状对、dtype和模块权重一致、设备一致。第三个最容易出事——模型搬到 GPU 上,h_0还留在 CPU,报错信息是设备不匹配,但错误栈会指向 forward 内部,第一眼看不出是自己构造的隐状态的问题。
还有一个隐蔽的坑是张量的内存连续性。如果你先做了x.transpose(0, 1)再取某一段当输入,得到的张量在内存里不是连续的,某些后端路径会直接拒绝,抛出类似 "Expected hidden to be contiguous" 的提示。解决办法很便宜:.contiguous()一下就行。
2. output 与 h_n 的关系:亲手跑一遍比背公式管用
2.1 单向单层配置下,两者数值完全重合
nn.GRU的forward返回的是一个二元组(output, h_n)。很多人第一次看到两个张量,直觉会以为一个是"汇总"一个是"细节",其实它们的区别只在时间维度的取值方式上:output保留了每一个时间步的信息,h_n只保留每个层、每个方向在最后一步的信息。
跑一段最短的验证代码:
import torch import torch.nn as nn torch.manual_seed(0) gru = nn.GRU(input_size=10, hidden_size=20) x = torch.randn(5, 3, 10) # (L=5, N=3, H_in=10) out, h = gru(x) print(out.shape) # torch.Size([5, 3, 20]) print(h.shape) # torch.Size([1, 3, 20]) print(torch.allclose(out[-1], h[-1])) # Trueout[-1]是第 4 个时间步的输出,h[-1]是第 0 层(唯一一层)处理完第 4 个时间步之后的隐状态。单向情况下它们物理上是同一份数据,我一般用torch.allclose而不是==来判断,因为极少数情况下不同算子融合路径会带来浮点级的尾差。
2.2 h_n 的堆叠顺序是"先层后方向"
h_n的第 0 维长度等于num_layers * num_directions,但索引规则不是"所有正向排前面、所有反向排后面",而是层优先:第 0 层占前两个位置(正向在前、反向在后),第 1 层再占后面两个,以此类推。写代码时可以用h_n.view(num_layers, num_directions, N, H)把它重排成一个更直观的四维张量,我个人在做多层双向模型时几乎都会加这一步,省得每次都要在脑子里数下标。
2.3 output 永远是"最后一层"的输出,这点要记牢
一个常见的误解是以为output把所有层的输出都堆在一起了,实际不是——output只包含最后一层在每个时间步上的输出。中间层的结果只体现在h_n里。所以当你需要中间层的表示(比如做多层特征融合、给不同层加辅助损失)时,只能一层一层手动调用,或者用nn.GRU之外的方式拆开。
顺带一个很实用的等式:在单向(不管多少层)的情况下,h_n[-1]永远等于output在最后一个时间步上的切片。多层单向时output[-1]依然是最后一层的最后一步,所以这个等式始终成立。这也是为什么很多分类代码里out[-1]和h[-1]可以互换使用。
3. 多层与双向开关一开,形状就开始连锁变化
3.1 num_layers 只撑大 h_n,不改 output 的宽度
把num_layers从 1 改成 2,output的最后一维仍然是hidden_size(因为只有最后一层会输出),但h_n的第 0 维从 1 变成 2。中间层的输入维度由上一层自动衔接,第一层吃input_size,之后的层都吃hidden_size,这些都不用你操心。
真正需要操心的是:如果你想把h_n从第一层传递到下一批数据继续用,得注意h_n[0]是第一层的状态、h_n[-1]是最后一层的状态,方向别搞反。
3.2 bidirectional 把输出宽度直接翻倍
打开bidirectional=True之后,D变成 2,output的最后一维变成2 * hidden_size。拼接顺序是正向在前、反向在后:前hidden_size维来自正向,后hidden_size维来自反向。
gru = nn.GRU(10, 20, num_layers=2, bidirectional=True) x = torch.randn(5, 3, 10) out, h = gru(x) print(out.shape) # torch.Size([5, 3, 40]) print(h.shape) # torch.Size([4, 3, 20]) 两层 x 两方向注意h_n的最后一维不会翻倍,仍然是hidden_size,翻倍只发生在output上。因为正向和反向各自维护一套独立的隐状态,只是在输出时被拼到了一起。
3.3 反向那一路的"最后状态"其实对应第一个时间步
这是双向 GRU 里最值得单独拿出来讲的一点。反向这一路是从序列末尾往前处理的,所以它处理完所有输入之后的隐状态,落在output的第 0 个时间步的后半段,而不是最后一个时间步。
验证一下:
gru = nn.GRU(10, 20, bidirectional=True) out, h = gru(torch.randn(5, 3, 10)) print(torch.allclose(out[-1, :, :20], h[0])) # True,正向的终态在末尾 print(torch.allclose(out[0, :, 20:], h[1])) # True,反向的终态在开头由此可以推出一个很实用的结论:如果你要从双向 GRU 里提取"看过整句话之后再给出的句子表示",最优取法是把output[-1]的前半段和output[0]的后半段拼起来,因为这两段分别是正向、反向各自看完整个序列之后的表示。而output[-1]的后半段只见过最后一个词,信息量很少。
如果懒得做这套拼接,最稳的做法是直接对output做池化:
sent_repr = out.mean(dim=0) # (N, 40) # 或者只用两侧终态拼接 sent_repr = torch.cat([h[0], h[1]], dim=-1) # (N, 40)后者在分类任务里用得最多,因为它取的是两个方向真正的"收敛态",不用管时间维下标怎么数。
4. 变长序列:padding 之后不做 pack,模型在悄悄学填充符
4.1 不 pack 到底错在哪
这是我认为 GRU 使用中危害最大、又最不容易被发现的一个问题。假设一个 batch 里有三条长度分别为 5、2、3 的序列,为了凑成矩阵你必须把短的补到 5。如果你直接把补零之后的张量喂给 GRU,那么对于第二条序列,GRU 在第 5 个时间步之后的隐状态,是它"读过三个填充符"之后的状态,而不是读完第二个真实词之后的状态。此时h_n里装的已经不是句子的语义了。
更糟的是这件事不会报错。模型照样训练,指标可能只是略微变差,你会以为是数据或者超参数的问题,很难定位到这里。
4.2 从 padding 到 pack 的完整链路
标准做法是用pack_padded_sequence把填充部分折叠掉,让 GRU 只在真实长度上计算,输出再用pad_packed_sequence还原成带填充的矩阵。
import torch import torch.nn as nn from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence emb = nn.Embedding(100, 32, padding_idx=0) gru = nn.GRU(32, 64, batch_first=True) x = torch.tensor([ [1, 2, 3, 4, 5], [6, 7, 0, 0, 0], [8, 9, 10, 0, 0], ]) lengths = torch.tensor([5, 2, 3]) # 必须是 CPU 上的 int64 e = emb(x) # (3, 5, 32) packed = pack_padded_sequence(e, lengths, batch_first=True, enforce_sorted=False) out_packed, h = gru(packed) out, out_lengths = pad_packed_sequence(out_packed, batch_first=True) print(out.shape) # torch.Size([3, 5, 64]) print(h.shape) # torch.Size([1, 3, 64]) print(out_lengths) # tensor([5, 2, 3])几个容易忽略的点:第一,PackedSequence内部永远是时间维优先的,所以即使模块设了batch_first=True,从 pack 得到的输出仍然要先pad_packed_sequence才能按 batch 维处理。第二,喂进去的是PackedSequence,那output出来也是PackedSequence,不能直接.mean(dim=0),必须先还原。第三,padding_idx=0在 embedding 里设上能进一步减少填充符带来的干扰,虽然 pack 之后它本来就参与不了计算,但万一某段逻辑绕过了 pack,这行配置能兜底。
4.3 enforce_sorted 与三类高频报错
pack_padded_sequence默认要求序列长度是降序排列的,如果不满足又不显式设置enforce_sorted=False,就会得到一个提示长度未排序的错误。传enforce_sorted=False之后,函数内部会自动排序、计算、再还原顺序,代价是一次额外的索引操作,绝大多数场景完全可以接受。
日常最容易撞上的三类报错,我整理成了表:
| 报错关键词 | 根本原因 | 处理方式 |
|---|---|---|
| lengths must be a CPU tensor | 长度张量被放到 GPU 上了 | 构造时不要.to(device),保持在 CPU |
| lengths must be of type torch.int64 | 用了默认的 int32 或浮点 | 构造时写dtype=torch.long |
| sorted_indices / decreasing order | 长度没降序且未关排序检查 | 加enforce_sorted=False |
| Expected all tensors on same device | 长度或 h_0 与权重设备不一致 | 统一 device,.contiguous()兜底 |
还有一个替代方案:不 pack,改为手动按真实长度取output。具体是构造index = (lengths - 1).view(-1, 1, 1).expand(-1, 1, hidden_size),再output.gather(0, index)。这种写法在自定义层数解耦、需要逐层干预的场合更灵活,但速度上不如 pack,我在序列不算长的项目里才会偶尔用它。
5. 把 GRU 接到下游任务:分类头和序列标注头的写法完全不同
5.1 句子分类:拿隐状态接一个线性层
最典型的结构是 embedding 加 GRU 加线性分类头。核心是决定用哪一份张量作为"句子表示"。
class GruClassifier(nn.Module): def __init__(self, vocab_size, emb_dim, hidden, num_cls): super().__init__() self.emb = nn.Embedding(vocab_size, emb_dim, padding_idx=0) self.gru = nn.GRU(emb_dim, hidden, num_layers=1, batch_first=True, bidirectional=True) self.fc = nn.Linear(hidden * 2, num_cls) def forward(self, x, lengths): e = self.emb(x) # (N, L, E) packed = pack_padded_sequence(e, lengths, batch_first=True, enforce_sorted=False) _, h = self.gru(packed) # (2, N, H) feat = torch.cat([h[0], h[1]], dim=-1) # (N, 2H) return self.fc(feat)这里用h[0]和h[1]而不是output[-1],原因是双向配置下output[-1]只含正向的收敛态,反向那半段几乎是没用的。这一点和单向模型的写法差别很大,改配置时千万别只改bidirectional忘了改取法。
5.2 序列标注:必须用完整的 output
做词性标注、实体识别这类逐时间步输出的任务时,h_n是完全不够用的,因为每一个时间步都要出结果。写法是还原后的output直接过一层Linear:
logits = self.fc(out) # (N, L, num_tags)这里有一个细节值得注意:还原后的output在填充位置上的值,是 GRU 对填充符计算出来的结果,不是零。所以计算损失时必须配合 mask,把填充位置排除掉:
mask = (x != 0) # (N, L) loss = criterion(logits.transpose(1, 2), y) loss = (loss * mask).sum() / mask.sum()如果直接在整段序列上算平均损失,填充位置会稀释掉梯度信号,而且在类别极不均衡的时候会明显拉偏模型。我在这上面吃过亏:一开始指标看起来"还行",加上 mask 之后同一个模型的 F1 直接涨了好几个点。
5.3 用权重形状反推门控顺序
如果你需要自己实现一个等价的前向、或者想确认权重到底存在哪,可以直接看参数形状:
gru = nn.GRU(10, 20) print(gru.weight_ih_l0.shape) # torch.Size([60, 10]) 3*H_out x input_size print(gru.weight_hh_l0.shape) # torch.Size([20*3, 20]) r, z, n = gru.weight_ih_l0.chunk(3, dim=0) # 顺序:重置门、更新门、新门形状是3 * hidden_size,对应三个门在输出方向上拼接。顺序是重置门、更新门、新候选,不是直觉上的"更新门在前"。多层的话权重名会带层号后缀,比如weight_ih_l1;双向则分_reverse后缀,比如weight_ih_l0_reverse。写自定义初始化的时候按这些名字去遍历,比手动枚举稳妥得多。
顺带说一句参数初始化:PyTorch 默认用均匀分布初始化 GRU 权重,范围由hidden_size决定。在序列较长、层数较深的时候,默认初始化有时会让前几个 epoch 的梯度偏小,我习惯把自己的 embedding 用正态分布初始化,GRU 部分保持默认,这样就够用了。
6. 训练循环里的三件小事:隐状态传递、detach 和 dropout
6.1 跨 batch 传隐状态必须 detach
做语言建模或者需要跨越 batch 边界延续状态的场景时,会把上一批的h_n当作下一批的h_0。这时如果不.detach(),计算图会一直往后延伸,跑不了几步就会抛出 "Trying to backward through the graph a second time" 的错误。
h = None for xb, yb in loader: out, h = model(xb, h) loss = criterion(out, yb) loss.backward() optimizer.step() optimizer.zero_grad(set_to_none=True) h = h.detach() # 关键一步:截断计算图set_to_none=True是个小细节,比清成零稍微省一点显存和带宽,在长序列任务里积少成多。
6.2 dropout 只在层与层之间生效
nn.GRU的dropout参数不是对输入或输出做丢弃,而是对每一层之间的输出做丢弃。所以当num_layers=1时这个参数实际上没有任何作用,PyTorch 会给出一个提醒。要让 dropout 真正生效,至少得两层。如果你的模型是单层 GRU,又想加正则化,正确的位置是在 embedding 之后、或者在线性分类头之前手动加nn.Dropout。
6.3 报错速查
把前面提到的坑集中成一张表,方便排查时对号入座:
| 现象或报错 | 大概率原因 | 处理方式 |
|---|---|---|
| 维度能跑通但效果异常差 | seq 与 batch 维写反,或未 pack | 打印 shape 核对,改用 pack |
| Expected hidden size 报错 | h_0 形状与 num_layers/direction 不匹配 | 用num_layers * num_directions计算 |
| Trying to backward through the graph a second time | 跨 batch 传递 h_n 未 detach | 在传回前h.detach() |
| dropout 参数似乎无效 | num_layers 为 1 | 增加层数或手动加 Dropout |
| GPU 上性能明显低于预期 | 隐状态非连续或未走高效路径 | .contiguous(),检查 pack 是否生效 |
| 输入张量形状正确但结果随机 | 反向那半段被误用为句子表示 | 改用h[0]/h[1]拼接或池化 |
我个人固化下来的习惯是:任何涉及 GRU 的改动,先在 CPU 上用torch.randn造一批小张量,把input、output、h_n三个形状打印出来,确认无误再切到真实数据上跑。这一步大概花三十秒,但能省掉后面几小时盯着 loss 曲线发呆的时间。另外,output和h_n的关系建议你也亲手跑一遍验证代码,比记十条笔记都管用——我自己就是被打脸之后才真正记住"双向配置下反向终态落在时间步 0"这件事的。