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
34 changes: 33 additions & 1 deletion tests/test_model_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -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
30 changes: 17 additions & 13 deletions trl/models/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down