From d15a88629cf973b760f0aa8d5efa650faf67eab2 Mon Sep 17 00:00:00 2001 From: bghira Date: Sun, 2 Aug 2026 13:25:50 -0600 Subject: [PATCH] Add Ideogram segmented checkpointing support --- .../helpers/models/ideogram/transformer.py | 23 ++++++++++++++++--- .../training/default_settings/safety_check.py | 2 ++ ...t_segmented_checkpointing_model_support.py | 16 +++++++++++++ 3 files changed, 38 insertions(+), 3 deletions(-) diff --git a/simpletuner/helpers/models/ideogram/transformer.py b/simpletuner/helpers/models/ideogram/transformer.py index de236f427..c30675f41 100644 --- a/simpletuner/helpers/models/ideogram/transformer.py +++ b/simpletuner/helpers/models/ideogram/transformer.py @@ -29,6 +29,7 @@ QWEN3_VL_ACTIVATION_LAYERS, ) from simpletuner.helpers.models.ideogram.quantized_loading import Fp8Linear +from simpletuner.helpers.training.gradient_checkpointing_interval import should_checkpoint_block @dataclass @@ -336,6 +337,8 @@ def __init__(self, config: Ideogram4Config) -> None: ) self.gradient_checkpointing = False self.gradient_checkpointing_backend = "torch" + self.gradient_checkpointing_interval = None + self.gradient_checkpointing_segment_stride = None def enable_gradient_checkpointing(self) -> None: self.gradient_checkpointing = True @@ -346,6 +349,12 @@ def disable_gradient_checkpointing(self) -> None: def set_gradient_checkpointing_backend(self, backend: str) -> None: self.gradient_checkpointing_backend = backend + def set_gradient_checkpointing_interval(self, interval: int) -> None: + self.gradient_checkpointing_interval = interval + + def set_gradient_checkpointing_segment_stride(self, segment_stride: int | None) -> None: + self.gradient_checkpointing_segment_stride = segment_stride + def enable_flowmap_time_conditioning(self, gate_value: float = 0.25, deltatime_type: str = "r") -> None: self.flowmap_deltatime_type = validate_flowmap_deltatime_type(deltatime_type, model_name="Ideogram") if self.delta_t_embedding is None: @@ -461,15 +470,23 @@ def forward( sin = sin.to(h.dtype) if torch.is_grad_enabled() and self.gradient_checkpointing: - 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 else: checkpoint_fn = torch.utils.checkpoint.checkpoint - for layer in self.layers: - h = checkpoint_fn(layer, h, segment_ids, cos, sin, adaln_input, use_reentrant=False) + for layer_idx, layer in enumerate(self.layers): + if should_checkpoint_block( + layer_idx, + True, + self.gradient_checkpointing_interval, + self.gradient_checkpointing_segment_stride, + ): + h = checkpoint_fn(layer, h, segment_ids, cos, sin, adaln_input, use_reentrant=False) + else: + h = layer(h, segment_ids=segment_ids, cos=cos, sin=sin, adaln_input=adaln_input) else: for layer in self.layers: h = layer(h, segment_ids=segment_ids, cos=cos, sin=sin, adaln_input=adaln_input) diff --git a/simpletuner/helpers/training/default_settings/safety_check.py b/simpletuner/helpers/training/default_settings/safety_check.py index db4c2bde5..fc879e5a4 100644 --- a/simpletuner/helpers/training/default_settings/safety_check.py +++ b/simpletuner/helpers/training/default_settings/safety_check.py @@ -155,6 +155,7 @@ def safety_check(args, accelerator): "cosmos", "flux2", "hidream", + "ideogram", ] gradient_checkpointing_segment_stride_supported_models = [ "ace_step", @@ -168,6 +169,7 @@ def safety_check(args, accelerator): "flux2", "hidream", "hunyuanvideo", + "ideogram", ] 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 4dd9c9944..bf5658ba9 100644 --- a/tests/test_segmented_checkpointing_model_support.py +++ b/tests/test_segmented_checkpointing_model_support.py @@ -340,3 +340,19 @@ def test_checkpointing_controls(self): ffn=False, attention_offload=True, ) + + +class IdeogramSegmentedCheckpointingSupportTests(unittest.TestCase): + def test_checkpointing_controls(self): + from simpletuner.helpers.models.ideogram.transformer import Ideogram4Transformer + + assert_checkpointing_controls( + self, + Ideogram4Transformer, + backend=True, + interval=True, + stride=True, + checkpoint_attention_offload=False, + ffn=False, + attention_offload=False, + )