Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 9 additions & 6 deletions src/accelerate/accelerator.py
Original file line number Diff line number Diff line change
Expand Up @@ -1646,12 +1646,15 @@ def _get_tensor_address(p):

for obj in result:
if isinstance(obj, torch.optim.Optimizer):
for param_group in obj.param_groups:
# Each param_group originally maps to model parameters (e.g., from model.parameters()).
# After _prepare_tp(), parameter references are replaced with DTensor instances.
# Therefore, we remap the parameter references to their new DTensor addresses
# so that the optimizer can correctly update the model parameters.
param_group["params"] = [mapping[_get_tensor_address(p)] for p in param_group["params"]]
# Each param_group originally maps to model parameters (e.g., from model.parameters()).
# After _prepare_tp(), parameter references are replaced with DTensor instances.
# Therefore, we remap the parameter references and their optimizer states to the new DTensors.
parameters_map = {
p: mapping[_get_tensor_address(p)]
for param_group in obj.param_groups
for p in param_group["params"]
}
obj._switch_parameters(parameters_map)

return result

Expand Down
3 changes: 3 additions & 0 deletions src/accelerate/optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -183,6 +183,9 @@ def step(self, closure=None):
def _switch_parameters(self, parameters_map):
for param_group in self.optimizer.param_groups:
param_group["params"] = [parameters_map.get(p, p) for p in param_group["params"]]
for old_parameter, new_parameter in parameters_map.items():
if old_parameter is not new_parameter and old_parameter in self.optimizer.state:
self.optimizer.state[new_parameter] = self.optimizer.state.pop(old_parameter)

@property
def step_was_skipped(self):
Expand Down
16 changes: 16 additions & 0 deletions tests/test_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,22 @@ def test_accelerated_optimizer_pickling(self):
except Exception as e:
self.fail(f"Accelerated optimizer pickling failed with {e}")

def test_switch_parameters_updates_optimizer_state(self):
model = torch.nn.Linear(10, 10)
optimizer = torch.optim.Adagrad(model.parameters(), 0.1)
accelerator = Accelerator()
optimizer = accelerator.prepare(optimizer)

old_parameters = list(optimizer.param_groups[0]["params"])
old_states = [optimizer.state[parameter] for parameter in old_parameters]
new_parameters = [torch.nn.Parameter(torch.empty_like(parameter)) for parameter in old_parameters]
optimizer._switch_parameters(dict(zip(old_parameters, new_parameters)))

assert all(actual is expected for actual, expected in zip(optimizer.param_groups[0]["params"], new_parameters))
assert set(optimizer.state) == set(new_parameters)
assert all(optimizer.state[parameter] is state for parameter, state in zip(new_parameters, old_states))
optimizer.state_dict()


@require_fp16
@require_non_cpu
Expand Down