【动手学深度学习PyTorch】神经网络的延后初始化(Deferred Initialization)
·
延后初始化(Deferred Initialization)
一、为什么需要延后初始化?
在构建神经网络时,通常我们需要给参数(权重、偏置)分配内存并初始化。但是有时候:
- 我们在定义
nn.Linear、nn.Conv2d等层的时候,并 不知道输入数据的形状(尤其是 batch 大小或特征维度)。 - 如果在创建时就立刻分配参数,必须明确输入维度;但在一些场景下(比如自定义网络、可变输入长度)我们希望推迟到真正收到数据时再分配。
延后初始化 就是指:PyTorch 在创建网络层的时候,可以先只保存层的结构信息(比如输出维度),而不立即分配参数;等到第一次前向传播时,根据输入数据的实际形状,才真正初始化参数。
二、直观例子
import torch
from torch import nn
net = nn.Sequential(
nn.Linear(20, 256),
nn.ReLU(),
nn.Linear(256, 10)
)
print(net)
这里的 nn.Linear(20, 256) 一开始就知道输入是 20 维,所以会立即初始化参数。
但如果你用 自定义网络,有时输入维度是不确定的,比如:
class MyNet(nn.Module):
def __init__(self, num_outputs):
super().__init__()
self.dense = None # 先不定义输入大小的层
self.num_outputs = num_outputs
def forward(self, x):
if self.dense is None: # 第一次看到输入数据时再初始化
self.dense = nn.Linear(x.shape[1], self.num_outputs)
return self.dense(x)
运行时:
X = torch.randn(5, 20) # 输入是 20 维
net = MyNet(10)
print(net(X).shape) # (5, 10)
➡️ 这里的 self.dense 是在 第一次 forward 时才根据 X.shape[1] (=20) 初始化的。
三、PyTorch 内置延后初始化的地方
其实你不写 if self.dense is None 也能做到,因为很多 PyTorch 层(比如 nn.LazyLinear, nn.LazyConv2d)内置了 Lazy Modules,就是延后初始化的实现:
lazy_linear = nn.LazyLinear(10) # 只指定输出 10
print(lazy_linear.weight) # 还没初始化
此时参数是 未初始化的张量。只有当你第一次给它输入时:
X = torch.randn(5, 20)
out = lazy_linear(X) # 这里自动推断输入维度是20
print(out.shape) # (5, 10)
print(lazy_linear.weight.shape) # (10, 20)
才会真正分配 weight(10, 20)。
四、延后初始化的好处
- 更灵活:不需要在写模型时死板地知道所有输入维度。
- 方便构建动态网络:比如 RNN/Transformers 里输入维度可变,延后初始化可以自动适配。
- 简化代码:避免手动根据输入形状去写参数大小。
五、普通 Linear 和 LazyLinear 对比实验
import torch
from torch import nn
# -------------------------------
# 1. 普通 Linear:初始化时就分配参数
# -------------------------------
linear = nn.Linear(20, 10) # 输入20,输出10
print("普通 Linear 初始化时:")
print("weight:", linear.weight.shape)
print("bias:", linear.bias.shape)
# -------------------------------
# 2. LazyLinear:延后初始化
# -------------------------------
lazy_linear = nn.LazyLinear(10) # 只指定输出,输入维度先不写
print("\nLazyLinear 初始化时:")
print("weight:", lazy_linear.weight) # 未初始化
print("bias:", lazy_linear.bias) # 未初始化
# -------------------------------
# 3. 第一次 forward:触发延后初始化
# -------------------------------
X = torch.randn(5, 20) # 输入batch=5, 特征维度=20
out = lazy_linear(X) # 这里才初始化
print("\n第一次 forward 之后:")
print("output:", out.shape)
print("weight:", lazy_linear.weight.shape)
print("bias:", lazy_linear.bias.shape)
普通 Linear 初始化时:
weight: torch.Size([10, 20])
bias: torch.Size([10])
LazyLinear 初始化时:
weight: <UninitializedParameter>
bias: <UninitializedParameter>
第一次 forward 之后:
output: torch.Size([5, 10])
weight: torch.Size([10, 20])
bias: torch.Size([10])
六、延后初始化的注意点
- 如果你在 forward 前尝试访问参数,会发现是未初始化状态,会报错。
- 适合输入维度未知的情况,如果输入固定,还是建议直接指定输入输出维度。
LazyLinear、LazyConv2d等模块就是专门为此设计的。
📌 总结一句话:
PyTorch 的延后初始化 = 层在创建时先不分配参数,等第一次 forward 时根据输入数据形状自动分配并初始化参数。
内容声明
本文基于开源教材《动手学深度学习》(Dive into Deep Learning, 作者:Aston Zhang、Zachary C. Lipton、Mu Li、Alexander J. Smola 等)整理,原始项目地址:https://github.com/d2l-ai/d2l-zh。
在整理过程中对部分内容进行了删改和补充,仅用于个人学习与交流,版权归原作者所有。
更多推荐


所有评论(0)