diff --git a/src/accelerate/optimizer.py b/src/accelerate/optimizer.py index c1f8faa1543..a318f2f1d1d 100644 --- a/src/accelerate/optimizer.py +++ b/src/accelerate/optimizer.py @@ -141,6 +141,13 @@ def eval(self): """ if hasattr(self.optimizer, "eval") and callable(self.optimizer.eval): self.optimizer.eval() + elif ( + hasattr(self.optimizer, "optimizer") + and hasattr(self.optimizer.optimizer, "eval") + and callable(self.optimizer.optimizer.eval) + ): + # the deepspeed optimizer further wraps the optimizer + self.optimizer.optimizer.eval() def step(self, closure=None): if is_lomo_available(): diff --git a/tests/test_optimizer.py b/tests/test_optimizer.py index 8bb324f0ec5..01665563d0b 100644 --- a/tests/test_optimizer.py +++ b/tests/test_optimizer.py @@ -17,6 +17,7 @@ import torch from accelerate import Accelerator +from accelerate.optimizer import AcceleratedOptimizer from accelerate.test_utils import require_cpu, require_fp16, require_non_cpu from accelerate.test_utils.testing import AccelerateTestCase @@ -34,6 +35,36 @@ def test_accelerated_optimizer_pickling(self): self.fail(f"Accelerated optimizer pickling failed with {e}") +class OptimizerModeTester(AccelerateTestCase): + def test_accelerated_optimizer_eval_reaches_deepspeed_wrapper(self): + class ScheduleFreeLikeOptimizer(torch.optim.SGD): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.mode = None + + def eval(self): + self.mode = "eval" + + class DeepSpeedLikeOptimizerWrapper: + def __init__(self, optimizer): + self.optimizer = optimizer + + def state_dict(self): + return self.optimizer.state_dict() + + def load_state_dict(self, state_dict): + self.optimizer.load_state_dict(state_dict) + + Accelerator() + model = torch.nn.Linear(10, 10) + inner = ScheduleFreeLikeOptimizer(model.parameters(), lr=0.1) + optimizer = AcceleratedOptimizer(DeepSpeedLikeOptimizerWrapper(inner), device_placement=False) + + optimizer.eval() + + assert inner.mode == "eval" + + @require_fp16 @require_non_cpu class OptimizerTester(AccelerateTestCase):