You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在保留已有参数的情况下修改torch.nn.Linear的输出维度?

解决PyTorch Linear层动态扩展输出维度并保留训练参数的问题

核心思路

PyTorch的Linear层无法直接修改out_features属性,我们可以通过拼接原有训练参数与新初始化参数的方式,创建新的Linear层替换原层,既能完整保留已训练的权重和偏置,又能让新增参数遵循标准初始化逻辑。

具体实现代码

假设你需要将self.basis_mat的输出维度扩展到new_out_features,可以按以下步骤操作:

# 1. 提取原Linear层的参数与设备信息
old_weight = self.basis_mat.weight.data
has_bias = self.basis_mat.bias is not None
old_bias = self.basis_mat.bias.data if has_bias else None
device = old_weight.device

# 2. 设置目标输出维度,计算需要新增的维度数量
new_out_features = 512  # 替换为你实际需要的目标输出维度
delta_out = new_out_features - self.basis_mat.out_features

# 3. 初始化新增参数(与PyTorch Linear层默认初始化逻辑对齐)
# 权重采用Kaiming均匀初始化,和Linear层默认初始化一致
new_weight = torch.nn.init.kaiming_uniform_(torch.empty(delta_out, old_weight.size(1), device=device))
# 偏置默认初始化为0(仅当原层包含偏置时)
new_bias = torch.zeros(delta_out, device=device) if has_bias else None

# 4. 拼接原有参数与新增参数
combined_weight = torch.cat([old_weight, new_weight], dim=0)
combined_bias = torch.cat([old_bias, new_bias], dim=0) if has_bias else None

# 5. 创建新的Linear层并赋值拼接后的参数
new_basis_mat = torch.nn.Linear(
    in_features=self.basis_mat.in_features,
    out_features=new_out_features,
    bias=has_bias
).to(device)

new_basis_mat.weight.data = combined_weight
if has_bias:
    new_basis_mat.bias.data = combined_bias

# 6. 替换原层
self.basis_mat = new_basis_mat

关键注意事项

  • 设备一致性:确保新增参数和原参数在同一设备(CPU/GPU)上,避免张量不匹配错误。
  • 初始化匹配:新增权重的初始化逻辑要和原Linear层保持一致,PyTorch默认用kaiming_uniform_初始化权重、zeros_初始化偏置,这样能保证训练的稳定性。
  • 梯度继承:替换后的新层会自动参与后续反向传播,原有参数的梯度会正常更新,新增参数也会从随机初始化状态开始训练。

内容的提问来源于stack exchange,提问作者Shiran Yuan

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.18 19:03:26