☰
PyTorch GRU输入输出形状详解:batch_first、双向与变长序列
2026/10/2 16:10:25 网站建设 项目流程

第一次把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])) # True

out[-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"这件事的。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询