diff --git a/modules/modelSetup/BaseAnimaSetup.py b/modules/modelSetup/BaseAnimaSetup.py index 8c0b9e0aa..7af31f09f 100644 --- a/modules/modelSetup/BaseAnimaSetup.py +++ b/modules/modelSetup/BaseAnimaSetup.py @@ -14,8 +14,6 @@ enable_checkpointing_for_qwen_transformer, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -23,7 +21,6 @@ from torch import Tensor -#TODO share more code with other models class BaseAnimaSetup( BaseModelSetup, ModelSetupDiffusionLossMixin, @@ -46,23 +43,10 @@ def setup_optimizations( model: AnimaModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_qwen_transformer(model.transformer, config, config.transformer) - model.text_encoder_offload_conductor = enable_checkpointing_for_qwen3_encoder_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_qwen_transformer) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_qwen3_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) self._set_attention_backend(model.transformer, config.attention_mechanism, mask=False) diff --git a/modules/modelSetup/BaseChromaSetup.py b/modules/modelSetup/BaseChromaSetup.py index 0083774a1..d4e882e44 100644 --- a/modules/modelSetup/BaseChromaSetup.py +++ b/modules/modelSetup/BaseChromaSetup.py @@ -16,15 +16,12 @@ enable_checkpointing_for_t5_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers from modules.util.TrainProgress import TrainProgress import torch from torch import Tensor -#TODO share more code with Flux and other models class BaseChromaSetup( BaseModelSetup, ModelSetupDiffusionLossMixin, @@ -47,23 +44,10 @@ def setup_optimizations( model: ChromaModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_chroma_transformer(model.transformer, config, config.transformer) - model.text_encoder_offload_conductor = enable_checkpointing_for_t5_encoder_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_chroma_transformer) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_t5_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) diff --git a/modules/modelSetup/BaseErnieSetup.py b/modules/modelSetup/BaseErnieSetup.py index d5050c49f..91d929827 100644 --- a/modules/modelSetup/BaseErnieSetup.py +++ b/modules/modelSetup/BaseErnieSetup.py @@ -14,8 +14,6 @@ enable_checkpointing_for_mistral_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers from modules.util.TrainProgress import TrainProgress import torch @@ -43,23 +41,10 @@ def setup_optimizations( model: ErnieModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_ernie_transformer(model.transformer, config, config.transformer) - model.text_encoder_offload_conductor = enable_checkpointing_for_mistral_encoder_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_ernie_transformer) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_mistral_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) diff --git a/modules/modelSetup/BaseFlux2Setup.py b/modules/modelSetup/BaseFlux2Setup.py index 63e8dfc70..1aa032946 100644 --- a/modules/modelSetup/BaseFlux2Setup.py +++ b/modules/modelSetup/BaseFlux2Setup.py @@ -16,8 +16,6 @@ enable_checkpointing_for_qwen3_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers from modules.util.TrainProgress import TrainProgress import torch @@ -43,26 +41,12 @@ def setup_optimizations( model: Flux2Model, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_flux2_transformer(model.transformer, config, config.transformer) - if model.is_dev(): - model.text_encoder_offload_conductor = enable_checkpointing_for_mistral_encoder_layers(model.text_encoder, config, config.text_encoder) - else: - model.text_encoder_offload_conductor = enable_checkpointing_for_qwen3_encoder_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) + text_encoder_checkpointing_fn = enable_checkpointing_for_mistral_encoder_layers if model.is_dev() \ + else enable_checkpointing_for_qwen3_encoder_layers + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_flux2_transformer) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, text_encoder_checkpointing_fn, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) self._set_attention_backend(model.transformer, config.attention_mechanism, mask=False) diff --git a/modules/modelSetup/BaseFluxSetup.py b/modules/modelSetup/BaseFluxSetup.py index 398d4a1b4..1d5fddb07 100644 --- a/modules/modelSetup/BaseFluxSetup.py +++ b/modules/modelSetup/BaseFluxSetup.py @@ -17,8 +17,6 @@ enable_checkpointing_for_t5_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -48,27 +46,11 @@ def setup_optimizations( model: FluxModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_flux_transformer(model.transformer, config, config.transformer) - if model.text_encoder_1 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_1, config, config.text_encoder) - if model.text_encoder_2 is not None: - model.text_encoder_2_offload_conductor = enable_checkpointing_for_t5_encoder_layers(model.text_encoder_2, config, config.text_encoder_2) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_2_autocast_context, model.text_encoder_2_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder_1, self.train_device, model.train_dtype, config) - quantize_layers(model.text_encoder_2, self.train_device, model.text_encoder_2_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_flux_transformer) + self._setup_model_part(model, config, "text_encoder_1", config.text_encoder, enable_checkpointing_for_clip_encoder_layers) + self._setup_model_part(model, config, "text_encoder_2", config.text_encoder_2, enable_checkpointing_for_t5_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) self._set_attention_backend(model.transformer, config.attention_mechanism, mask=False) diff --git a/modules/modelSetup/BaseHiDreamSetup.py b/modules/modelSetup/BaseHiDreamSetup.py index 301574535..2bd864b4b 100644 --- a/modules/modelSetup/BaseHiDreamSetup.py +++ b/modules/modelSetup/BaseHiDreamSetup.py @@ -17,8 +17,6 @@ enable_checkpointing_for_t5_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers from modules.util.TrainProgress import TrainProgress import torch @@ -47,41 +45,13 @@ def setup_optimizations( model: HiDreamModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_hi_dream_transformer(model.transformer, config, config.transformer) - if model.text_encoder_1 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_1, config, config.text_encoder) - if model.text_encoder_2 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_2, config, config.text_encoder_2) - if model.text_encoder_3 is not None: - model.text_encoder_3_offload_conductor = enable_checkpointing_for_t5_encoder_layers(model.text_encoder_3, config, config.text_encoder_3) - if model.text_encoder_4 is not None: - model.text_encoder_4_offload_conductor = enable_checkpointing_for_llama_encoder_layers(model.text_encoder_4, config, config.text_encoder_4) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_3_autocast_context, model.text_encoder_3_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - model.transformer_autocast_context, model.transformer_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder_1, self.train_device, model.train_dtype, config) - quantize_layers(model.text_encoder_2, self.train_device, model.train_dtype, config) - quantize_layers(model.text_encoder_3, self.train_device, model.text_encoder_3_train_dtype, config) - quantize_layers(model.text_encoder_4, self.train_device, model.train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.transformer_train_dtype, config) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_hi_dream_transformer, disable_fp16_autocast=True) + self._setup_model_part(model, config, "text_encoder_1", config.text_encoder, enable_checkpointing_for_clip_encoder_layers) + self._setup_model_part(model, config, "text_encoder_2", config.text_encoder_2, enable_checkpointing_for_clip_encoder_layers) + self._setup_model_part(model, config, "text_encoder_3", config.text_encoder_3, enable_checkpointing_for_t5_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "text_encoder_4", config.text_encoder_4, enable_checkpointing_for_llama_encoder_layers) + self._setup_model_part(model, config, "vae", config.vae) self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) diff --git a/modules/modelSetup/BaseHunyuanVideoSetup.py b/modules/modelSetup/BaseHunyuanVideoSetup.py index 4f5713008..ceca94e0c 100644 --- a/modules/modelSetup/BaseHunyuanVideoSetup.py +++ b/modules/modelSetup/BaseHunyuanVideoSetup.py @@ -17,8 +17,6 @@ enable_checkpointing_for_llama_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers from modules.util.TrainProgress import TrainProgress import torch @@ -47,27 +45,11 @@ def setup_optimizations( model: HunyuanVideoModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_hunyuan_video_transformer(model.transformer, config, config.transformer) - if model.text_encoder_1 is not None: - model.text_encoder_1_offload_conductor = enable_checkpointing_for_llama_encoder_layers(model.text_encoder_1, config, config.text_encoder) - if model.text_encoder_2 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_2, config, config.text_encoder_2) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.transformer_autocast_context, model.transformer_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder_1, self.train_device, model.train_dtype, config) - quantize_layers(model.text_encoder_2, self.train_device, model.train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.transformer_train_dtype, config) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_hunyuan_video_transformer, disable_fp16_autocast=True) + self._setup_model_part(model, config, "text_encoder_1", config.text_encoder, enable_checkpointing_for_llama_encoder_layers) + self._setup_model_part(model, config, "text_encoder_2", config.text_encoder_2, enable_checkpointing_for_clip_encoder_layers) + self._setup_model_part(model, config, "vae", config.vae) model.vae.enable_tiling() self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) diff --git a/modules/modelSetup/BaseIdeogramSetup.py b/modules/modelSetup/BaseIdeogramSetup.py index 7fb1a0bff..6f23b5ef3 100644 --- a/modules/modelSetup/BaseIdeogramSetup.py +++ b/modules/modelSetup/BaseIdeogramSetup.py @@ -14,8 +14,6 @@ enable_checkpointing_for_qwen3vl_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers from modules.util.TrainProgress import TrainProgress import torch @@ -43,33 +41,13 @@ def setup_optimizations( model: IdeogramModel, config: TrainConfig, ): - # Only the conditional transformer is trained, so gradient checkpointing applies there. - model.transformer_offload_conductor = \ - enable_checkpointing_for_ideogram_transformer(model.transformer, config, config.transformer) - - # The unconditional transformer is frozen, but it still benefits from layer offloading - # since both transformers need to fit in VRAM during sampling. It is optional, so may be unloaded. - if model.unconditional_transformer is not None: - model.unconditional_transformer_offload_conductor = \ - enable_checkpointing_for_ideogram_transformer(model.unconditional_transformer, config, config.unconditional_transformer) - - model.text_encoder_offload_conductor = enable_checkpointing_for_qwen3vl_encoder_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) - quantize_layers(model.unconditional_transformer, self.train_device, model.train_dtype, config) + super().setup_optimizations(model, config) + # The unconditional transformer is frozen but still layer-offloaded, so both transformers fit in VRAM + # during sampling. + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_ideogram_transformer) + self._setup_model_part(model, config, "unconditional_transformer", config.unconditional_transformer, enable_checkpointing_for_ideogram_transformer) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_qwen3vl_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) self._set_attention_backend(model.transformer, config.attention_mechanism, mask=False) if model.unconditional_transformer is not None: diff --git a/modules/modelSetup/BaseKrea2Setup.py b/modules/modelSetup/BaseKrea2Setup.py index 1851fae3e..fa0c9ec97 100644 --- a/modules/modelSetup/BaseKrea2Setup.py +++ b/modules/modelSetup/BaseKrea2Setup.py @@ -14,8 +14,6 @@ enable_checkpointing_for_qwen3vl_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers from modules.util.TrainProgress import TrainProgress import torch @@ -45,23 +43,10 @@ def setup_optimizations( model: Krea2Model, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_krea2_transformer(model.transformer, config, config.transformer) - model.text_encoder_offload_conductor = enable_checkpointing_for_qwen3vl_encoder_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_krea2_transformer) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_qwen3vl_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) diff --git a/modules/modelSetup/BaseModelSetup.py b/modules/modelSetup/BaseModelSetup.py index 89b337340..2b6543c49 100644 --- a/modules/modelSetup/BaseModelSetup.py +++ b/modules/modelSetup/BaseModelSetup.py @@ -3,10 +3,12 @@ from modules.model.BaseModel import BaseModel from modules.util.config.TrainConfig import TrainConfig, TrainEmbeddingConfig, TrainModelPartConfig +from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context from modules.util.enum.AttentionMechanism import AttentionMechanism from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.ModuleFilter import ModuleFilter from modules.util.NamedParameterGroup import NamedParameterGroup, NamedParameterGroupCollection +from modules.util.quantization_util import quantize_layers from modules.util.TimedActionMixin import TimedActionMixin from modules.util.TrainProgress import TrainProgress @@ -47,7 +49,9 @@ def setup_optimizations( model: BaseModel, config: TrainConfig, ): - pass + # Leaves call super() first, so model.train_dtype is set before their _setup_model_part calls read it. + model.train_dtype = config.train_dtype + model.autocast_context = create_autocast_context(self.train_device, config.train_dtype, config.enable_autocast_cache) @abstractmethod def setup_model( @@ -237,6 +241,35 @@ def _setup_model_part_requires_grad( for param in self.frozen_parameters[unique_name]: param.requires_grad_(False) + def _setup_model_part( + self, + model, + config: TrainConfig, + attr: str, + config_part: TrainModelPartConfig, + checkpointing_fn=None, + *, + disable_fp16_autocast: bool = False, + ): + module = getattr(model, attr) + if module is None: + return + + if checkpointing_fn is not None: + conductor = checkpointing_fn(module, config, config_part) + if conductor is not None: + setattr(model, f"{attr}_offload_conductor", conductor) + + if disable_fp16_autocast: + autocast_context, train_dtype = disable_fp16_autocast_context( + self.train_device, config.train_dtype, config.fallback_train_dtype, config.enable_autocast_cache) + setattr(model, f"{attr}_autocast_context", autocast_context) + setattr(model, f"{attr}_train_dtype", train_dtype) + else: + train_dtype = model.train_dtype + + quantize_layers(module, self.train_device, train_dtype, config) + @staticmethod def _set_attention_backend(component, attn: AttentionMechanism, mask: bool): match attn: diff --git a/modules/modelSetup/BasePixArtAlphaSetup.py b/modules/modelSetup/BasePixArtAlphaSetup.py index 2fccc30f6..c698969af 100644 --- a/modules/modelSetup/BasePixArtAlphaSetup.py +++ b/modules/modelSetup/BasePixArtAlphaSetup.py @@ -16,8 +16,6 @@ enable_checkpointing_for_t5_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers from modules.util.TrainProgress import TrainProgress import torch @@ -49,22 +47,11 @@ def setup_optimizations( model: PixArtAlphaModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_basic_transformer_blocks(model.transformer, config, config.transformer, offload_enabled=True) - model.text_encoder_offload_conductor = enable_checkpointing_for_t5_encoder_layers(model.text_encoder, config, config.text_encoder) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_basic_transformer_blocks) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_t5_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) def _setup_embeddings( diff --git a/modules/modelSetup/BaseQwenSetup.py b/modules/modelSetup/BaseQwenSetup.py index 0e6cd1c55..64d617e25 100644 --- a/modules/modelSetup/BaseQwenSetup.py +++ b/modules/modelSetup/BaseQwenSetup.py @@ -14,15 +14,12 @@ enable_checkpointing_for_qwen_transformer, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers from modules.util.TrainProgress import TrainProgress import torch from torch import Tensor -#TODO share more code with other models class BaseQwenSetup( BaseModelSetup, ModelSetupDiffusionLossMixin, @@ -44,23 +41,11 @@ def setup_optimizations( model: QwenModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_qwen_transformer(model.transformer, config, config.transformer) - model.text_encoder_offload_conductor = enable_checkpointing_for_qwen25vl_encoder_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_qwen_transformer) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_qwen25vl_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) def predict( diff --git a/modules/modelSetup/BaseSanaSetup.py b/modules/modelSetup/BaseSanaSetup.py index ef4dff54f..761104238 100644 --- a/modules/modelSetup/BaseSanaSetup.py +++ b/modules/modelSetup/BaseSanaSetup.py @@ -16,8 +16,7 @@ enable_checkpointing_for_sana_transformer, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers +from modules.util.dtype_util import disable_fp16_autocast_context from modules.util.TrainProgress import TrainProgress import torch @@ -50,19 +49,10 @@ def setup_optimizations( config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_sana_transformer(model.transformer, config, config.transformer) - model.text_encoder_offload_conductor = enable_checkpointing_for_gemma_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) + super().setup_optimizations(model, config) + # Sana's vae runs under its own fp16-disabled autocast in predict(). Inconsistently, vae_train_dtype + # is never read and the vae below is quantized with model.train_dtype, not this fp16-disabled dtype. model.vae_autocast_context, model.vae_train_dtype = disable_fp16_autocast_context( self.train_device, config.train_dtype, @@ -70,9 +60,10 @@ def setup_optimizations( config.enable_autocast_cache, ) - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_sana_transformer) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_gemma_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) + self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) def _setup_embeddings( diff --git a/modules/modelSetup/BaseStableDiffusion3Setup.py b/modules/modelSetup/BaseStableDiffusion3Setup.py index dd7cf0841..3689f0710 100644 --- a/modules/modelSetup/BaseStableDiffusion3Setup.py +++ b/modules/modelSetup/BaseStableDiffusion3Setup.py @@ -17,8 +17,6 @@ enable_checkpointing_for_t5_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers from modules.util.TrainProgress import TrainProgress import torch @@ -46,30 +44,13 @@ def setup_optimizations( model: StableDiffusion3Model, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_stable_diffusion_3_transformer(model.transformer, config, config.transformer) - if model.text_encoder_1 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_1, config, config.text_encoder) - if model.text_encoder_2 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_2, config, config.text_encoder_2) - if model.text_encoder_3 is not None: - model.text_encoder_3_offload_conductor = enable_checkpointing_for_t5_encoder_layers(model.text_encoder_3, config, config.text_encoder_3) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_3_autocast_context, model.text_encoder_3_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_stable_diffusion_3_transformer) + self._setup_model_part(model, config, "text_encoder_1", config.text_encoder, enable_checkpointing_for_clip_encoder_layers) + self._setup_model_part(model, config, "text_encoder_2", config.text_encoder_2, enable_checkpointing_for_clip_encoder_layers) + self._setup_model_part(model, config, "text_encoder_3", config.text_encoder_3, enable_checkpointing_for_t5_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) - quantize_layers(model.text_encoder_1, self.train_device, model.train_dtype, config) - quantize_layers(model.text_encoder_2, self.train_device, model.train_dtype, config) - quantize_layers(model.text_encoder_3, self.train_device, model.text_encoder_3_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) self._set_attention_backend(model.transformer, config.attention_mechanism, mask=False) def _setup_embeddings( diff --git a/modules/modelSetup/BaseStableDiffusionSetup.py b/modules/modelSetup/BaseStableDiffusionSetup.py index b6a3facc0..c027745cc 100644 --- a/modules/modelSetup/BaseStableDiffusionSetup.py +++ b/modules/modelSetup/BaseStableDiffusionSetup.py @@ -17,7 +17,6 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.conv_util import apply_circular_padding_to_conv2d -from modules.util.dtype_util import create_autocast_context from modules.util.quantization_util import quantize_layers from modules.util.TrainProgress import TrainProgress @@ -49,9 +48,11 @@ def setup_optimizations( model: StableDiffusionModel, config: TrainConfig, ): + super().setup_optimizations(model, config) + if config.unet.checkpointing_enabled(): model.unet.enable_gradient_checkpointing() - enable_checkpointing_for_basic_transformer_blocks(model.unet, config, config.unet, offload_enabled=False) + enable_checkpointing_for_basic_transformer_blocks(model.unet, config, config.unet, supports_offloading=False) enable_checkpointing_for_clip_encoder_layers(model.text_encoder, config, config.text_encoder) if config.force_circular_padding: @@ -60,9 +61,6 @@ def setup_optimizations( if model.unet_lora is not None: apply_circular_padding_to_conv2d(model.unet_lora) - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - quantize_layers(model.text_encoder, self.train_device, model.train_dtype, config) quantize_layers(model.vae, self.train_device, model.train_dtype, config) quantize_layers(model.unet, self.train_device, model.train_dtype, config) diff --git a/modules/modelSetup/BaseStableDiffusionXLSetup.py b/modules/modelSetup/BaseStableDiffusionXLSetup.py index 54a080579..aeecbd4c6 100644 --- a/modules/modelSetup/BaseStableDiffusionXLSetup.py +++ b/modules/modelSetup/BaseStableDiffusionXLSetup.py @@ -17,7 +17,7 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.conv_util import apply_circular_padding_to_conv2d -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context +from modules.util.dtype_util import disable_fp16_autocast_context from modules.util.quantization_util import quantize_layers from modules.util.TrainProgress import TrainProgress @@ -46,9 +46,11 @@ def setup_optimizations( model: StableDiffusionXLModel, config: TrainConfig, ): + super().setup_optimizations(model, config) + if config.unet.checkpointing_enabled(): model.unet.enable_gradient_checkpointing() - enable_checkpointing_for_basic_transformer_blocks(model.unet, config, config.unet, offload_enabled=False) + enable_checkpointing_for_basic_transformer_blocks(model.unet, config, config.unet, supports_offloading=False) enable_checkpointing_for_clip_encoder_layers(model.text_encoder_1, config, config.text_encoder) enable_checkpointing_for_clip_encoder_layers(model.text_encoder_2, config, config.text_encoder_2) @@ -58,9 +60,6 @@ def setup_optimizations( if model.unet_lora is not None: apply_circular_padding_to_conv2d(model.unet_lora) - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - model.vae_autocast_context, model.vae_train_dtype = disable_fp16_autocast_context( self.train_device, config.train_dtype, diff --git a/modules/modelSetup/BaseWuerstchenSetup.py b/modules/modelSetup/BaseWuerstchenSetup.py index dc0c48528..54453f026 100644 --- a/modules/modelSetup/BaseWuerstchenSetup.py +++ b/modules/modelSetup/BaseWuerstchenSetup.py @@ -15,7 +15,6 @@ from modules.util.config.TrainConfig import TrainConfig from modules.util.conv_util import apply_circular_padding_to_conv2d from modules.util.dtype_util import ( - create_autocast_context, disable_bf16_on_fp16_autocast_context, disable_fp16_autocast_context, ) @@ -52,6 +51,11 @@ def setup_optimizations( model: WuerstchenModel, config: TrainConfig, ): + super().setup_optimizations(model, config) + + # Hand-wired rather than via _setup_model_part: Wuerstchen's parts don't match the + # transformer/text_encoder/vae shape it assumes, and take bespoke per-part contexts + # (stable-cascade prior fp16-disable, effnet bf16-on-fp16). if config.prior.checkpointing_enabled(): model.prior_prior.enable_gradient_checkpointing() enable_checkpointing_for_clip_encoder_layers(model.prior_text_encoder, config, config.text_encoder) @@ -63,9 +67,6 @@ def setup_optimizations( if model.prior_prior_lora is not None: apply_circular_padding_to_conv2d(model.prior_prior_lora) - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - if model.model_type.is_stable_cascade(): model.prior_autocast_context, model.prior_train_dtype = disable_fp16_autocast_context( self.train_device, diff --git a/modules/modelSetup/BaseZImageSetup.py b/modules/modelSetup/BaseZImageSetup.py index b77f9bb72..c7b23cbe2 100644 --- a/modules/modelSetup/BaseZImageSetup.py +++ b/modules/modelSetup/BaseZImageSetup.py @@ -15,8 +15,6 @@ enable_checkpointing_for_z_image_transformer, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers from modules.util.TrainProgress import TrainProgress import torch @@ -45,24 +43,11 @@ def setup_optimizations( model: ZImageModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_z_image_transformer(model.transformer, config, config.transformer) - model.text_encoder_offload_conductor = enable_checkpointing_for_qwen3_encoder_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - #TODO necessary if we don't train it? - model.text_encoder_autocast_context, model.text_encoder_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_z_image_transformer) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_qwen3_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) def predict( diff --git a/modules/util/checkpointing_util.py b/modules/util/checkpointing_util.py index 1f669e9a5..2df40065e 100644 --- a/modules/util/checkpointing_util.py +++ b/modules/util/checkpointing_util.py @@ -243,13 +243,14 @@ def enable_checkpointing( part: TrainModelPartConfig, compile: bool, lists, # if there are multiple entries in this list, they must be in the exact order they are executed - otherwise offloading fails - offload_enabled: bool = True, + supports_offloading: bool = True, ) -> LayerOffloadConductor | None: if not part.checkpointing_or_offloading_enabled() and not compile: return None - # a conductor exists iff this part actually offloads (and the component supports conductor offloading) - offload = offload_enabled and part.offloading_enabled() + # a conductor exists iff this part actually offloads: the user enabled it (part.offloading_enabled()) and the + # architecture can be driven by the conductor (supports_offloading). + offload = supports_offloading and part.offloading_enabled() conductor = LayerOffloadConductor(model, config, part) if offload else None checkpointing = part.checkpointing_enabled() @@ -298,12 +299,12 @@ def enable_checkpointing_for_basic_transformer_blocks( model: nn.Module, config: TrainConfig, part: TrainModelPartConfig, - offload_enabled: bool, + supports_offloading: bool = True, ) -> LayerOffloadConductor | None: return enable_checkpointing(model, config, part, config.compile, [ (BasicTransformerBlock , []), ], - offload_enabled = offload_enabled, + supports_offloading = supports_offloading, ) def enable_checkpointing_for_clip_encoder_layers( @@ -313,7 +314,7 @@ def enable_checkpointing_for_clip_encoder_layers( ): return enable_checkpointing(model, config, part, False, [ (CLIPEncoderLayer, []), # No activation offloading for text encoders, because the output might be taken from the middle of the network - ], offload_enabled=False) # CLIP is non-offloadable; keep it plain-checkpointed so a migrated offload_fraction can't build a self-activating conductor + ], supports_offloading=False) # CLIP is non-offloadable; keep it plain-checkpointed so a migrated offload_fraction can't build a self-activating conductor def enable_checkpointing_for_t5_encoder_layers( model: nn.Module, diff --git a/modules/util/dtype_util.py b/modules/util/dtype_util.py index b6d6b2643..7dde18d45 100644 --- a/modules/util/dtype_util.py +++ b/modules/util/dtype_util.py @@ -20,7 +20,7 @@ def create_autocast_context( device: torch.device, train_dtype: DataType | None, enable_autocast_cache: bool, -) -> tuple[torch.autocast | nullcontext, DataType]: +) -> torch.autocast | nullcontext: torch_train_dtype = train_dtype.torch_dtype() if torch_train_dtype in (torch.float16, torch.bfloat16): @@ -34,13 +34,13 @@ def create_autocast_context( # MPS additionally needs macOS >= 14. print(f"Warning: Mixed precision training is untested on device type '{device.type}'.") return torch.autocast(device_type=device.type, dtype=torch_train_dtype, - cache_enabled=enable_autocast_cache), train_dtype + cache_enabled=enable_autocast_cache) elif device.type == "cuda": # float32/tfloat32 on CUDA (and ROCm, which also reports device type "cuda"): # CUDA accepts float32 as an autocast dtype and upcasts lower-precision weights # on the fly (this is undocumented but works). return torch.autocast(device_type=device.type, dtype=torch_train_dtype, - cache_enabled=enable_autocast_cache), train_dtype + cache_enabled=enable_autocast_cache) else: # float32/tfloat32 on a non-CUDA backend (cpu, mps, xpu, ...): those backends # reject fp32 autocast, so disable autocast and let the model run at its weight @@ -49,7 +49,7 @@ def create_autocast_context( print("Warning: float32 training does not upcast lower-precision weights on this device " "(only CUDA can autocast to float32); the model runs at its weight dtype. " "Set the weight data types to float32 for full precision.") - return torch.autocast(device_type=device.type, enabled=False), train_dtype + return torch.autocast(device_type=device.type, enabled=False) def disable_fp16_autocast_context(