当前位置: 首页 > 图灵资讯 > 行业资讯> Python中PyTorch如何对嵌入层Embedding进行权重共享?

Python中PyTorch如何对嵌入层Embedding进行权重共享?

来源:图灵python
时间: 2026-09-03 16:16:14
直接赋值 lm_head.weight = embedding_layer.weight 因两者而合法有效 shape 均为 [V, d] 且 lm_head 无 bias,PyTorch 中 nn.Linear 权重存储为 [out_features, in_features],与 nn.Embedding 的 [vocab_size, embedding_dim] 自然对齐,赋值后共享内存和梯度。

直接赋值 lm_head.weight = embedding_layer.weight 好吧,前提是两者 shape 完全一致([V, d]),且 lm_head 不带 bias。

为什么 lm_head.weight = embedding_layer.weight 合法有效?

PyTorch 的 nn.Embeddingnn.Linear 在权重 shape 天然对齐:天然对齐:

  • embedding_layer.weight 形状是 [vocab_size, embedding_dim](即 [V, d]
  • lm_head = nn.Linear(embedding_dim, vocab_size, bias=False)weight 形状也是 [vocab_size, embedding_dim](PyTorch 中 Linear 存储为 [out_features, in_features]
  • 因此,赋值后,两者指向相同的内存;当反向传输时,梯度会自动累积到相同的内存中 .grad
  • 调用 lm_head(hidden_states) 实际执行是 hidden_states @ embedding_layer.weight.T,语义上是“计算隐藏状态和各种隐藏状态” token embedding 的点积”
必须避免的三个典型错误

在实践中,由于细节不一致,很容易导致共享失效或报错:

PyTorch Linux版本 2.11.0

PyTorch 2.11.0 下载历史版,来自 PyPI 适用于旧项目兼容、实验复制和指定环境安装的官方发布。

下载

  • lm_head 加了 bias=True:权重 shape 变成 [V, d] + [V],无法和 embedding_layer.weight 对齐;一定要设置 bias=False
  • 逆转初始化顺序:先 self.lm_head = ...self.lm_head.weight = self.wte.weight;不能反过来(wte 在创建之前引用)
  • forward 内部动态赋值:例如写成 lm_head.weight.data = embedding_layer.weight.data —— 这只是复制值,不共享梯度;必须在那里 __init__ 中等对象级引用
transformers 加载库时如何保持权重共享?

HF 的 transformers 模型(如 GPT2LMHeadModel)默认启用权重共享,但加载逻辑有一个隐含的协议:

立即学习“Python免费学习笔记(深入);

  • 保存时只保存一份 embedding 权重(model.state_dict()["transformer.wte.weight"]),不单独存 lm_head.weight
  • 加载时跳过 lm_head.weight 加载步骤,然后在 post_init_tie_weights 中执行 self.lm_head.weight = self.transformer.wte.weight
  • 假如你手动修改 state dict,漏掉这个绑定,模型会退化为两套独立权重——参数翻倍,收敛慢,logits 偏移

真正的关键不是“如何写那行赋值”,而是确保整个生命周期从初始化、训练到保存和加载 embedding_layer.weightlm_head.weight 总是一样的 Python 对象。任何中间环节的深度复制,.data 赋值或分步加载可能会悄然破坏这种引用关系。