Skip to content

Commit e5bda62

Browse files
authored
[CherryPick][DCP] Fix Optimizer Learning Rate not being loaded correctly (#129398) (#129683)
[DCP] Fix Optimizer Learning Rate not being loaded correctly (#129398) Fixes #129079 Currently, the tensor object is loading correctly in-place, but the non-tensor object such as learning rate is not load correctly after f518cf8, which is a regression introduced in 2.3. This PR replaces tree_map_only and manual replacement of the state dict items with _tree_map_only and fixes the regression of non-tensor loading. Test: ``` python3 test/distributed/checkpoint/e2e/test_e2e_save_and_load.py -k test_init_state_dict python3 test/distributed/checkpoint/test_tp_checkpoint.py -k test_tp_checkpoint_load_on_meta_device ``` Pull Request resolved: #129398 Approved by: https://github.com/fegin (cherry picked from commit 8b8e2fc)
1 parent 705e3ae commit e5bda62

2 files changed

Lines changed: 46 additions & 8 deletions

File tree

‎test/distributed/checkpoint/e2e/test_e2e_save_and_load.py‎

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,9 @@
1818
_patch_model_state_dict,
1919
_patch_optimizer_state_dict,
2020
get_model_state_dict,
21+
get_optimizer_state_dict,
2122
get_state_dict,
23+
set_state_dict,
2224
)
2325
from torch.distributed.checkpoint.state_dict_loader import _load_state_dict_from_keys
2426
from torch.distributed.checkpoint.utils import CheckpointException
@@ -417,6 +419,48 @@ def test_no_cpu(self):
417419
f.result()
418420

419421

422+
class TestInitStateDict(DTensorTestBase):
423+
@with_temp_dir
424+
def test_init_state_dict(self):
425+
temp_dir = self.temp_dir
426+
model = TestDummyModel()
427+
optim = torch.optim.Adam(model.parameters(), lr=0.1)
428+
429+
state_dict_to_save = {
430+
"model": get_model_state_dict(model),
431+
"optimizer": get_optimizer_state_dict(model, optim),
432+
}
433+
DCP.save(state_dict_to_save, checkpoint_id=temp_dir)
434+
435+
torch.manual_seed(0)
436+
model_2 = TestDummyModel()
437+
# Changing the learning rate for optimizer, which is not a tensor.
438+
optim_2 = torch.optim.Adam(model_2.parameters(), lr=0.2)
439+
440+
msd = get_model_state_dict(model_2)
441+
osd = get_optimizer_state_dict(model_2, optim_2)
442+
443+
state_dict_to_load = {"model": msd, "optimizer": osd}
444+
DCP.load(state_dict_to_load, checkpoint_id=temp_dir)
445+
446+
# We need to check that the two variables point to the same object in memory,
447+
# since we claim DCP is in-place loading.
448+
self.assertTrue(msd is state_dict_to_load["model"])
449+
self.assertTrue(osd is state_dict_to_load["optimizer"])
450+
451+
# set_state_dict calls load_state_dict for model and optimizer.
452+
# so we should see the optim_2.param_groups learning rate is 0.1 instead of 0.2 now.
453+
set_state_dict(
454+
model_2,
455+
optim_2,
456+
model_state_dict=state_dict_to_load["model"],
457+
optim_state_dict=state_dict_to_load["optimizer"],
458+
)
459+
self.assertEqual(msd, get_model_state_dict(model_2))
460+
self.assertEqual(osd, get_optimizer_state_dict(model_2, optim_2))
461+
self.assertEqual(optim_2.param_groups[0]["lr"], 0.1)
462+
463+
420464
instantiate_parametrized_tests(TestE2ESaveAndLoad)
421465
if __name__ == "__main__":
422466
run_tests()

‎torch/distributed/checkpoint/planner_helpers.py‎

Lines changed: 2 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@
1111
from torch.distributed._tensor._utils import compute_local_shape_and_global_offset
1212
from torch.distributed.checkpoint.planner import _Checkpointable
1313

14-
from torch.utils._pytree import tree_map_only
14+
from torch.utils._pytree import tree_map_only_
1515

1616
from .metadata import (
1717
BytesStorageMetadata,
@@ -295,13 +295,7 @@ def _create_read_items(fqn: str, md: STORAGE_TYPES, obj: Any) -> List[ReadItem]:
295295

296296

297297
def _init_state_dict(state_dict: STATE_DICT_TYPE) -> None:
298-
state_dict_assigned_storage = tree_map_only(
299-
torch.Tensor, lambda v: _init_meta_tensor(v), state_dict
300-
)
301-
# The inplace version of tree_map_only, tree_map_only_ doesn't seem to work.
302-
# So we need to temporariy update the each element in the state dict with meta tensor.
303-
for k in state_dict.keys():
304-
state_dict[k] = state_dict_assigned_storage[k]
298+
tree_map_only_(torch.Tensor, _init_meta_tensor, state_dict)
305299

306300

307301
def _init_meta_tensor(value: Any) -> Any:

0 commit comments

Comments
 (0)