ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyTorch nn.GRU 输入输出形状与 h_n/output 区别详解

PyTorch nn.GRU 输入输出形状与 h_n/output 区别详解 写循环神经网络这块代码时torch.nn.GRU的输入输出形状是最容易反复回查的地方。我自己的习惯是把 shape 注释直接写在每行 forward 代码后面原因很简单这个模块的输入有两个张量输出也有两个张量四个形状里任何一个对不上报错信息读起来都很相似等你隔两个月再回来看自己的代码脑子里那些维度顺序早就乱了。这篇内容就是围绕torch.nn.GRU的输入尺寸、输出含义、权重顺序和实际落地写法展开的把能直接跑起来的示例、算得清楚的维度换算、以及我在变长序列和多层双向场景里踩过的坑都摊开讲。适合刚接触序列建模的新手也适合用惯了batch_firstTrue却说不清h_n和output差别的老手对照排查。1. 先把 nn.GRU 的构造参数摊开看构造一个 GRU 只需要两三个参数就能跑但真正影响后面所有 shape 推理的恰恰是那几个带默认值的参数。先把参数表列清楚后面所有形状推导都从这里出发。1.1 构造参数逐项说明与取值逻辑参数默认值作用什么时候必须改input_size无每个时间步输入特征维度必填由数据决定hidden_size无隐藏状态维度必填自己定num_layers1堆叠层数需要更深特征交互时biasTrue是否带偏置一般不动batch_firstFalse输入是否 batch 在前数据来自 DataLoader 时几乎必改dropout0层间 dropout只在 num_layers1 时生效bidirectionalFalse是否双向需要完整上下文时input_size和hidden_size的区别值得强调。前者描述每个时刻喂进来的向量有多长后者描述网络自己维护的记忆有多长。举个具体场景用 128 维词向量、隐藏层 256 维的单向 GRU那input_size128、hidden_size256输入张量最后一维必须是 128输出张量最后一维一定是 256。这两个数字混起来是新手最常见的错误来源因为它们都是维度但在语义上完全不同——一个是被动的输入宽度一个是主动的容量。num_layers是层数第 0 层的输出会变成第 1 层的输入。这里有个容易忽略的点多层时第 1 层的input_size没有显式传它由第一层的输出维度推出。如果你开了双向第二层的输入宽度就变成了2 * hidden_size因为上一层把正反两个方向的结果拼在了一起。这个换算在写自定义模块时经常需要手算。1.2 两个默认值埋下的隐形坑先说batch_first。默认是False意味着输入是(seq_len, batch, input_size)这个顺序。这个默认设定在早期是从 cuDNN 的接口习惯延续下来的但对绝大多数写着DataLoader、习惯(N, C, ...)的人来说(seq, batch, feat)非常反直觉。我自己在第一次用的时候就是把 batch 放在第一维模型能跑通、loss 也能降只是收敛得莫名其妙地慢直到打印出来x.shape才发现把 32 个样本当成了 32 个时间步序列长度搞反了。这类错误不会报 shape 冲突因为(32, 64, 128)和(64, 32, 128)在维度数量上完全一致只是语义相反非常隐蔽。再说dropout。这个参数只在num_layers 1时才真正起作用如果层数是 1 又传了非零 dropout新版 PyTorch 会给出警告提示你non-zero dropout expects num_layers greater than 1。很多人以为设了 dropout 就有正则效果其实单层时它被静默忽略。正确的用法是在多层场景下让每一层输出到下一层之前做一次 dropout而不是在每个时间步内部做这也是它被称作层间 dropout的原因。另外补一句参数初始化。nn.GRU内部走的是reset_parameters()权重按U(-1/sqrt(hidden_size), 1/sqrt(hidden_size))的均匀分布初始化。这个区间不是随便定的它保证了初始状态下每步输出的方差大致稳定不会随着时间步展开爆炸或消失。如果你在迁移学习时手动覆盖了权重记得留意这个量级直接用randn初始化会让前几十步的梯度非常难看。2. 输入张量的形状规则三个维度分别代表什么2.1 单层单向的最小可运行示例先把最小的例子跑通把所有形状打印出来这是理解一切后续变体的基础。import torch import torch.nn as nn torch.manual_seed(42) input_size 4 hidden_size 3 seq_len 5 batch 2 gru nn.GRU(input_size, hidden_size) x torch.randn(seq_len, batch, input_size) h0 torch.zeros(1, batch, hidden_size) output, hn gru(x, h0) print(x.shape) # torch.Size([5, 2, 4]) (seq, batch, feat) print(output.shape) # torch.Size([5, 2, 3]) 每个时间步的隐状态 print(hn.shape) # torch.Size([1, 2, 3]) 只有最后一步这里三个数字要逐个对上。输入的 5 是序列长度也就是有 5 个时间步2 是 batch即同时处理 2 条样本4 是每个时间步的特征维度。输出的 5 和 2 跟输入保持一致只有最后一维从input_size换成了hidden_size这就是 GRU 在做的事把每个时刻的 4 维输入结合历史记忆映射成 3 维的隐藏表示。有意思的是h0的形状(1, 2, 3)。第一个 1 是num_layers * num_directions因为现在是单层单向所以是 1。很多人会以为h0就是(batch, hidden)多出来的这一维经常导致报错。如果你不想显式传h0直接不传或者传None也可以PyTorch 会自动用全零初始化效果等价。2.2 h_0 的形状为什么带一个层数维度h_0的形状是(num_layers * num_directions, batch, hidden_size)。多出来的第一维本质上是给每一层、每个方向各准备一份初始记忆。因为不同层是独立循环的它们的隐状态互不共享所以需要分别初始化。第 0 层的初始状态是h_0[0]第 1 层是h_0[1]依此类推。注意这个索引顺序是按层优先不是按方向优先双向时h_0[0]是第 0 层正向h_0[1]是第 0 层反向h_0[2]才是第 1 层正向h_0[3]是第 1 层反向。这个顺序有实际意义。当你做 encoder-decoder 结构想把 encoder 最后一层的h_n拿来初始化 decoder 的第一层时需要做一次切片和 reshape。正确的做法是取h_n中对应最后一层的部分而不是直接整个塞过去——层数不一样时形状就对不上了。还有一点h_0必须和模型在同一设备、同一 dtype 上。混用float32和float64、或者在 CPU 上建h0却在 GPU 上跑模型报错信息都是expected scalar type或者设备不匹配看起来跟形状无关实际排查起来要绕一圈。我现在的习惯是统一用x.new_zeros(...)来创建这样设备、dtype 自动和输入对齐省掉一个潜在故障点。2.3 batch_first 切换后的错位陷阱改成batch_firstTrue之后输入变成(batch, seq_len, input_size)输出变成(batch, seq_len, hidden_size)。但h_0和h_n的形状完全不变仍然是(num_layers * num_directions, batch, hidden_size)batch 维始终在中间。这个不对称性是最容易搞混的地方输入输出的 batch 位置变了唯独隐状态没变。gru nn.GRU(input_size, hidden_size, batch_firstTrue) x torch.randn(batch, seq_len, input_size) # (2, 5, 4) output, hn gru(x) # output (2, 5, 3), hn (1, 2, 3)我见过不少人为了让h_n也变成 batch 在前特意去 transpose 一下结果反而把正确的代码改错了。记住这条规则就够了batch_first只影响input和output这两个张量不影响h_0和h_n。提示如果你一边用batch_firstTrue一边按(seq, batch, feat)去构造输入模型不会报错只要另外两个维度凑巧对得上但训练结果会完全错误。养成打印x.shape并和注释对照的习惯比任何调试技巧都管用。3. 输出张量的两条线output 与 h_n 到底有什么区别3.1 output 记录每一步的完整轨迹output是 GRU 在每一个时间步上的隐藏状态堆起来的序列。形状是(seq_len, batch, num_directions * hidden_size)batch_firstTrue时前两维互换。它的长度和输入序列长度完全一致这是一个关键性质无论你输入多长的序列output的长度都和你一一对应。理解output最直观的方式是把它想成逐帧录像。序列有 5 个时刻它就存 5 帧每一帧都是那一时刻网络对整段历史的压缩表示。做序列标注任务时用的就是它因为你需要每个时刻都有一个预测结果比如词性标注、命名实体识别输出的长度必须和输入句子长度一致。单向单层的情况下output[-1]和hn[0]是同一个张量数值上相等是同一份数据的不同视图。这不是巧合最后一步的隐状态当然就是整个序列走完之后的最终状态。3.2 h_n 只是最后一步的快照h_n的形状是(num_layers * num_directions, batch, hidden_size)注意最后一维没有乘num_directions。这是因为h_n是按层、按方向分开存的不像output那样把双向结果拼在一起。做分类任务时用的是h_n。把一整段文本压缩成一个向量、然后接一个全连接层做情感分类这是最经典的用法。取h_n而不是output[-1]的原因是多层或者双向的时候h_n直接就是整理好的结构不需要你再做切片拼接。这里有个细节值得说清楚。单向单层时取output[-1]和取hn[0]完全等价我建议统一用hn的写法因为一旦后面把模型改成双向或多层用output[-1]的代码就得改而用hn的代码逻辑上只需要调整后续的拼接处理。写代码时少留一点改起来要动全身的隐患长期看是划算的。3.3 多层双向场景下两者的切片关系双向的时候output的最后一维是2 * hidden_size前半段是正向结果后半段是反向结果。output[:, :, :hidden_size]是正向每个时刻的输出output[:, :, hidden_size:]是反向每个时刻的输出。h_n在双向时第一维是 2h_n[0]是正向走完整个序列的最终状态对应output的最后一步h_n[1]是反向走完整个序列的最终状态对应output的第一步。为什么反向对应第一步因为反向 GRU 是从序列末尾往前读的它读完整段序列时停在位置 0所以output的第 0 个时刻正好是反向的终点。# num_layers2, bidirectionalTrue gru nn.GRU(input_size, hidden_size, num_layers2, bidirectionalTrue, batch_firstTrue) x torch.randn(batch, seq_len, input_size) output, hn gru(x) print(output.shape) # (2, 5, 6) hidden_size3双向拼成 6 print(hn.shape) # (4, 2, 3) 2层 x 2方向 4做分类时如果想把双向信息合并常见写法是torch.cat([hn[-2], hn[-1]], dim1)得到(batch, 2 * hidden_size)。注意这里取的是hn[-2]和hn[-1]也就是最后一层的正向和反向而不是第 0 层的。写hn[0]和hn[1]也能跑但拿到的是第一层的输出信息没经过第二层加工效果上会有差距而且这个错误不会报任何异常只能从最终指标上看出来。4. 手写一版 GRU 与官方实现对拍4.1 从权重形状反推门控顺序要真正搞懂输入输出的含义最有效的办法是用纯张量运算手写一版 GRU然后和官方实现对拍。先看权重形状gru nn.GRU(input_size4, hidden_size3) print(gru.weight_ih_l0.shape) # torch.Size([9, 4]) 3*hidden, input print(gru.weight_hh_l0.shape) # torch.Size([9, 3]) 3*hidden, hidden print(gru.bias_ih_l0.shape) # torch.Size([9]) print(gru.bias_hh_l0.shape) # torch.Size([9])9 就是3 * hidden_size。为什么是 3 倍因为 GRU 有三个门重置门 r、更新门 z、候选状态 n。PyTorch 把这三组参数在输出维度上拼成了一个大矩阵而不是分成三个独立的Linear这样一次矩阵乘法就能算完三个门效率更高。关键问题是这 9 行里哪 3 行属于哪个门顺序是r、z、n不是很多人以为的 z、r、n。这个顺序记住就行因为在第 4.2 节的代码里一旦顺序搞错结果会和官方实现对不上而且是那种误差不大但就是不收敛的错法。4.2 逐时间步复现的完整代码GRU 的更新公式是重置门r_t sigmoid(x_t W_ir.T b_ir h W_hr.T b_hr)更新门z_t sigmoid(x_t W_iz.T b_iz h W_hz.T b_hz)候选态n_t tanh(x_t W_in.T b_in r_t * (h W_hn.T b_hn))新隐状态h_t (1 - z_t) * n_t z_t * h_{t-1}最后一步的h_t (1 - z) * n z * h是 GRU 最精妙的地方。更新门 z 输出的每个元素都在 0 到 1 之间它像一个刻度盘控制着新信息和旧记忆各占多少。当 z 接近 0网络更倾向于接受新的候选状态相当于把当前输入写进记忆当 z 接近 1网络几乎保留原来的隐状态不动相当于选择性地遗忘新输入。这种机制让 GRU 在长序列上比朴素 RNN 更容易保住早期信息。import torch import torch.nn as nn torch.manual_seed(0) input_size, hidden_size, seq_len, batch 4, 3, 5, 2 gru nn.GRU(input_size, hidden_size) x torch.randn(seq_len, batch, input_size) h0 torch.zeros(1, batch, hidden_size) output, hn gru(x, h0) # 按 r, z, n 的顺序切分权重 W_ir, W_iz, W_in gru.weight_ih_l0.chunk(3, dim0) W_hr, W_hz, W_hn gru.weight_hh_l0.chunk(3, dim0) b_ir, b_iz, b_in gru.bias_ih_l0.chunk(3) b_hr, b_hz, b_hn gru.bias_hh_l0.chunk(3) h h0[0] outs [] for t in range(seq_len): xt x[t] r torch.sigmoid(xt W_ir.T b_ir h W_hr.T b_hr) z torch.sigmoid(xt W_iz.T b_iz h W_hz.T b_hz) n torch.tanh(xt W_in.T b_in r * (h W_hn.T b_hn)) h (1 - z) * n z * h outs.append(h) manual_out torch.stack(outs, dim0) print(torch.allclose(manual_out, output, atol1e-6)) # True print(torch.allclose(h, hn[0], atol1e-6)) # True4.3 结果对齐与数值误差说明跑出来allclose是True说明对输入输出的理解是对的。这里用atol1e-6而不是完全相等原因是浮点运算的累加顺序不同会带来末位差异官方实现内部可能做了融合或者换了个加法次序属于正常现象。如果误差在1e-5量级以内基本可以确认逻辑没错如果差到1e-2级别大概率是门的顺序搞反了或者某个 bias 漏加了。这个对拍实验的价值在于它证明了一件事nn.GRU没有任何魔法它就是一个逐时间步的循环只不过内部用 cuDNN 或 fused kernel 做了加速。你理解了这一层再去看output和h_n的关系就不会再觉得它们是两个独立的东西h_n只是output在特定位置上的取值而已。注意手写版本只是为了理解和验证实际训练务必用官方实现。手写的 Python 循环在 GPU 上慢几十倍而且不会自动利用cudnn的优化路径。5. 实战里绕不开的几个变体5.1 多层双向的维度换算练习多层双向的维度换算有个固定套路第 k 层的input_size等于第 k-1 层输出的最后一维。第 0 层是input_size之后每层的输入都是num_directions * hidden_size。gru nn.GRU(input_size4, hidden_size3, num_layers2, bidirectionalTrue, batch_firstTrue) x torch.randn(2, 5, 4) output, hn gru(x) # 第 0 层: 输入 4 - 输出 6 (双向各 3 拼起来) # 第 1 层: 输入 6 - 输出 6 print(output.shape) # (2, 5, 6) print(hn.shape) # (4, 2, 3)注意output的最后一维在每一层都是 6所以从外面看永远是 6看不出层数。层数信息只体现在hn的第一维上。这也是为什么很多人搞不清第几层的输出怎么取——output只暴露最后一层的结果中间层的结果是拿不到的除非你在每层后面插 hook 或者手动逐层跑。如果确实需要中间层表示一个实用做法是把多层 GRU 拆开、每层单独跑然后自己保存每层的output。代价是失去 cuDNN 的整图优化但换来的是灵活的特征提取能力做多层特征融合时经常这么干。5.2 变长序列与 padding 处理真实数据里句子长度不一通常的做法是 padding 到同一个长度。这时候如果直接用output[-1]取最后一步取到的可能是 padding 位置的输出那个位置的信息毫无意义。正确做法是用pack_padded_sequence告诉模型哪些位置是真实数据。from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence gru nn.GRU(input_size4, hidden_size3, batch_firstTrue) lengths torch.tensor([5, 3, 1]) x torch.randn(3, 5, 4) # 全部 padding 到长度 5 packed pack_padded_sequence(x, lengths, batch_firstTrue, enforce_sortedFalse) packed_out, hn gru(packed) out, out_lengths pad_packed_sequence(packed_out, batch_firstTrue) print(out.shape) # (3, 5, 3)重新补齐成统一长度 print(out_lengths) # tensor([5, 3, 1]) print(hn.shape) # (1, 3, 3)用pack之后hn里的每个位置都是各样本真实序列走完后的状态不受 padding 影响这一点是它比手动取output[-1]更可靠的地方。用enforce_sortedFalse可以让内部自动按长度排序省去自己排的麻烦但要注意返回的隐状态顺序和原始 batch 的下标可能不再一一对应如果后续计算 loss 需要对齐标签这点必须确认清楚。我自己的习惯是把长度按降序排好再传进去enforce_sorted保持默认这样顺序从头到尾都是确定的调试时少一类幽灵 bug。5.3 取最后一步的三种写法对比写法适用场景风险点output[-1]单向、无 padding 或长度整齐双向或多层时语义不对output[:, -1, :]batch_firstTrue的单向同上hn[-1]/hn.permute(1,0,2)单向多层取最后一层双向需再拼一次torch.cat([hn[-2], hn[-1]], 1)双向取最后一层索引写错就取到浅层这四种写法在单向单层时结果一样但一旦结构变了正确性就分岔了。我一般会写一个get_last_hidden的小函数把双向判断、batch_first 处理都封在里面这样模型结构改动时只改一处不用满项目找output[-1]。6. 常见报错与排查速查表6.1 维度类报错与对应原因报错信息关键字最可能的原因修复方式Expected hidden size (1, N, H)h0第一维写成了 batch改成(num_layers*dir, batch, H)input.size(-1) must be equal to input_size特征维和构造时不一致检查x.shape[-1]Expected 3D input输入是 2D缺了时间步或 batch用x.unsqueeze(0)补维Input and hidden tensors are not the same dtypeh0用默认 float32输入是 float64用x.new_zeros创建dropout expects num_layers greater than 1单层传了非零 dropout层数设为 2 或把 dropout 设为 0维度类报错里最值得警惕的是能跑通但结果错的那一类。比如输入维度全对但序列长度和 batch 反了或者用了batch_firstTrue却按旧格式构造输入。这类问题不会抛异常只会让模型学得莫名其妙。我的做法是在forward开头固定加一行断言把期望的形状写死出错时第一时间定位而不是等 loss 曲线不对劲再回头查。6.2 数值与训练类问题排查梯度爆炸在 RNN 系里是老问题。GRU 因为有更新门的门控机制比朴素 RNN 稳很多但层数一多、学习率一大还是会炸。判断信号很直接loss 突然跳到nan或者变成一个巨大的数。处理手段是torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)这一行基本是标配我自己写序列模型时从不会漏掉。另一个常见现象是loss 降得下去但验证集不涨。这时候先怀疑的不是模型结构而是h_n的取法。分类任务里如果误取了浅层的h_n[0]而不是最后一层模型依然能训但表达能力被截断了一部分。可以打印一下hn.shape和output.shape确认两者的维度关系和预期一致再继续调别的超参数。还有一类是 dtype 混用导致的隐性降速。如果h0是 float64 而模型是 float32PyTorch 会做类型提升整个前向变成 double 精度速度掉一半以上而且不会有任何警告。这个坑在从 numpy 转数据的时候特别容易踩因为np.zeros默认就是 float64。统一用torch.zeros(..., dtypetorch.float32)或者x.new_zeros(...)能规避掉。7. 我自己的一些实操体会写了这么多序列模型我对nn.GRU的输入输出有个比较实用的心法把input想成外面送进来的原料h_0想成开工前给的初始库存output是每个时间点盘点的库存流水h_n是收工时的期末库存。这个类比基本上能覆盖所有形状推导因为库存流水有时间维度、有批次维度期末库存没有时间维度只有批次和层数方向。真正让我少走了很多弯路的习惯是在每个模型的forward里把四个形状写成注释并且用assert把最容易出错的假设固定下来。例如双向分类模型里固定断言output.size(-1) 2 * self.hidden_size一旦有人改了bidirectional参数而忘记改后续拼接代码第一时间就会炸出来而不是等到准确率莫名其妙下降两个点再去追。这类提前失败的写法比任何事后调试都省时间。最后分享一个小细节。如果你只需要序列的最终表示而且序列长度整齐直接传None当h_0是最省事的PyTorch 内部会做零初始化还省了一次显式分配。但如果你的 batch 里混着不同长度的 padding 数据就别偷这个懒老老实实用pack_padded_sequence配合显式的长度张量。这两条路我都走过图省事的那次最终在验证集上多花了一天时间排查教训挺深的。
返回列表