by_group: Dict[int, List[UpdateUnit]] = defaultdict(list)
for unit in batch:
by_group[unit.adapter_index].append(unit)
for group_index, units in by_group.items():
self._step_group(self.optimizer.param_groups[group_index], units)
def _step_group(self, group: Dict[str, Any], units: List[UpdateUnit]) -> None:
"""Run one group's parameters through the matching functional Adam/AdamW."""
params, grads, exp_avgs, exp_avg_sqs, max_exp_avg_sqs, state_steps = (
self._collect_group_args(units, group)
)
if not params:
return
if self.is_new_adamw:
self._step_new_adamw(group, params, grads, exp_avgs, exp_avg_sqs, max_exp_avg_sqs)
return
func = getattr(torch.optim._functional, self.functional_name)
kwargs = {
"amsgrad": group["amsgrad"],
"beta1": group["betas"][0],
"beta2": group["betas"][1],
"lr": group["lr"],