diff --git a/modules/modelSetup/BaseAnimaSetup.py b/modules/modelSetup/BaseAnimaSetup.py index 79179cdfb..0ed9bbb19 100644 --- a/modules/modelSetup/BaseAnimaSetup.py +++ b/modules/modelSetup/BaseAnimaSetup.py @@ -15,7 +15,6 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -50,22 +49,14 @@ def setup_optimizations( 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.weight_dtypes().transformer, - config.weight_dtypes().text_encoder, - config.weight_dtypes().vae, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - ], config.enable_autocast_cache) + 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.weight_dtypes().text_encoder, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - ], config.enable_autocast_cache, ) diff --git a/modules/modelSetup/BaseChromaSetup.py b/modules/modelSetup/BaseChromaSetup.py index 006587b71..01dcfab0c 100644 --- a/modules/modelSetup/BaseChromaSetup.py +++ b/modules/modelSetup/BaseChromaSetup.py @@ -17,7 +17,6 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -52,24 +51,14 @@ def setup_optimizations( 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.weight_dtypes().transformer, - config.weight_dtypes().text_encoder, - config.weight_dtypes().vae, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache) + 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.weight_dtypes().text_encoder, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache, ) diff --git a/modules/modelSetup/BaseErnieSetup.py b/modules/modelSetup/BaseErnieSetup.py index 703d1f6db..c412073e8 100644 --- a/modules/modelSetup/BaseErnieSetup.py +++ b/modules/modelSetup/BaseErnieSetup.py @@ -15,7 +15,6 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -48,22 +47,14 @@ def setup_optimizations( 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.weight_dtypes().transformer, - config.weight_dtypes().text_encoder, - config.weight_dtypes().vae, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - ], config.enable_autocast_cache) + 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.weight_dtypes().text_encoder, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - ], config.enable_autocast_cache, ) diff --git a/modules/modelSetup/BaseFlux2Setup.py b/modules/modelSetup/BaseFlux2Setup.py index ef22c6ef6..6b2c1bb3e 100644 --- a/modules/modelSetup/BaseFlux2Setup.py +++ b/modules/modelSetup/BaseFlux2Setup.py @@ -17,7 +17,6 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -51,22 +50,14 @@ def setup_optimizations( 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.weight_dtypes().transformer, - config.weight_dtypes().text_encoder, - config.weight_dtypes().vae, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - ], config.enable_autocast_cache) + 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.weight_dtypes().text_encoder, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - ], config.enable_autocast_cache, ) diff --git a/modules/modelSetup/BaseFluxSetup.py b/modules/modelSetup/BaseFluxSetup.py index 5ec88c0c7..55999fb20 100644 --- a/modules/modelSetup/BaseFluxSetup.py +++ b/modules/modelSetup/BaseFluxSetup.py @@ -18,7 +18,6 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -55,25 +54,14 @@ def setup_optimizations( 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.weight_dtypes().transformer, - config.weight_dtypes().text_encoder, - config.weight_dtypes().text_encoder_2, - config.weight_dtypes().vae, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache) + 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.weight_dtypes().text_encoder_2, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache, ) diff --git a/modules/modelSetup/BaseHiDreamSetup.py b/modules/modelSetup/BaseHiDreamSetup.py index 687e12431..1a06f961b 100644 --- a/modules/modelSetup/BaseHiDreamSetup.py +++ b/modules/modelSetup/BaseHiDreamSetup.py @@ -18,7 +18,6 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -59,27 +58,14 @@ def setup_optimizations( 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.weight_dtypes().transformer, - config.weight_dtypes().text_encoder, - config.weight_dtypes().text_encoder_2, - config.weight_dtypes().text_encoder_3, - config.weight_dtypes().text_encoder_4, - config.weight_dtypes().vae, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache) + 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.weight_dtypes().text_encoder_3, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache, ) @@ -88,11 +74,6 @@ def setup_optimizations( self.train_device, config.train_dtype, config.fallback_train_dtype, - [ - config.weight_dtypes().transformer, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache, ) diff --git a/modules/modelSetup/BaseHunyuanVideoSetup.py b/modules/modelSetup/BaseHunyuanVideoSetup.py index 3c003779d..8e99dce84 100644 --- a/modules/modelSetup/BaseHunyuanVideoSetup.py +++ b/modules/modelSetup/BaseHunyuanVideoSetup.py @@ -18,7 +18,6 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -55,25 +54,14 @@ def setup_optimizations( 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.weight_dtypes().transformer, - config.weight_dtypes().text_encoder, - config.weight_dtypes().text_encoder_2, - config.weight_dtypes().vae, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache) + 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.weight_dtypes().transformer, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache, ) diff --git a/modules/modelSetup/BaseIdeogramSetup.py b/modules/modelSetup/BaseIdeogramSetup.py index 279c09388..f067e8c2e 100644 --- a/modules/modelSetup/BaseIdeogramSetup.py +++ b/modules/modelSetup/BaseIdeogramSetup.py @@ -15,7 +15,6 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -57,22 +56,14 @@ def setup_optimizations( 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.weight_dtypes().transformer, - config.weight_dtypes().text_encoder, - config.weight_dtypes().vae, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - ], config.enable_autocast_cache) + 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.weight_dtypes().text_encoder, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - ], config.enable_autocast_cache, ) diff --git a/modules/modelSetup/BaseKrea2Setup.py b/modules/modelSetup/BaseKrea2Setup.py index 441dc21ee..52e9b5959 100644 --- a/modules/modelSetup/BaseKrea2Setup.py +++ b/modules/modelSetup/BaseKrea2Setup.py @@ -15,7 +15,6 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -50,22 +49,14 @@ def setup_optimizations( 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.weight_dtypes().transformer, - config.weight_dtypes().text_encoder, - config.weight_dtypes().vae, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - ], config.enable_autocast_cache) + 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.weight_dtypes().text_encoder, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - ], config.enable_autocast_cache, ) diff --git a/modules/modelSetup/BasePixArtAlphaSetup.py b/modules/modelSetup/BasePixArtAlphaSetup.py index 311109842..57bf3a40b 100644 --- a/modules/modelSetup/BasePixArtAlphaSetup.py +++ b/modules/modelSetup/BasePixArtAlphaSetup.py @@ -17,7 +17,6 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -54,23 +53,13 @@ def setup_optimizations( 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) - model.autocast_context, model.train_dtype = create_autocast_context(self.train_device, config.train_dtype, [ - config.weight_dtypes().transformer, - config.weight_dtypes().text_encoder, - config.weight_dtypes().vae, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache) + 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.weight_dtypes().text_encoder, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache, ) diff --git a/modules/modelSetup/BaseQwenSetup.py b/modules/modelSetup/BaseQwenSetup.py index c229fd695..a618dc28f 100644 --- a/modules/modelSetup/BaseQwenSetup.py +++ b/modules/modelSetup/BaseQwenSetup.py @@ -15,7 +15,6 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -49,22 +48,14 @@ def setup_optimizations( 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.weight_dtypes().transformer, - config.weight_dtypes().text_encoder, - config.weight_dtypes().vae, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - ], config.enable_autocast_cache) + 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.weight_dtypes().text_encoder, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - ], config.enable_autocast_cache, ) diff --git a/modules/modelSetup/BaseSanaSetup.py b/modules/modelSetup/BaseSanaSetup.py index 20d9eaed5..a96afd770 100644 --- a/modules/modelSetup/BaseSanaSetup.py +++ b/modules/modelSetup/BaseSanaSetup.py @@ -17,7 +17,6 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -55,23 +54,13 @@ def setup_optimizations( 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.weight_dtypes().transformer, - config.weight_dtypes().text_encoder, - config.weight_dtypes().vae, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache) + 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.weight_dtypes().text_encoder, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache, ) @@ -79,9 +68,6 @@ def setup_optimizations( self.train_device, config.train_dtype, config.fallback_train_dtype, - [ - config.weight_dtypes().vae, - ], config.enable_autocast_cache, ) diff --git a/modules/modelSetup/BaseStableDiffusion3Setup.py b/modules/modelSetup/BaseStableDiffusion3Setup.py index 4f14d9f14..678d727ff 100644 --- a/modules/modelSetup/BaseStableDiffusion3Setup.py +++ b/modules/modelSetup/BaseStableDiffusion3Setup.py @@ -18,7 +18,6 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -56,26 +55,14 @@ def setup_optimizations( 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.weight_dtypes().transformer, - config.weight_dtypes().text_encoder, - config.weight_dtypes().text_encoder_2, - config.weight_dtypes().text_encoder_3, - config.weight_dtypes().vae, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache) + 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.weight_dtypes().text_encoder_3, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache, ) diff --git a/modules/modelSetup/BaseStableDiffusionSetup.py b/modules/modelSetup/BaseStableDiffusionSetup.py index f8aa5dec3..5be9fc97e 100644 --- a/modules/modelSetup/BaseStableDiffusionSetup.py +++ b/modules/modelSetup/BaseStableDiffusionSetup.py @@ -18,7 +18,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.enum.TrainingMethod import TrainingMethod from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -62,13 +61,8 @@ 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.weight_dtypes().text_encoder, - config.weight_dtypes().unet, - config.weight_dtypes().vae, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache) + 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) diff --git a/modules/modelSetup/BaseStableDiffusionXLSetup.py b/modules/modelSetup/BaseStableDiffusionXLSetup.py index 17d9cf4e0..e5c7bf0c3 100644 --- a/modules/modelSetup/BaseStableDiffusionXLSetup.py +++ b/modules/modelSetup/BaseStableDiffusionXLSetup.py @@ -18,7 +18,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_fp16_autocast_context -from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -60,22 +59,13 @@ 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.weight_dtypes().unet, - config.weight_dtypes().text_encoder, - config.weight_dtypes().text_encoder_2, - config.weight_dtypes().vae, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache) + 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, config.fallback_train_dtype, - [ - config.weight_dtypes().vae, - ], config.enable_autocast_cache, ) diff --git a/modules/modelSetup/BaseWuerstchenSetup.py b/modules/modelSetup/BaseWuerstchenSetup.py index 3ce486966..078e14fff 100644 --- a/modules/modelSetup/BaseWuerstchenSetup.py +++ b/modules/modelSetup/BaseWuerstchenSetup.py @@ -19,7 +19,6 @@ disable_bf16_on_fp16_autocast_context, disable_fp16_autocast_context, ) -from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -65,26 +64,14 @@ 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.weight_dtypes().decoder_text_encoder, - config.weight_dtypes().decoder, - config.weight_dtypes().decoder_vqgan, - config.weight_dtypes().effnet_encoder, - config.weight_dtypes().text_encoder, - config.weight_dtypes().prior, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - config.weight_dtypes().embedding if config.train_any_embedding() else None, - ], config.enable_autocast_cache) + 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, config.train_dtype, config.fallback_train_dtype, - [ - config.weight_dtypes().prior, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - ], config.enable_autocast_cache, ) else: diff --git a/modules/modelSetup/BaseZImageSetup.py b/modules/modelSetup/BaseZImageSetup.py index 3ea818952..a180f57d2 100644 --- a/modules/modelSetup/BaseZImageSetup.py +++ b/modules/modelSetup/BaseZImageSetup.py @@ -16,7 +16,6 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -50,12 +49,8 @@ def setup_optimizations( 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.weight_dtypes().transformer, - config.weight_dtypes().text_encoder, - config.weight_dtypes().vae, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - ], config.enable_autocast_cache) + 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 = \ @@ -63,10 +58,6 @@ def setup_optimizations( self.train_device, config.train_dtype, config.fallback_train_dtype, - [ - config.weight_dtypes().text_encoder, - config.weight_dtypes().lora if config.training_method == TrainingMethod.LORA else None, - ], config.enable_autocast_cache, ) diff --git a/modules/util/dtype_util.py b/modules/util/dtype_util.py index d0df10c59..b6d6b2643 100644 --- a/modules/util/dtype_util.py +++ b/modules/util/dtype_util.py @@ -1,20 +1,11 @@ from contextlib import nullcontext -from modules.util.config.TrainConfig import TrainConfig from modules.util.enum.DataType import DataType import torch from torch.nn import Parameter -def allow_mixed_precision(train_config: TrainConfig): - all_dtypes = list(train_config.weight_dtypes().all_dtypes() + [train_config.train_dtype]) - all_dtypes = list(filter(lambda dtype: dtype != DataType.NONE, all_dtypes)) - all_dtypes = set(all_dtypes) - - return len(all_dtypes) != 1 - - def enable_grad_scaling(train_dtype: DataType, parameters: list[Parameter]): trainable_parameter_dtype = list({parameter.dtype for parameter in parameters}) return train_dtype == DataType.FLOAT_16 and all(dtype == torch.float32 for dtype in trainable_parameter_dtype) @@ -28,46 +19,61 @@ def create_grad_scaler(): def create_autocast_context( device: torch.device, train_dtype: DataType | None, - weight_dtypes: list[DataType | None], enable_autocast_cache: bool, ) -> tuple[torch.autocast | nullcontext, DataType]: - if torch.backends.mps.is_available(): - if any(train_dtype != dt for dt in weight_dtypes if dt is not None): - print("Warning: Mixed precision training is untested on macOS. Consider setting all dtypes to be the same.") - else: - return nullcontext(), train_dtype - - weight_dtypes = list(weight_dtypes) - weight_dtypes = list(filter(lambda dtype: dtype != DataType.NONE and dtype is not None, weight_dtypes)) - weight_dtypes = list(set(weight_dtypes)) - - if len(weight_dtypes) == 1 and train_dtype == weight_dtypes[0]: - return torch.autocast(device_type=device.type, enabled=False), train_dtype - else: - return torch.autocast(device_type=device.type, dtype=train_dtype.torch_dtype(), + torch_train_dtype = train_dtype.torch_dtype() + + if torch_train_dtype in (torch.float16, torch.bfloat16): + # fp16/bf16 autocast is supported on every backend. autocast casts the operands + # of matmul/conv-type ops to train_dtype (precision-sensitive ops like norms stay + # in fp32), so a weight stored at a different dtype is cast on the fly rather than + # mismatching in the matmul. + if device.type != "cuda": + # CUDA (incl. ROCm) is the tested backend. fp16/bf16 autocast works on + # other backends too (mps, xpu, cpu, ...) but is untested here; bf16 on + # 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 + 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 + 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 + # dtype. Disable explicitly (not nullcontext) so any enclosing autocast is + # suppressed too. + 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 def disable_fp16_autocast_context( device: torch.device, train_dtype: DataType | None, fallback_train_dtype: DataType | None, - weight_dtypes: list[DataType | None], enable_autocast_cache: bool, ) -> tuple[torch.autocast | nullcontext, DataType]: - weight_dtypes = list(filter(lambda dtype: dtype != DataType.NONE and dtype is not None, weight_dtypes)) - weight_dtypes = list(set(weight_dtypes)) - if train_dtype != DataType.FLOAT_16: - # train dtype is not fp16 -> nothing to disable + # the main autocast context isn't fp16 -> nothing to override, defer to it return nullcontext(), train_dtype - if len(weight_dtypes) == 1 and fallback_train_dtype == weight_dtypes[0]: - # fallback_train_dtype is the same as all weights -> disable autocast - return torch.autocast(device_type=device.type, enabled=False), weight_dtypes[0] - - return torch.autocast(device_type=device.type, dtype=fallback_train_dtype.torch_dtype(), - cache_enabled=enable_autocast_cache), fallback_train_dtype + # fp16 training but this component is unstable in fp16 -> override the outer fp16 + # autocast and run it at the fallback precision. A bf16 fallback works on every + # backend; a float32 fallback can only be applied via autocast on CUDA. + fallback_torch_dtype = fallback_train_dtype.torch_dtype() + if fallback_torch_dtype in (torch.float16, torch.bfloat16) or device.type == "cuda": + return torch.autocast(device_type=device.type, dtype=fallback_torch_dtype, + cache_enabled=enable_autocast_cache), fallback_train_dtype + else: + raise RuntimeError( + f"A float32 fallback for fp16-unstable layers cannot be applied on device type " + f"'{device.type}' (only CUDA can autocast to float32). Use a bfloat16 fallback dtype." + ) def disable_bf16_on_fp16_autocast_context( @@ -76,6 +82,9 @@ def disable_bf16_on_fp16_autocast_context( weight_dtypes: list[DataType | None], enable_autocast_cache: bool, ) -> tuple[torch.autocast | nullcontext, DataType]: + # Only used for the Wuerstchen / Stable Cascade effnet encoder. The rationale for + # this special case is unknown, so its original behavior is deliberately kept + # unchanged rather than migrated to the create_autocast_context approach above. weight_dtypes = list(filter(lambda dtype: dtype != DataType.NONE and dtype is not None, weight_dtypes)) weight_dtypes = list(set(weight_dtypes))