diff --git a/tests/test_model_utils.py b/tests/test_model_utils.py index 51354135af1..ce492a7ba12 100644 --- a/tests/test_model_utils.py +++ b/tests/test_model_utils.py @@ -19,7 +19,7 @@ from transformers import AutoModelForCausalLM from trl.import_utils import is_deepspeed_available -from trl.models.utils import disable_gradient_checkpointing, prepare_deepspeed +from trl.models.utils import _unwrap_model_for_generation, disable_gradient_checkpointing, prepare_deepspeed @pytest.mark.skipif(not is_deepspeed_available(), reason="deepspeed is not installed") @@ -73,3 +73,35 @@ def test_when_enabled(self): with disable_gradient_checkpointing(model): assert model.is_gradient_checkpointing is False assert model.is_gradient_checkpointing is True + + +class TestUnwrapModelForGeneration: + def test_restores_gradient_checkpointing_on_error(self): + class DummyModel: + def __init__(self): + self.is_gradient_checkpointing = True + self.enable_calls = 0 + + def gradient_checkpointing_disable(self): + self.is_gradient_checkpointing = False + + def gradient_checkpointing_enable(self): + self.is_gradient_checkpointing = True + self.enable_calls += 1 + + class FakeDistributedBackend: + def __init__(self, accelerator): + self.is_zero3 = False + + unwrapped = DummyModel() + accelerator = types.SimpleNamespace(unwrap_model=lambda model: unwrapped) + + with ( + patch("trl.distributed.DistributedBackend", FakeDistributedBackend), + pytest.raises(RuntimeError, match="boom"), + ): + with _unwrap_model_for_generation(unwrapped, accelerator): + raise RuntimeError("boom") + + assert unwrapped.enable_calls == 1 + assert unwrapped.is_gradient_checkpointing is True diff --git a/trl/models/utils.py b/trl/models/utils.py index 2b1c510273c..07ba2b4e14c 100644 --- a/trl/models/utils.py +++ b/trl/models/utils.py @@ -130,20 +130,24 @@ def _unwrap_model_for_generation( unwrapped_model.gradient_checkpointing_disable() from ..distributed import DistributedBackend - if DistributedBackend(accelerator).is_zero3: - if not gather_deepspeed3_params: - yield accelerator.unwrap_model(model) - else: - import deepspeed - - with deepspeed.zero.GatheredParameters(model.parameters()): - remove_hooks(model) + try: + if DistributedBackend(accelerator).is_zero3: + if not gather_deepspeed3_params: yield accelerator.unwrap_model(model) - add_hooks(model) - else: - yield unwrapped_model - if is_gradient_checkpointing: - unwrapped_model.gradient_checkpointing_enable() + else: + import deepspeed + + with deepspeed.zero.GatheredParameters(model.parameters()): + remove_hooks(model) + try: + yield accelerator.unwrap_model(model) + finally: + add_hooks(model) + else: + yield unwrapped_model + finally: + if is_gradient_checkpointing: + unwrapped_model.gradient_checkpointing_enable() @contextmanager