更新优化器

在模型训练中,我们经常会遇到参数数量动态变化的场景。像 3D Gaussian Splatting 中对点云进行克隆和分裂时,模型参数的形状(Shape)会发生变化。

参数变了,负责更新参数的“大管家”——优化器(Optimizer),自然也需要进行相应的调整。如果直接用旧的优化器去更新新维度的参数,程序当场就会因为维度不匹配而崩溃报错。

在这篇文章中,我们就来实现 update_optimizer_state 函数,使得在保留“前人记忆”的同时,平滑过渡到新架构。


优化器底层参数

我们先弄懂 PyTorch 优化器(如 torch.optim.Adam)内部三个核心概念。我们可以把优化器想象成一个“档案管理处”,专门记录和管理参数该如何更新。

1. optimizer.param_groups

它是一个列表(List),里面装的是一个个字典(Dict)。

模型里的参数经常被分成不同的“组”,以便享受不同的待遇(比如特征提取层学习率小,全连接层学习率大)。这个列表就是用来管理这些分组配置的。

  1. # param_groups 的结构示例
  2. [
  3.     {
  4.         'params': [Parameter_1, Parameter_2], # 属于这组的 Tensor 对象
  5.         'lr': 0.001,                          # 这组参数的学习率
  6.         'weight_decay': 0.01,                 # 这组参数的权重衰减
  7.     },
  8.     {
  9.         'params': [Parameter_3],
  10.         'lr': 0.01,                           # 另一组参数使用不同的学习率
  11.         'weight_decay': 0.0,
  12.     }
  13. ]

2. optimizer.state

它是一个巨大的字典(Dict)。我们可以把它当作优化器的“核心记忆库”,用来存所有参数的历史训练状态。

它的特殊之处在于:它的 Key(键)不是字符串,而是具体的“参数 Tensor 对象本身”。

3. optimizer.state[param]

这就相当于精准查字典。去记忆库里,把某一个具体参数(比如 param)的私人历史档案给调出来。拿到的会是这样一个包含动量和步数的小字典:

  1. {
  2.     'step': 100,               # 该参数已经更新了 100 次
  3.     'exp_avg': Tensor([...]),  # 积攒的一阶动量 (形状与参数本身一致)
  4.     'exp_avg_sq': Tensor([...])# 积攒的二阶动量 (形状与参数本身一致)
  5. }

代码实现

有了上面的基础,我们再来实现代码就会显得非常直观了。这个函数的核心功能就是:在训练过程中改变模型结构时,如何完美迁移 Adam 优化器的动量状态。

  1. @torch.no_grad()
  2. def update_optimizer_state(optimizer, new_parameters, map_state_fn):
  3.     """
  4.     Recreate the optimizer with new parameters and transfer/map the states.
  5.     """
  6.     new_optimizer = makeOptimizer(new_parameters)
  7.     for old_group, new_group in zip(optimizer.param_groups, new_optimizer.param_groups):
  8.         old_param = old_group["params"][0]
  9.         new_param = new_group["params"][0]
  10.         if old_param in optimizer.state:
  11.             old_state = optimizer.state[old_param]
  12.             new_state = {}
  13.             if "step" in old_state:
  14.                 new_state["step"] = old_state["step"].clone() if isinstance(old_state["step"], torch.Tensor) else old_state["step"]
  15.             if "exp_avg" in old_state:
  16.                 new_state["exp_avg"] = map_state_fn(old_state["exp_avg"], new_param)
  17.             if "exp_avg_sq" in old_state:
  18.                 new_state["exp_avg_sq"] = map_state_fn(old_state["exp_avg_sq"], new_param)
  19.             new_optimizer.state[new_param] = new_state
  20.     return new_optimizer

代码里直接写成 ["params"][0],是因为我们直接把每个 Tensor 独立当成一个组来分别管理。

step (步数 / t):代表这个参数被更新了多少次。在较新版本的 PyTorch 中,为了 GPU 加速,它通常是一个一维 Tensor;在老版本中可能只是个整数 int。所以代码严谨地使用了 isinstance(..., torch.Tensor) 做判断。

exp_avg (一阶动量 / m_t):梯度的一阶导数(动量)。它的 Shape 与其绑定的参数完全一模一样。里面存的是历史梯度的移动平均值。

exp_avg_sq (二阶动量 / v_t):梯度的二阶导数。Shape 同样与参数一模一样。它衡量梯度的波动程度,Adam 优化器用它来实现学习率的自适应缩放。

既然 exp_avgexp_avg_sq 的 Shape 死死绑定在参数上,当新参数 new_param 的形状改变时,旧的动量矩阵无法直接套用。

这就使用到了 map_state_fn,以高斯克隆流程为例:

  1. @torch.no_grad()
  2. def clone_gaussians(mask_clone, parameters, optimizer):
  3.     """
  4.     Clone selected Gaussians.
  5.     """
  6.     if not mask_clone.any():
  7.         return parameters, optimizer
  8.  
  9.     for name, param in parameters.items():
  10.         old_val = param.detach()
  11.         cloned_val = old_val[mask_clone]
  12.         new_val = torch.cat([old_val, cloned_val], dim=0)
  13.         parameters[name] = torch.nn.Parameter(new_val, requires_grad=True)
  14.  
  15.     def map_state_fn(old_state, new_param):
  16.         new_state = torch.zeros_like(new_param)
  17.         new_state[:old_state.shape[0]] = old_state
  18.         return new_state
  19.  
  20.     new_optimizer = update_optimizer_state(optimizer, parameters, map_state_fn)
  21.     return parameters, new_optimizer

因为我们旧的参数放前面,新的参数放后面。所以只需要把旧参数的状态复制到前面,后面的参数初始化为零。