From f141c4e07331a1e06d6ec28e275acb3746e551c2 Mon Sep 17 00:00:00 2001 From: bghira Date: Sun, 2 Aug 2026 13:27:23 -0600 Subject: [PATCH] Add Sana segmented checkpointing support --- simpletuner/helpers/models/sana/transformer.py | 14 ++++++++++++-- .../training/default_settings/safety_check.py | 1 + ...test_segmented_checkpointing_model_support.py | 16 ++++++++++++++++ 3 files changed, 29 insertions(+), 2 deletions(-) diff --git a/simpletuner/helpers/models/sana/transformer.py b/simpletuner/helpers/models/sana/transformer.py index cc93c61ae..b21370ba7 100644 --- a/simpletuner/helpers/models/sana/transformer.py +++ b/simpletuner/helpers/models/sana/transformer.py @@ -35,6 +35,7 @@ validate_flowmap_deltatime_type, ) from simpletuner.helpers.musubi_block_swap import MusubiBlockSwapManager +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block from simpletuner.helpers.training.grounding.gligen_layers import apply_grounding_fuser from simpletuner.helpers.training.tread import TREADRouter from simpletuner.helpers.utils.patching import CallableDict, MutableModuleList, PatchableModule @@ -472,6 +473,7 @@ def __init__( self.gradient_checkpointing = False self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None self.gradient_checkpointing_backend = "torch" # tread support @@ -500,6 +502,9 @@ def set_gradient_checkpointing_interval(self, interval: int): """ self.gradient_checkpointing_interval = interval + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None): + self.gradient_checkpointing_segment_stride = segment_stride + def set_gradient_checkpointing_backend(self, backend: str): self.gradient_checkpointing_backend = backend @@ -742,7 +747,12 @@ def forward( if ( self.training and self.gradient_checkpointing - and (self.gradient_checkpointing_interval is None or i % self.gradient_checkpointing_interval == 0) + and should_checkpoint_block( + i, + True, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ) ): def create_custom_forward(module): @@ -751,7 +761,7 @@ def custom_forward(*inputs): return custom_forward - if self.gradient_checkpointing_backend == "unsloth": + if self.gradient_checkpointing_backend.startswith("unsloth"): from simpletuner.helpers.training.offloaded_gradient_checkpointer import offloaded_checkpoint checkpoint_fn = offloaded_checkpoint diff --git a/simpletuner/helpers/training/default_settings/safety_check.py b/simpletuner/helpers/training/default_settings/safety_check.py index 9ef8d207c..7f448ed31 100644 --- a/simpletuner/helpers/training/default_settings/safety_check.py +++ b/simpletuner/helpers/training/default_settings/safety_check.py @@ -191,6 +191,7 @@ def safety_check(args, accelerator): "mageflow", "pixart", "qwen_image", + "sana", ] attention_activation_offload_supported_models = [ "chroma", diff --git a/tests/test_segmented_checkpointing_model_support.py b/tests/test_segmented_checkpointing_model_support.py index bd81f1f53..1b3b8186d 100644 --- a/tests/test_segmented_checkpointing_model_support.py +++ b/tests/test_segmented_checkpointing_model_support.py @@ -527,3 +527,19 @@ def test_checkpointing_controls(self): ffn=False, attention_offload=False, ) + + +class SanaSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.sana.transformer import SanaTransformer2DModel + + assert_checkpointing_controls( + self, + SanaTransformer2DModel, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + )