延后初始化(Deferred Initialization)


一、为什么需要延后初始化?

在构建神经网络时,通常我们需要给参数(权重、偏置)分配内存并初始化。但是有时候:

  • 我们在定义 nn.Linearnn.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 前尝试访问参数,会发现是未初始化状态,会报错。
  • 适合输入维度未知的情况,如果输入固定,还是建议直接指定输入输出维度。
  • LazyLinearLazyConv2d 等模块就是专门为此设计的。

📌 总结一句话:
PyTorch 的延后初始化 = 层在创建时先不分配参数,等第一次 forward 时根据输入数据形状自动分配并初始化参数



内容声明
本文基于开源教材《动手学深度学习》(Dive into Deep Learning, 作者:Aston Zhang、Zachary C. Lipton、Mu Li、Alexander J. Smola 等)整理,原始项目地址:https://github.com/d2l-ai/d2l-zh。
在整理过程中对部分内容进行了删改和补充,仅用于个人学习与交流,版权归原作者所有。


Logo

开源鸿蒙跨平台开发社区汇聚开发者与厂商,共建“一次开发,多端部署”的开源生态,致力于降低跨端开发门槛,推动万物智联创新。

更多推荐