更新优化器
在模型训练中,我们经常会遇到参数数量动态变化的场景。像 3D Gaussian Splatting 中对点云进行克隆和分裂时,模型参数的形状(Shape)会发生变化。
参数变了,负责更新参数的“大管家”——优化器(Optimizer),自然也需要进行相应的调整。如果直接用旧的优化器去更新新维度的参数,程序当场就会因为维度不匹配而崩溃报错。
在这篇文章中,我们就来实现 update_optimizer_state 函数,使得在保留“前人记忆”的同时,平滑过渡到新架构。
优化器底层参数
我们先弄懂 PyTorch 优化器(如 torch.optim.Adam)内部三个核心概念。我们可以把优化器想象成一个“档案管理处”,专门记录和管理参数该如何更新。
1. optimizer.param_groups
它是一个列表(List),里面装的是一个个字典(Dict)。
模型里的参数经常被分成不同的“组”,以便享受不同的待遇(比如特征提取层学习率小,全连接层学习率大)。这个列表就是用来管理这些分组配置的。
- # param_groups 的结构示例
- [
- {
- 'params': [Parameter_1, Parameter_2], # 属于这组的 Tensor 对象
- 'lr': 0.001, # 这组参数的学习率
- 'weight_decay': 0.01, # 这组参数的权重衰减
- },
- {
- 'params': [Parameter_3],
- 'lr': 0.01, # 另一组参数使用不同的学习率
- 'weight_decay': 0.0,
- }
- ]
2. optimizer.state
它是一个巨大的字典(Dict)。我们可以把它当作优化器的“核心记忆库”,用来存所有参数的历史训练状态。
它的特殊之处在于:它的 Key(键)不是字符串,而是具体的“参数 Tensor 对象本身”。
3. optimizer.state[param]
这就相当于精准查字典。去记忆库里,把某一个具体参数(比如 param)的私人历史档案给调出来。拿到的会是这样一个包含动量和步数的小字典:
- {
- 'step': 100, # 该参数已经更新了 100 次
- 'exp_avg': Tensor([...]), # 积攒的一阶动量 (形状与参数本身一致)
- 'exp_avg_sq': Tensor([...])# 积攒的二阶动量 (形状与参数本身一致)
- }
代码实现
有了上面的基础,我们再来实现代码就会显得非常直观了。这个函数的核心功能就是:在训练过程中改变模型结构时,如何完美迁移 Adam 优化器的动量状态。
- @torch.no_grad()
- def update_optimizer_state(optimizer, new_parameters, map_state_fn):
- """
- Recreate the optimizer with new parameters and transfer/map the states.
- """
- new_optimizer = makeOptimizer(new_parameters)
- for old_group, new_group in zip(optimizer.param_groups, new_optimizer.param_groups):
- old_param = old_group["params"][0]
- new_param = new_group["params"][0]
- if old_param in optimizer.state:
- old_state = optimizer.state[old_param]
- new_state = {}
- if "step" in old_state:
- new_state["step"] = old_state["step"].clone() if isinstance(old_state["step"], torch.Tensor) else old_state["step"]
- if "exp_avg" in old_state:
- new_state["exp_avg"] = map_state_fn(old_state["exp_avg"], new_param)
- if "exp_avg_sq" in old_state:
- new_state["exp_avg_sq"] = map_state_fn(old_state["exp_avg_sq"], new_param)
- new_optimizer.state[new_param] = new_state
- return new_optimizer
代码里直接写成 ["params"][0],是因为我们直接把每个 Tensor 独立当成一个组来分别管理。
step (步数 /
exp_avg (一阶动量 /
exp_avg_sq (二阶动量 /
既然 exp_avg 和 exp_avg_sq 的 Shape 死死绑定在参数上,当新参数 new_param 的形状改变时,旧的动量矩阵无法直接套用。
这就使用到了 map_state_fn,以高斯克隆流程为例:
- @torch.no_grad()
- def clone_gaussians(mask_clone, parameters, optimizer):
- """
- Clone selected Gaussians.
- """
- if not mask_clone.any():
- return parameters, optimizer
- for name, param in parameters.items():
- old_val = param.detach()
- cloned_val = old_val[mask_clone]
- new_val = torch.cat([old_val, cloned_val], dim=0)
- parameters[name] = torch.nn.Parameter(new_val, requires_grad=True)
- def map_state_fn(old_state, new_param):
- new_state = torch.zeros_like(new_param)
- new_state[:old_state.shape[0]] = old_state
- return new_state
- new_optimizer = update_optimizer_state(optimizer, parameters, map_state_fn)
- return parameters, new_optimizer
因为我们旧的参数放前面,新的参数放后面。所以只需要把旧参数的状态复制到前面,后面的参数初始化为零。