Skip to content

Commit 9725d31

Browse files
authored
remove force gradient checkpointing (#1589)
1 parent ea0bf33 commit 9725d31

4 files changed

Lines changed: 4 additions & 24 deletions

File tree

‎examples/lingbot_video/model_training/train.py‎

Lines changed: 1 addition & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
import torch, os, argparse, accelerate, warnings
1+
import torch, os, argparse, accelerate
22
from diffsynth.core import UnifiedDataset
33
from diffsynth.pipelines.lingbot_video import LingBotVideoPipeline
44
from diffsynth.diffusion import *
@@ -26,11 +26,6 @@ def __init__(
2626
min_timestep_boundary=0.0,
2727
):
2828
super().__init__()
29-
# Warning
30-
if not use_gradient_checkpointing:
31-
warnings.warn("Gradient checkpointing is detected as disabled. To prevent out-of-memory errors, the training framework will forcibly enable gradient checkpointing.")
32-
use_gradient_checkpointing = True
33-
3429
# Load models. The Qwen3-VL processor (tokenizer + image/video processor) is
3530
# passed separately via `processor_config`, mirroring the inference pipeline.
3631
model_configs = self.parse_model_configs(model_paths, model_id_with_origin_paths, fp8_models=fp8_models, offload_models=offload_models, device=device)

‎examples/ltx2/model_training/train.py‎

Lines changed: 1 addition & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
import torch, os, argparse, accelerate, warnings
1+
import torch, os, argparse, accelerate
22
from diffsynth.core import UnifiedDataset
33
from diffsynth.core.data.operators import LoadAudioWithTorchaudio, ToAbsolutePath, RouteByType, SequencialProcess
44
from diffsynth.pipelines.ltx2_audio_video import LTX2AudioVideoPipeline, ModelConfig
@@ -24,11 +24,6 @@ def __init__(
2424
task="sft",
2525
):
2626
super().__init__()
27-
# Warning
28-
if not use_gradient_checkpointing:
29-
warnings.warn("Gradient checkpointing is detected as disabled. To prevent out-of-memory errors, the training framework will forcibly enable gradient checkpointing.")
30-
use_gradient_checkpointing = True
31-
3227
# Load models
3328
model_configs = self.parse_model_configs(model_paths, model_id_with_origin_paths, fp8_models=fp8_models, offload_models=offload_models, device=device)
3429
tokenizer_config = ModelConfig(model_id="google/gemma-3-12b-it-qat-q4_0-unquantized") if tokenizer_path is None else ModelConfig(tokenizer_path)

‎examples/minimax_h3/model_training/train.py‎

Lines changed: 1 addition & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
import torch, os, argparse, accelerate, warnings
1+
import torch, os, argparse, accelerate
22
from diffsynth.core import UnifiedDataset
33
from diffsynth.core.data.operators import LoadAudioWithTorchaudio, ToAbsolutePath
44
from diffsynth.utils.data.minimax_h3 import MiniMaxH3ReferenceLoader
@@ -31,11 +31,6 @@ def __init__(
3131
task="sft",
3232
):
3333
super().__init__()
34-
# Warning
35-
if not use_gradient_checkpointing:
36-
warnings.warn("Gradient checkpointing is detected as disabled. To prevent out-of-memory errors, the training framework will forcibly enable gradient checkpointing.")
37-
use_gradient_checkpointing = True
38-
3934
# Load models
4035
model_configs = self.parse_model_configs(model_paths, model_id_with_origin_paths, fp8_models=fp8_models, offload_models=offload_models, device=device)
4136
pipe_kwargs = {}

‎examples/mova/model_training/train.py‎

Lines changed: 1 addition & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
import torch, os, argparse, accelerate, warnings
1+
import torch, os, argparse, accelerate
22
from diffsynth.core import UnifiedDataset
33
from diffsynth.core.data.operators import LoadAudioWithTorchaudio, ToAbsolutePath, RouteByType, SequencialProcess
44
from diffsynth.pipelines.mova_audio_video import MovaAudioVideoPipeline, ModelConfig
@@ -26,11 +26,6 @@ def __init__(
2626
min_timestep_boundary=0.0,
2727
):
2828
super().__init__()
29-
# Warning
30-
if not use_gradient_checkpointing:
31-
warnings.warn("Gradient checkpointing is detected as disabled. To prevent out-of-memory errors, the training framework will forcibly enable gradient checkpointing.")
32-
use_gradient_checkpointing = True
33-
3429
# Load models
3530
model_configs = self.parse_model_configs(model_paths, model_id_with_origin_paths, fp8_models=fp8_models, offload_models=offload_models, device=device)
3631
tokenizer_config = ModelConfig(model_id="google/gemma-3-12b-it-qat-q4_0-unquantized") if tokenizer_path is None else ModelConfig(tokenizer_path)

0 commit comments

Comments
 (0)