Skip to content
Merged
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
4 changes: 3 additions & 1 deletion simpletuner/helpers/training/trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -103,7 +103,7 @@
from simpletuner.helpers.training.sdnq_compile import configure_sdnq_compile_mode
from simpletuner.helpers.training.state_tracker import StateTracker
from simpletuner.helpers.training.validation import Validation, prepare_validation_prompt_list
from simpletuner.helpers.training.wrappers import unwrap_model
from simpletuner.helpers.training.wrappers import rebind_prepared_forward, unwrap_model
from simpletuner.helpers.utils import ramtorch as ramtorch_utils
from simpletuner.helpers.utils.checkpoint_manager import (
CHECKPOINT_GUARD_FILENAME,
Expand Down Expand Up @@ -4574,6 +4574,7 @@ def init_prepare_models(self, lr_scheduler):
accelerator_state.parallelism_config = standalone_cp_prepare_config
for label, prepared in zip(prepared_labels, results):
if label == "primary_model":
prepared = rebind_prepared_forward(prepared, primary_model)
self.model.set_prepared_model(prepared)
# If we skipped device placement for block swap, move the model incrementally
if musubi_block_swap_active:
Expand Down Expand Up @@ -4636,6 +4637,7 @@ def prepare_model_for_inference(self) -> None:
self.model.before_accelerator_prepare()

prepared_model = self.accelerator.prepare_model(primary_model, evaluation_mode=True)
prepared_model = rebind_prepared_forward(prepared_model, primary_model)
self.model.set_prepared_model(prepared_model)
prepared_model.eval()

Expand Down
16 changes: 16 additions & 0 deletions simpletuner/helpers/training/wrappers.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
from types import MethodType

from diffusers.utils.torch_utils import is_compiled_module


Expand Down Expand Up @@ -28,6 +30,20 @@ def unwrap_model(accelerator, model, keep_fp32_wrapper: bool = True):
return _unwrap_execution_wrappers(model)


def rebind_prepared_forward(prepared, original):
if original is None or prepared is original:
return prepared
target = prepared
while target is not None and target is not original:
for attribute in ("forward", "_original_forward"):
bound = target.__dict__.get(attribute)
function = getattr(bound, "__func__", None)
if function is not None and getattr(bound, "__self__", None) is original:
setattr(target, attribute, MethodType(function, target))
target = target.__dict__.get("_modules", {}).get("module")
return prepared


def gather_dict_of_tensors_shapes(tensors: dict) -> dict:
def _shape(value):
if isinstance(value, dict):
Expand Down
227 changes: 227 additions & 0 deletions tests/test_wrappers_rebind.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,227 @@
import unittest
from types import MethodType

import torch
import torch._dynamo
from torch import nn

from accelerate.utils.operations import convert_outputs_to_fp32
from accelerate.utils.other import compile_regions

from simpletuner.helpers.training.wrappers import rebind_prepared_forward

OptimizedModule = torch._dynamo.eval_frame.OptimizedModule


class _Block(nn.Module):
def __init__(self):
super().__init__()
self.linear = nn.Linear(4, 4, bias=False)
with torch.no_grad():
self.linear.weight.copy_(torch.eye(4))

def forward(self, x):
return self.linear(x)


class _Tiny(nn.Module):
def __init__(self):
super().__init__()
self.blocks = nn.ModuleList([_Block(), _Block()])
self.gradient_checkpointing = False
self.executed_by = []

def forward(self, x):
self.executed_by.append(self)
if self.gradient_checkpointing:
x = x * 2
for block in self.blocks:
x = block(x)
return x


def _install_mixed_precision_forward(model):
autocast_context = torch.autocast(device_type="cpu", dtype=torch.bfloat16)
model._original_forward = model.forward
model_forward_func = model.forward.__func__
new_forward = autocast_context(model_forward_func)
model.forward = MethodType(new_forward, model)
model.forward = MethodType(convert_outputs_to_fp32(model.forward.__func__), model)
return model


def _prepare_like_accelerate(backend="eager"):
original = _Tiny()
_install_mixed_precision_forward(original)
twin = compile_regions(original, backend=backend)
return original, twin


class RebindPreparedForwardTests(unittest.TestCase):
def test_compile_regions_twin_inherits_forward_bound_to_the_original(self):
original, twin = _prepare_like_accelerate()

self.assertIsNot(twin, original)
self.assertIsInstance(twin.blocks[0], OptimizedModule)
self.assertNotIsInstance(original.blocks[0], OptimizedModule)
self.assertIs(twin.__dict__["forward"].__self__, original)
self.assertIs(twin.__dict__["_original_forward"].__self__, original)

twin.gradient_checkpointing = True
x = torch.ones(1, 4)
out = twin(x)

self.assertIs(twin.executed_by[-1], original)
self.assertFalse(original.gradient_checkpointing)
torch.testing.assert_close(out, x)

def test_rebind_makes_the_twin_the_executing_module(self):
original, twin = _prepare_like_accelerate()

returned = rebind_prepared_forward(twin, original)

self.assertIs(returned, twin)
self.assertIs(twin.__dict__["forward"].__self__, twin)
self.assertIs(twin.__dict__["_original_forward"].__self__, twin)
self.assertIs(original.__dict__["forward"].__self__, original)

twin.gradient_checkpointing = True
x = torch.ones(1, 4)
out = twin(x)

self.assertIs(twin.executed_by[-1], twin)
self.assertIsInstance(twin.executed_by[-1].blocks[0], OptimizedModule)
torch.testing.assert_close(out, x * 2)

def test_rebind_applies_with_the_inductor_backend(self):
original, twin = _prepare_like_accelerate(backend="inductor")

self.assertIs(twin.__dict__["forward"].__self__, original)
rebind_prepared_forward(twin, original)
self.assertIs(twin.__dict__["forward"].__self__, twin)

def test_noop_when_prepared_is_original(self):
original = _install_mixed_precision_forward(_Tiny())
before = original.__dict__["forward"]

self.assertIs(rebind_prepared_forward(original, original), original)
self.assertIs(original.__dict__["forward"], before)
self.assertIs(original.__dict__["forward"].__self__, original)

def test_noop_when_original_is_none(self):
original = _install_mixed_precision_forward(_Tiny())
twin = compile_regions(original, backend="eager")
before = twin.__dict__["forward"]

self.assertIs(rebind_prepared_forward(twin, None), twin)
self.assertIs(twin.__dict__["forward"], before)
self.assertIs(twin.__dict__["forward"].__self__, original)

def test_noop_when_accelerate_already_rebound_the_twin(self):
original, twin = _prepare_like_accelerate()
twin.forward = MethodType(twin.__dict__["forward"].__func__, twin)
twin._original_forward = MethodType(twin.__dict__["_original_forward"].__func__, twin)
before_forward = twin.__dict__["forward"]
before_original_forward = twin.__dict__["_original_forward"]

self.assertIs(rebind_prepared_forward(twin, original), twin)
self.assertIs(twin.__dict__["forward"], before_forward)
self.assertIs(twin.__dict__["_original_forward"], before_original_forward)
self.assertIs(twin.__dict__["forward"].__self__, twin)
self.assertIs(original.__dict__["forward"].__self__, original)

def test_noop_without_an_instance_forward(self):
original = _Tiny()
twin = compile_regions(original, backend="eager")

self.assertNotIn("forward", twin.__dict__)
rebind_prepared_forward(twin, original)
self.assertNotIn("forward", twin.__dict__)
self.assertNotIn("_original_forward", twin.__dict__)

def test_noop_for_plain_torch_compile_optimized_module(self):
original = _install_mixed_precision_forward(_Tiny())
compiled = torch.compile(original, backend="eager")
before = compiled.__dict__.get("forward")

self.assertIsInstance(compiled, OptimizedModule)
self.assertIs(rebind_prepared_forward(compiled, original), compiled)
self.assertIs(compiled.__dict__.get("forward"), before)
self.assertIs(original.__dict__["forward"].__self__, original)

def test_rebinds_the_inner_twin_under_ddp(self):
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel

dist.init_process_group("gloo", store=dist.HashStore(), rank=0, world_size=1)
try:
original = _install_mixed_precision_forward(_Tiny())
twin = compile_regions(DistributedDataParallel(original), backend="eager")

self.assertIsInstance(twin, DistributedDataParallel)
self.assertIsNot(twin.module, original)
self.assertIs(twin.module.__dict__["forward"].__self__, original)

self.assertIs(rebind_prepared_forward(twin, original), twin)
self.assertIs(twin.module.__dict__["forward"].__self__, twin.module)
self.assertIs(twin.module.__dict__["_original_forward"].__self__, twin.module)
self.assertIs(original.__dict__["forward"].__self__, original)

twin.module.gradient_checkpointing = True
x = torch.ones(1, 4)
out = twin(x)

self.assertIs(twin.module.executed_by[-1], twin.module)
self.assertIsInstance(twin.module.executed_by[-1].blocks[0], OptimizedModule)
torch.testing.assert_close(out, x * 2)
finally:
dist.destroy_process_group()

def test_noop_when_accelerate_already_rebound_the_inner_twin_under_ddp(self):
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel

dist.init_process_group("gloo", store=dist.HashStore(), rank=0, world_size=1)
try:
original = _install_mixed_precision_forward(_Tiny())
twin = compile_regions(DistributedDataParallel(original), backend="eager")
inner = twin.module
inner.forward = MethodType(inner.__dict__["forward"].__func__, inner)
inner._original_forward = MethodType(inner.__dict__["_original_forward"].__func__, inner)
before_forward = inner.__dict__["forward"]
before_original_forward = inner.__dict__["_original_forward"]

self.assertIs(rebind_prepared_forward(twin, original), twin)
self.assertIs(inner.__dict__["forward"], before_forward)
self.assertIs(inner.__dict__["_original_forward"], before_original_forward)
self.assertNotIn("forward", twin.__dict__)
self.assertIs(original.__dict__["forward"].__self__, original)

inner.gradient_checkpointing = True
x = torch.ones(1, 4)
torch.testing.assert_close(twin(x), x * 2)
self.assertIs(inner.executed_by[-1], inner)
finally:
dist.destroy_process_group()

def test_noop_for_a_module_wrapper_with_dot_module(self):
original = _install_mixed_precision_forward(_Tiny())

class _Wrapper(nn.Module):
def __init__(self, module):
super().__init__()
self.module = module

def forward(self, *args, **kwargs):
return self.module(*args, **kwargs)

wrapper = _Wrapper(original)

self.assertNotIn("forward", wrapper.__dict__)
self.assertIs(rebind_prepared_forward(wrapper, original), wrapper)
self.assertNotIn("forward", wrapper.__dict__)
self.assertIs(original.__dict__["forward"].__self__, original)


if __name__ == "__main__":
unittest.main()
Loading