From 414ca281f25b30178de715f77045ddbd32b11776 Mon Sep 17 00:00:00 2001 From: dxqb <183307934+dxqb@users.noreply.github.com> Date: Sat, 20 Jun 2026 17:30:53 +0200 Subject: [PATCH 1/8] Centralize model composition and training methods in ModelType Add a _MODEL_PARTS table + ModelType.model_parts() as the single source of truth for which components each model type has, keyed by TrainConfig field names, and a ModelType.supported_training_methods() that enumerates every type explicitly, raising on an unknown type rather than defaulting. Collapse ModelTab's per-type __setup_*_ui methods into one __setup_ui that derives the has_* widget flags from model_parts(), and collapse TopBar's per-type training-method dispatch to build its dropdown from supported_training_methods(). Co-Authored-By: Claude Sonnet 4.6 --- modules/ui/ModelTab.py | 335 +++++---------------------------- modules/ui/TopBar.py | 38 +--- modules/util/enum/ModelType.py | 55 ++++++ 3 files changed, 109 insertions(+), 319 deletions(-) diff --git a/modules/ui/ModelTab.py b/modules/ui/ModelTab.py index ff17ea3ba..5bc3092d3 100644 --- a/modules/ui/ModelTab.py +++ b/modules/ui/ModelTab.py @@ -45,307 +45,66 @@ def refresh_ui(self): base_frame.grid_columnconfigure(3, weight=0) base_frame.grid_columnconfigure(4, weight=1) - if self.train_config.model_type.is_stable_diffusion(): #TODO simplify - self.__setup_stable_diffusion_ui(base_frame) - if self.train_config.model_type.is_stable_diffusion_3(): - self.__setup_stable_diffusion_3_ui(base_frame) - elif self.train_config.model_type.is_stable_diffusion_xl(): - self.__setup_stable_diffusion_xl_ui(base_frame) - elif self.train_config.model_type.is_wuerstchen(): - self.__setup_wuerstchen_ui(base_frame) - elif self.train_config.model_type.is_pixart(): - self.__setup_pixart_alpha_ui(base_frame) - elif self.train_config.model_type.is_flux_1(): - self.__setup_flux_ui(base_frame) - elif self.train_config.model_type.is_flux_2(): - self.__setup_flux_2_ui(base_frame) - elif self.train_config.model_type.is_z_image(): - self.__setup_z_image_ui(base_frame) - elif self.train_config.model_type.is_chroma(): - self.__setup_chroma_ui(base_frame) - elif self.train_config.model_type.is_qwen(): - self.__setup_qwen_ui(base_frame) - elif self.train_config.model_type.is_sana(): - self.__setup_sana_ui(base_frame) - elif self.train_config.model_type.is_hunyuan_video(): - self.__setup_hunyuan_video_ui(base_frame) - elif self.train_config.model_type.is_hi_dream(): - self.__setup_hi_dream_ui(base_frame) - elif self.train_config.model_type.is_ernie(): - self.__setup_ernie_ui(base_frame) - - def __setup_stable_diffusion_ui(self, frame): - row = 0 - row = self.__create_base_dtype_components(frame, row) - row = self.__create_base_components( - frame, - row, - has_unet=True, - has_text_encoder=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - allow_safetensors=True, - allow_diffusers=self.train_config.training_method in [ - TrainingMethod.FINE_TUNE, - TrainingMethod.FINE_TUNE_VAE, - ], - allow_legacy_safetensors=self.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_stable_diffusion_3_ui(self, frame): - row = 0 - row = self.__create_base_dtype_components(frame, row) - row = self.__create_base_components( - frame, - row, - has_transformer=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_text_encoder_3=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - allow_safetensors=True, - allow_diffusers=self.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=self.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_flux_ui(self, frame): - row = 0 - row = self.__create_base_dtype_components(frame, row) - row = self.__create_base_components( - frame, - row, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - allow_safetensors=True, - allow_diffusers=self.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=self.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_flux_2_ui(self, frame): - row = 0 - row = self.__create_base_dtype_components(frame, row) - row = self.__create_base_components( - frame, - row, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - allow_safetensors=True, - allow_diffusers=self.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=self.train_config.training_method == TrainingMethod.LORA, - ) + self.__setup_ui(base_frame) - def __setup_z_image_ui(self, frame): - row = 0 - row = self.__create_base_dtype_components(frame, row) - row = self.__create_base_components( - frame, - row, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - allow_safetensors=True, - allow_diffusers=self.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=self.train_config.training_method == TrainingMethod.LORA, - ) + def __setup_ui(self, frame): + model_type = self.train_config.model_type + training_method = self.train_config.training_method + parts = model_type.model_parts() - def __setup_ernie_ui(self, frame): - row = 0 - row = self.__create_base_dtype_components(frame, row) - row = self.__create_base_components( - frame, - row, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - allow_safetensors=True, - allow_diffusers=self.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=self.train_config.training_method == TrainingMethod.LORA, + # The transformer override path exists only for these architectures; SD3, PixArt, Sana + # and HiDream have a transformer but expose no override field. + allow_override_transformer = ( + model_type.is_flux() + or model_type.is_z_image() + or model_type.is_ernie() + or model_type.is_chroma() + or model_type.is_qwen() + or model_type.is_hunyuan_video() ) - def __setup_chroma_ui(self, frame): row = 0 row = self.__create_base_dtype_components(frame, row) row = self.__create_base_components( frame, row, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - allow_safetensors=True, - allow_diffusers=self.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=self.train_config.training_method == TrainingMethod.LORA, - ) + has_unet="unet" in parts, + has_prior="prior" in parts, + allow_override_prior=model_type.is_stable_cascade(), + has_transformer="transformer" in parts, + allow_override_transformer=allow_override_transformer, + has_text_encoder=not model_type.has_multiple_text_encoders(), + has_text_encoder_1=model_type.has_multiple_text_encoders(), + has_text_encoder_2="text_encoder_2" in parts, + has_text_encoder_3="text_encoder_3" in parts, + has_text_encoder_4="text_encoder_4" in parts, + allow_override_text_encoder_4="text_encoder_4" in parts, + has_vae="vae" in parts, + ) + if "effnet_encoder" in parts: + row = self.__create_effnet_encoder_components(frame, row) + if "decoder" in parts: + row = self.__create_decoder_components(frame, row, "decoder_text_encoder" in parts) + + if model_type.is_sana(): + allow_safetensors = training_method != TrainingMethod.FINE_TUNE + elif model_type.is_wuerstchen(): + allow_safetensors = training_method != TrainingMethod.FINE_TUNE \ + or model_type.is_stable_cascade() + else: + allow_safetensors = True + + if model_type.is_stable_diffusion(): + allow_diffusers = training_method in [TrainingMethod.FINE_TUNE, TrainingMethod.FINE_TUNE_VAE] + else: + allow_diffusers = training_method == TrainingMethod.FINE_TUNE - def __setup_qwen_ui(self, frame): - row = 0 - row = self.__create_base_dtype_components(frame, row) - row = self.__create_base_components( - frame, - row, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - allow_safetensors=True, - allow_diffusers=self.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=self.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_stable_diffusion_xl_ui(self, frame): - row = 0 - row = self.__create_base_dtype_components(frame, row) - row = self.__create_base_components( - frame, - row, - has_unet=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - allow_safetensors=True, - allow_diffusers=self.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=self.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_wuerstchen_ui(self, frame): - row = 0 - row = self.__create_base_dtype_components(frame, row) - row = self.__create_base_components( - frame, - row, - has_prior=True, - allow_override_prior=self.train_config.model_type.is_stable_cascade(), - has_text_encoder=True, - ) - row = self.__create_effnet_encoder_components(frame, row) - row = self.__create_decoder_components(frame, row, self.train_config.model_type.is_wuerstchen_v2()) - row = self.__create_output_components( - frame, - row, - allow_safetensors=self.train_config.training_method != TrainingMethod.FINE_TUNE - or self.train_config.model_type.is_stable_cascade(), - allow_diffusers=self.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=self.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_pixart_alpha_ui(self, frame): - row = 0 - row = self.__create_base_dtype_components(frame, row) - row = self.__create_base_components( - frame, - row, - has_transformer=True, - has_text_encoder=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - allow_safetensors=True, - allow_diffusers=self.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=self.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_sana_ui(self, frame): - row = 0 - row = self.__create_base_dtype_components(frame, row) - row = self.__create_base_components( - frame, - row, - has_transformer=True, - has_text_encoder=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - allow_safetensors=self.train_config.training_method != TrainingMethod.FINE_TUNE, - allow_diffusers=self.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=self.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_hunyuan_video_ui(self, frame): - row = 0 - row = self.__create_base_dtype_components(frame, row) - row = self.__create_base_components( - frame, - row, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - allow_safetensors=True, - allow_diffusers=self.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=self.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_hi_dream_ui(self, frame): - row = 0 - row = self.__create_base_dtype_components(frame, row) - row = self.__create_base_components( - frame, - row, - has_transformer=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_text_encoder_3=True, - has_text_encoder_4=True, - allow_override_text_encoder_4=True, - has_vae=True, - ) row = self.__create_output_components( frame, row, - allow_safetensors=True, - allow_diffusers=self.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=self.train_config.training_method == TrainingMethod.LORA, + allow_safetensors=allow_safetensors, + allow_diffusers=allow_diffusers, + allow_legacy_safetensors=training_method == TrainingMethod.LORA, ) def __create_dtype_options(self, include_gguf: bool=False, include_a8: bool=False) -> list[tuple[str, DataType]]: diff --git a/modules/ui/TopBar.py b/modules/ui/TopBar.py index 820fdb71a..d2ca78f7b 100644 --- a/modules/ui/TopBar.py +++ b/modules/ui/TopBar.py @@ -112,37 +112,13 @@ def __create_training_method(self): if self.training_method: self.training_method.destroy() - values = [] - #TODO simplify - if self.train_config.model_type.is_stable_diffusion(): - values = [ - ("Fine Tune", TrainingMethod.FINE_TUNE), - ("LoRA", TrainingMethod.LORA), - ("Embedding", TrainingMethod.EMBEDDING), - ("Fine Tune VAE", TrainingMethod.FINE_TUNE_VAE), - ] - elif self.train_config.model_type.is_stable_diffusion_3() \ - or self.train_config.model_type.is_stable_diffusion_xl() \ - or self.train_config.model_type.is_wuerstchen() \ - or self.train_config.model_type.is_pixart() \ - or self.train_config.model_type.is_flux_1() \ - or self.train_config.model_type.is_sana() \ - or self.train_config.model_type.is_hunyuan_video() \ - or self.train_config.model_type.is_hi_dream() \ - or self.train_config.model_type.is_chroma(): - values = [ - ("Fine Tune", TrainingMethod.FINE_TUNE), - ("LoRA", TrainingMethod.LORA), - ("Embedding", TrainingMethod.EMBEDDING), - ] - elif self.train_config.model_type.is_qwen() \ - or self.train_config.model_type.is_z_image() \ - or self.train_config.model_type.is_flux_2() \ - or self.train_config.model_type.is_ernie(): - values = [ - ("Fine Tune", TrainingMethod.FINE_TUNE), - ("LoRA", TrainingMethod.LORA), - ] + labels = { + TrainingMethod.FINE_TUNE: "Fine Tune", + TrainingMethod.LORA: "LoRA", + TrainingMethod.EMBEDDING: "Embedding", + TrainingMethod.FINE_TUNE_VAE: "Fine Tune VAE", + } + values = [(labels[m], m) for m in self.train_config.model_type.supported_training_methods()] # training method self.training_method = components.options_kv( diff --git a/modules/util/enum/ModelType.py b/modules/util/enum/ModelType.py index a3ad940ec..a6e6a2a96 100644 --- a/modules/util/enum/ModelType.py +++ b/modules/util/enum/ModelType.py @@ -1,5 +1,7 @@ from enum import Enum +from modules.util.enum.TrainingMethod import TrainingMethod + class ModelType(Enum): STABLE_DIFFUSION_15 = 'STABLE_DIFFUSION_15' @@ -166,6 +168,59 @@ def is_flow_matching(self) -> bool: def is_video_model(self) -> bool: return self.is_hunyuan_video() #incase we add more video models in the future + def model_parts(self) -> tuple[str, ...]: + return _MODEL_PARTS[self] + + def supported_training_methods(self) -> tuple[TrainingMethod, ...]: + if self.is_stable_diffusion(): + return (TrainingMethod.FINE_TUNE, TrainingMethod.LORA, TrainingMethod.EMBEDDING, TrainingMethod.FINE_TUNE_VAE) + if self.is_stable_diffusion_3() \ + or self.is_stable_diffusion_xl() \ + or self.is_wuerstchen() \ + or self.is_pixart() \ + or self.is_flux_1() \ + or self.is_sana() \ + or self.is_hunyuan_video() \ + or self.is_hi_dream() \ + or self.is_chroma(): + return (TrainingMethod.FINE_TUNE, TrainingMethod.LORA, TrainingMethod.EMBEDDING) + if self.is_qwen() or self.is_z_image() or self.is_flux_2() or self.is_ernie(): + return (TrainingMethod.FINE_TUNE, TrainingMethod.LORA) + raise ValueError(f"No supported training methods defined for model type {self}") + + +# The first text encoder is always "text_encoder" here (matching the config field), even for +# multi-encoder models that refer to it as "text_encoder_1" elsewhere in the code. +_MODEL_PARTS: dict[ModelType, tuple[str, ...]] = { + ModelType.STABLE_DIFFUSION_15: ("text_encoder", "unet", "vae"), + ModelType.STABLE_DIFFUSION_15_INPAINTING: ("text_encoder", "unet", "vae"), + ModelType.STABLE_DIFFUSION_20: ("text_encoder", "unet", "vae"), + ModelType.STABLE_DIFFUSION_20_BASE: ("text_encoder", "unet", "vae"), + ModelType.STABLE_DIFFUSION_20_INPAINTING: ("text_encoder", "unet", "vae"), + ModelType.STABLE_DIFFUSION_20_DEPTH: ("text_encoder", "unet", "vae"), + ModelType.STABLE_DIFFUSION_21: ("text_encoder", "unet", "vae"), + ModelType.STABLE_DIFFUSION_21_BASE: ("text_encoder", "unet", "vae"), + ModelType.STABLE_DIFFUSION_3: ("text_encoder", "text_encoder_2", "text_encoder_3", "transformer", "vae"), + ModelType.STABLE_DIFFUSION_35: ("text_encoder", "text_encoder_2", "text_encoder_3", "transformer", "vae"), + ModelType.STABLE_DIFFUSION_XL_10_BASE: ("text_encoder", "text_encoder_2", "unet", "vae"), + ModelType.STABLE_DIFFUSION_XL_10_BASE_INPAINTING: ("text_encoder", "text_encoder_2", "unet", "vae"), + # Only Würstchen v2's decoder has its own text encoder; Stable Cascade's decoder does not. + ModelType.WUERSTCHEN_2: ("text_encoder", "prior", "effnet_encoder", "decoder", "decoder_text_encoder", "decoder_vqgan"), + ModelType.STABLE_CASCADE_1: ("text_encoder", "prior", "effnet_encoder", "decoder", "decoder_vqgan"), + ModelType.PIXART_ALPHA: ("text_encoder", "transformer", "vae"), + ModelType.PIXART_SIGMA: ("text_encoder", "transformer", "vae"), + ModelType.FLUX_DEV_1: ("text_encoder", "text_encoder_2", "transformer", "vae"), + ModelType.FLUX_FILL_DEV_1: ("text_encoder", "text_encoder_2", "transformer", "vae"), + ModelType.FLUX_2: ("text_encoder", "transformer", "vae"), + ModelType.SANA: ("text_encoder", "transformer", "vae"), + ModelType.HUNYUAN_VIDEO: ("text_encoder", "text_encoder_2", "transformer", "vae"), + ModelType.HI_DREAM_FULL: ("text_encoder", "text_encoder_2", "text_encoder_3", "text_encoder_4", "transformer", "vae"), + ModelType.CHROMA_1: ("text_encoder", "transformer", "vae"), + ModelType.QWEN: ("text_encoder", "transformer", "vae"), + ModelType.Z_IMAGE: ("text_encoder", "transformer", "vae"), + ModelType.ERNIE: ("text_encoder", "transformer", "vae"), +} + class PeftType(Enum): LORA = 'LORA' From f4b7d503703fedb38f38f9bd7038dfabbc7ce9f5 Mon Sep 17 00:00:00 2001 From: dxqb <183307934+dxqb@users.noreply.github.com> Date: Sat, 20 Jun 2026 20:39:00 +0200 Subject: [PATCH 2/8] Add per-component offloading/checkpointing from PR #1476 Rebased onto centralize-model-type: this is the offloading-only part of split-offload, with the model-composition centralization (ModelType, ModelTab, TopBar) excluded since it already landed separately. --- modules/model/ChromaModel.py | 6 +- modules/model/ErnieModel.py | 6 +- modules/model/Flux2Model.py | 6 +- modules/model/FluxModel.py | 6 +- modules/model/HiDreamModel.py | 9 +- modules/model/HunyuanVideoModel.py | 6 +- modules/model/PixArtAlphaModel.py | 6 +- modules/model/QwenModel.py | 6 +- modules/model/SanaModel.py | 6 +- modules/model/StableDiffusion3Model.py | 6 +- modules/model/ZImageModel.py | 6 +- modules/modelSetup/BaseChromaSetup.py | 9 +- modules/modelSetup/BaseErnieSetup.py | 7 +- modules/modelSetup/BaseFlux2Setup.py | 16 +-- modules/modelSetup/BaseFluxSetup.py | 13 +- modules/modelSetup/BaseHiDreamSetup.py | 22 ++- modules/modelSetup/BaseHunyuanVideoSetup.py | 13 +- modules/modelSetup/BasePixArtAlphaSetup.py | 8 +- modules/modelSetup/BaseQwenSetup.py | 9 +- modules/modelSetup/BaseSanaSetup.py | 8 +- .../modelSetup/BaseStableDiffusion3Setup.py | 17 +-- .../modelSetup/BaseStableDiffusionSetup.py | 9 +- .../modelSetup/BaseStableDiffusionXLSetup.py | 8 +- modules/modelSetup/BaseWuerstchenSetup.py | 4 +- modules/modelSetup/BaseZImageSetup.py | 9 +- modules/ui/OffloadingWindow.py | 75 ---------- modules/ui/TrainUI.py | 4 + modules/ui/TrainingTab.py | 135 ++++++++++++------ modules/util/LayerOffloadConductor.py | 11 +- modules/util/checkpointing_util.py | 134 ++++++++++------- modules/util/config/TrainConfig.py | 76 ++++++++-- modules/util/create.py | 9 +- .../util/enum/GradientCheckpointingMethod.py | 17 --- training_presets/#chroma Finetune 16GB.json | 7 +- training_presets/#chroma Finetune 8GB.json | 7 +- training_presets/#chroma LoRA 8GB.json | 7 +- training_presets/#ernie LoRA 8GB.json | 7 +- training_presets/#flux2 Finetune 16GB.json | 7 +- training_presets/#flux2 LoRA 8GB.json | 10 +- training_presets/#hidream LoRA.json | 11 +- training_presets/#hunyuan video LoRA.json | 8 +- training_presets/#qwen Finetune 16GB.json | 5 +- training_presets/#qwen Finetune 24GB.json | 5 +- training_presets/#qwen LoRA 16GB.json | 5 +- training_presets/#qwen LoRA 24GB.json | 5 +- .../#z-image DeTurbo LoRA 8GB.json | 5 +- training_presets/#z-image Finetune 16GB.json | 5 +- training_presets/#z-image LoRA 8GB.json | 5 +- 48 files changed, 383 insertions(+), 398 deletions(-) delete mode 100644 modules/ui/OffloadingWindow.py delete mode 100644 modules/util/enum/GradientCheckpointingMethod.py diff --git a/modules/model/ChromaModel.py b/modules/model/ChromaModel.py index 59967c8dc..3fc96d8f9 100644 --- a/modules/model/ChromaModel.py +++ b/modules/model/ChromaModel.py @@ -113,8 +113,7 @@ def vae_to(self, device: torch.device): def text_encoder_to(self, device: torch.device): if self.text_encoder is not None: - if self.text_encoder_offload_conductor is not None and \ - self.text_encoder_offload_conductor.layer_offload_activated(): + if self.text_encoder_offload_conductor is not None: self.text_encoder_offload_conductor.to(device) else: self.text_encoder.to(device=device) @@ -123,8 +122,7 @@ def text_encoder_to(self, device: torch.device): self.text_encoder_lora.to(device) def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None and \ - self.transformer_offload_conductor.layer_offload_activated(): + if self.transformer_offload_conductor is not None: self.transformer_offload_conductor.to(device) else: self.transformer.to(device=device) diff --git a/modules/model/ErnieModel.py b/modules/model/ErnieModel.py index fb44425e8..3d15e5385 100644 --- a/modules/model/ErnieModel.py +++ b/modules/model/ErnieModel.py @@ -72,15 +72,13 @@ def vae_to(self, device: torch.device): def text_encoder_to(self, device: torch.device): if self.text_encoder is not None: - if self.text_encoder_offload_conductor is not None and \ - self.text_encoder_offload_conductor.layer_offload_activated(): + if self.text_encoder_offload_conductor is not None: self.text_encoder_offload_conductor.to(device) else: self.text_encoder.to(device=device) def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None and \ - self.transformer_offload_conductor.layer_offload_activated(): + if self.transformer_offload_conductor is not None: self.transformer_offload_conductor.to(device) else: self.transformer.to(device=device) diff --git a/modules/model/Flux2Model.py b/modules/model/Flux2Model.py index 004c70bee..a501bb116 100644 --- a/modules/model/Flux2Model.py +++ b/modules/model/Flux2Model.py @@ -122,15 +122,13 @@ def vae_to(self, device: torch.device): def text_encoder_to(self, device: torch.device): if self.text_encoder is not None: - if self.text_encoder_offload_conductor is not None and \ - self.text_encoder_offload_conductor.layer_offload_activated(): + if self.text_encoder_offload_conductor is not None: self.text_encoder_offload_conductor.to(device) else: self.text_encoder.to(device=device) def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None and \ - self.transformer_offload_conductor.layer_offload_activated(): + if self.transformer_offload_conductor is not None: self.transformer_offload_conductor.to(device) else: self.transformer.to(device=device) diff --git a/modules/model/FluxModel.py b/modules/model/FluxModel.py index b981865c4..b790335ac 100644 --- a/modules/model/FluxModel.py +++ b/modules/model/FluxModel.py @@ -149,8 +149,7 @@ def text_encoder_1_to(self, device: torch.device): def text_encoder_2_to(self, device: torch.device): if self.text_encoder_2 is not None: - if self.text_encoder_2_offload_conductor is not None and \ - self.text_encoder_2_offload_conductor.layer_offload_activated(): + if self.text_encoder_2_offload_conductor is not None: self.text_encoder_2_offload_conductor.to(device) else: self.text_encoder_2.to(device=device) @@ -159,8 +158,7 @@ def text_encoder_2_to(self, device: torch.device): self.text_encoder_2_lora.to(device) def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None and \ - self.transformer_offload_conductor.layer_offload_activated(): + if self.transformer_offload_conductor is not None: self.transformer_offload_conductor.to(device) else: self.transformer.to(device=device) diff --git a/modules/model/HiDreamModel.py b/modules/model/HiDreamModel.py index 049b0e6d3..9e9684851 100644 --- a/modules/model/HiDreamModel.py +++ b/modules/model/HiDreamModel.py @@ -220,8 +220,7 @@ def text_encoder_2_to(self, device: torch.device): def text_encoder_3_to(self, device: torch.device): if self.text_encoder_3 is not None: - if self.text_encoder_3_offload_conductor is not None and \ - self.text_encoder_3_offload_conductor.layer_offload_activated(): + if self.text_encoder_3_offload_conductor is not None: self.text_encoder_3_offload_conductor.to(device) else: self.text_encoder_3.to(device=device) @@ -231,8 +230,7 @@ def text_encoder_3_to(self, device: torch.device): def text_encoder_4_to(self, device: torch.device): if self.text_encoder_4 is not None: - if self.text_encoder_4_offload_conductor is not None and \ - self.text_encoder_4_offload_conductor.layer_offload_activated(): + if self.text_encoder_4_offload_conductor is not None: self.text_encoder_4_offload_conductor.to(device) else: self.text_encoder_4.to(device=device) @@ -241,8 +239,7 @@ def text_encoder_4_to(self, device: torch.device): self.text_encoder_4_lora.to(device) def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None and \ - self.transformer_offload_conductor.layer_offload_activated(): + if self.transformer_offload_conductor is not None: self.transformer_offload_conductor.to(device) else: self.transformer.to(device=device) diff --git a/modules/model/HunyuanVideoModel.py b/modules/model/HunyuanVideoModel.py index 2107e55d7..bfc07ca95 100644 --- a/modules/model/HunyuanVideoModel.py +++ b/modules/model/HunyuanVideoModel.py @@ -157,8 +157,7 @@ def text_encoder_to(self, device: torch.device): def text_encoder_1_to(self, device: torch.device): if self.text_encoder_1 is not None: - if self.text_encoder_1_offload_conductor is not None and \ - self.text_encoder_1_offload_conductor.layer_offload_activated(): + if self.text_encoder_1_offload_conductor is not None: self.text_encoder_1_offload_conductor.to(device) else: self.text_encoder_1.to(device=device) @@ -174,8 +173,7 @@ def text_encoder_2_to(self, device: torch.device): self.text_encoder_2_lora.to(device) def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None and \ - self.transformer_offload_conductor.layer_offload_activated(): + if self.transformer_offload_conductor is not None: self.transformer_offload_conductor.to(device) else: self.transformer.to(device=device) diff --git a/modules/model/PixArtAlphaModel.py b/modules/model/PixArtAlphaModel.py index 466cc61f9..42c6621f8 100644 --- a/modules/model/PixArtAlphaModel.py +++ b/modules/model/PixArtAlphaModel.py @@ -114,8 +114,7 @@ def vae_to(self, device: torch.device): self.vae.to(device=device) def text_encoder_to(self, device: torch.device): - if self.text_encoder_offload_conductor is not None and \ - self.text_encoder_offload_conductor.layer_offload_activated(): + if self.text_encoder_offload_conductor is not None: self.text_encoder_offload_conductor.to(device) else: self.text_encoder.to(device=device) @@ -124,8 +123,7 @@ def text_encoder_to(self, device: torch.device): self.text_encoder_lora.to(device) def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None and \ - self.transformer_offload_conductor.layer_offload_activated(): + if self.transformer_offload_conductor is not None: self.transformer_offload_conductor.to(device) else: self.transformer.to(device=device) diff --git a/modules/model/QwenModel.py b/modules/model/QwenModel.py index afa6c24fe..1aadcd467 100644 --- a/modules/model/QwenModel.py +++ b/modules/model/QwenModel.py @@ -81,8 +81,7 @@ def vae_to(self, device: torch.device): def text_encoder_to(self, device: torch.device): #TODO share more code between models if self.text_encoder is not None: - if self.text_encoder_offload_conductor is not None and \ - self.text_encoder_offload_conductor.layer_offload_activated(): + if self.text_encoder_offload_conductor is not None: self.text_encoder_offload_conductor.to(device) else: self.text_encoder.to(device=device) @@ -91,8 +90,7 @@ def text_encoder_to(self, device: torch.device): #TODO share more code between m self.text_encoder_lora.to(device) def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None and \ - self.transformer_offload_conductor.layer_offload_activated(): + if self.transformer_offload_conductor is not None: self.transformer_offload_conductor.to(device) else: self.transformer.to(device=device) diff --git a/modules/model/SanaModel.py b/modules/model/SanaModel.py index 9e8008219..c1d563cb2 100644 --- a/modules/model/SanaModel.py +++ b/modules/model/SanaModel.py @@ -116,8 +116,7 @@ def vae_to(self, device: torch.device): self.vae.to(device=device) def text_encoder_to(self, device: torch.device): - if self.text_encoder_offload_conductor is not None and \ - self.text_encoder_offload_conductor.layer_offload_activated(): + if self.text_encoder_offload_conductor is not None: self.text_encoder_offload_conductor.to(device) else: self.text_encoder.to(device=device) @@ -126,8 +125,7 @@ def text_encoder_to(self, device: torch.device): self.text_encoder_lora.to(device) def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None and \ - self.transformer_offload_conductor.layer_offload_activated(): + if self.transformer_offload_conductor is not None: self.transformer_offload_conductor.to(device) else: self.transformer.to(device=device) diff --git a/modules/model/StableDiffusion3Model.py b/modules/model/StableDiffusion3Model.py index 8f6cf5818..a2bae894d 100644 --- a/modules/model/StableDiffusion3Model.py +++ b/modules/model/StableDiffusion3Model.py @@ -180,8 +180,7 @@ def text_encoder_2_to(self, device: torch.device): def text_encoder_3_to(self, device: torch.device): if self.text_encoder_3 is not None: - if self.text_encoder_3_offload_conductor is not None and \ - self.text_encoder_3_offload_conductor.layer_offload_activated(): + if self.text_encoder_3_offload_conductor is not None: self.text_encoder_3_offload_conductor.to(device) else: self.text_encoder_3.to(device=device) @@ -190,8 +189,7 @@ def text_encoder_3_to(self, device: torch.device): self.text_encoder_3_lora.to(device) def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None and \ - self.transformer_offload_conductor.layer_offload_activated(): + if self.transformer_offload_conductor is not None: self.transformer_offload_conductor.to(device) else: self.transformer.to(device=device) diff --git a/modules/model/ZImageModel.py b/modules/model/ZImageModel.py index 7fd9e52cb..167b58f47 100644 --- a/modules/model/ZImageModel.py +++ b/modules/model/ZImageModel.py @@ -83,15 +83,13 @@ def vae_to(self, device: torch.device): def text_encoder_to(self, device: torch.device): #TODO share more code between models if self.text_encoder is not None: - if self.text_encoder_offload_conductor is not None and \ - self.text_encoder_offload_conductor.layer_offload_activated(): + if self.text_encoder_offload_conductor is not None: self.text_encoder_offload_conductor.to(device) else: self.text_encoder.to(device=device) def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None and \ - self.transformer_offload_conductor.layer_offload_activated(): + if self.transformer_offload_conductor is not None: self.transformer_offload_conductor.to(device) else: self.transformer.to(device=device) diff --git a/modules/modelSetup/BaseChromaSetup.py b/modules/modelSetup/BaseChromaSetup.py index 0eb623399..cf2918bc1 100644 --- a/modules/modelSetup/BaseChromaSetup.py +++ b/modules/modelSetup/BaseChromaSetup.py @@ -49,12 +49,9 @@ def setup_optimizations( model: ChromaModel, config: TrainConfig, ): - if config.gradient_checkpointing.enabled(): - model.transformer_offload_conductor = \ - enable_checkpointing_for_chroma_transformer(model.transformer, config) - if model.text_encoder is not None: - model.text_encoder_offload_conductor = \ - enable_checkpointing_for_t5_encoder_layers(model.text_encoder, config) + model.transformer_offload_conductor = enable_checkpointing_for_chroma_transformer(model.transformer, config, config.transformer) + if model.text_encoder is not None: + 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, diff --git a/modules/modelSetup/BaseErnieSetup.py b/modules/modelSetup/BaseErnieSetup.py index 3d630b8cf..da222a5ed 100644 --- a/modules/modelSetup/BaseErnieSetup.py +++ b/modules/modelSetup/BaseErnieSetup.py @@ -45,11 +45,8 @@ def setup_optimizations( model: ErnieModel, config: TrainConfig, ): - if config.gradient_checkpointing.enabled(): - model.transformer_offload_conductor = \ - enable_checkpointing_for_ernie_transformer(model.transformer, config) - model.text_encoder_offload_conductor = \ - enable_checkpointing_for_mistral_encoder_layers(model.text_encoder, config) + 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, diff --git a/modules/modelSetup/BaseFlux2Setup.py b/modules/modelSetup/BaseFlux2Setup.py index b91cd8af8..019f0f179 100644 --- a/modules/modelSetup/BaseFlux2Setup.py +++ b/modules/modelSetup/BaseFlux2Setup.py @@ -45,16 +45,12 @@ def setup_optimizations( model: Flux2Model, config: TrainConfig, ): - if config.gradient_checkpointing.enabled(): - model.transformer_offload_conductor = \ - enable_checkpointing_for_flux2_transformer(model.transformer, config) - if model.text_encoder is not None: - if model.is_dev(): - model.text_encoder_offload_conductor = \ - enable_checkpointing_for_mistral_encoder_layers(model.text_encoder, config) - else: - model.text_encoder_offload_conductor = \ - enable_checkpointing_for_qwen3_encoder_layers(model.text_encoder, config) + model.transformer_offload_conductor = enable_checkpointing_for_flux2_transformer(model.transformer, config, config.transformer) + if model.text_encoder is not None: + 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.weight_dtypes().transformer, diff --git a/modules/modelSetup/BaseFluxSetup.py b/modules/modelSetup/BaseFluxSetup.py index 9bae83cde..8eae6211d 100644 --- a/modules/modelSetup/BaseFluxSetup.py +++ b/modules/modelSetup/BaseFluxSetup.py @@ -49,14 +49,11 @@ def setup_optimizations( model: FluxModel, config: TrainConfig, ): - if config.gradient_checkpointing.enabled(): - model.transformer_offload_conductor = \ - enable_checkpointing_for_flux_transformer(model.transformer, config) - if model.text_encoder_1 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_1, config) - 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) + 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.weight_dtypes().transformer, diff --git a/modules/modelSetup/BaseHiDreamSetup.py b/modules/modelSetup/BaseHiDreamSetup.py index 17fbcc0d6..b5d742fb8 100644 --- a/modules/modelSetup/BaseHiDreamSetup.py +++ b/modules/modelSetup/BaseHiDreamSetup.py @@ -49,19 +49,15 @@ def setup_optimizations( model: HiDreamModel, config: TrainConfig, ): - if config.gradient_checkpointing.enabled(): - model.transformer_offload_conductor = \ - enable_checkpointing_for_hi_dream_transformer(model.transformer, config) - if model.text_encoder_1 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_1, config) - if model.text_encoder_2 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_2, config) - 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) - 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) + 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.weight_dtypes().transformer, diff --git a/modules/modelSetup/BaseHunyuanVideoSetup.py b/modules/modelSetup/BaseHunyuanVideoSetup.py index b072bf4ba..b1dbebca8 100644 --- a/modules/modelSetup/BaseHunyuanVideoSetup.py +++ b/modules/modelSetup/BaseHunyuanVideoSetup.py @@ -49,14 +49,11 @@ def setup_optimizations( model: HunyuanVideoModel, config: TrainConfig, ): - if config.gradient_checkpointing.enabled(): - model.transformer_offload_conductor = \ - enable_checkpointing_for_hunyuan_video_transformer(model.transformer, config) - 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) - if model.text_encoder_2 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_2, config) + 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.weight_dtypes().transformer, diff --git a/modules/modelSetup/BasePixArtAlphaSetup.py b/modules/modelSetup/BasePixArtAlphaSetup.py index 3069b4884..0d6dc2000 100644 --- a/modules/modelSetup/BasePixArtAlphaSetup.py +++ b/modules/modelSetup/BasePixArtAlphaSetup.py @@ -51,12 +51,8 @@ def setup_optimizations( model: PixArtAlphaModel, config: TrainConfig, ): - if config.gradient_checkpointing.enabled(): - model.vae.enable_gradient_checkpointing() - model.transformer_offload_conductor = \ - enable_checkpointing_for_basic_transformer_blocks(model.transformer, config, offload_enabled=True) - model.text_encoder_offload_conductor = \ - enable_checkpointing_for_t5_encoder_layers(model.text_encoder, config) + 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, diff --git a/modules/modelSetup/BaseQwenSetup.py b/modules/modelSetup/BaseQwenSetup.py index a8a7be8f6..0805358be 100644 --- a/modules/modelSetup/BaseQwenSetup.py +++ b/modules/modelSetup/BaseQwenSetup.py @@ -46,12 +46,9 @@ def setup_optimizations( model: QwenModel, config: TrainConfig, ): - if config.gradient_checkpointing.enabled(): - model.transformer_offload_conductor = \ - enable_checkpointing_for_qwen_transformer(model.transformer, config) - if model.text_encoder is not None: - model.text_encoder_offload_conductor = \ - enable_checkpointing_for_qwen25vl_encoder_layers(model.text_encoder, config) + model.transformer_offload_conductor = enable_checkpointing_for_qwen_transformer(model.transformer, config, config.transformer) + if model.text_encoder is not None: + 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, diff --git a/modules/modelSetup/BaseSanaSetup.py b/modules/modelSetup/BaseSanaSetup.py index 0c9ea6da0..d265fee2a 100644 --- a/modules/modelSetup/BaseSanaSetup.py +++ b/modules/modelSetup/BaseSanaSetup.py @@ -52,12 +52,8 @@ def setup_optimizations( config: TrainConfig, ): - if config.gradient_checkpointing.enabled(): - # model.vae.enable_gradient_checkpointing() - model.transformer_offload_conductor = \ - enable_checkpointing_for_sana_transformer(model.transformer, config) - model.text_encoder_offload_conductor = \ - enable_checkpointing_for_gemma_layers(model.text_encoder, config) + 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, diff --git a/modules/modelSetup/BaseStableDiffusion3Setup.py b/modules/modelSetup/BaseStableDiffusion3Setup.py index de5dc04e8..ee70847fa 100644 --- a/modules/modelSetup/BaseStableDiffusion3Setup.py +++ b/modules/modelSetup/BaseStableDiffusion3Setup.py @@ -48,16 +48,13 @@ def setup_optimizations( model: StableDiffusion3Model, config: TrainConfig, ): - if config.gradient_checkpointing.enabled(): - model.transformer_offload_conductor = \ - enable_checkpointing_for_stable_diffusion_3_transformer(model.transformer, config) - if model.text_encoder_1 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_1, config) - if model.text_encoder_2 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_2, config) - 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) + 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.weight_dtypes().transformer, diff --git a/modules/modelSetup/BaseStableDiffusionSetup.py b/modules/modelSetup/BaseStableDiffusionSetup.py index 8cf63ac07..3060816f7 100644 --- a/modules/modelSetup/BaseStableDiffusionSetup.py +++ b/modules/modelSetup/BaseStableDiffusionSetup.py @@ -51,11 +51,12 @@ def setup_optimizations( model: StableDiffusionModel, config: TrainConfig, ): - if config.gradient_checkpointing.enabled(): - model.vae.enable_gradient_checkpointing() + if config.unet.checkpointing_or_offloading_enabled(): model.unet.enable_gradient_checkpointing() - enable_checkpointing_for_basic_transformer_blocks(model.unet, config, offload_enabled=False) - enable_checkpointing_for_clip_encoder_layers(model.text_encoder, config) + enable_checkpointing_for_basic_transformer_blocks(model.unet, config, config.unet, offload_enabled=False) + if config.vae.checkpointing_enabled(): + model.vae.enable_gradient_checkpointing() + enable_checkpointing_for_clip_encoder_layers(model.text_encoder, config, config.text_encoder) if config.force_circular_padding: apply_circular_padding_to_conv2d(model.vae) diff --git a/modules/modelSetup/BaseStableDiffusionXLSetup.py b/modules/modelSetup/BaseStableDiffusionXLSetup.py index 61cb1e457..2d0057b8b 100644 --- a/modules/modelSetup/BaseStableDiffusionXLSetup.py +++ b/modules/modelSetup/BaseStableDiffusionXLSetup.py @@ -48,11 +48,11 @@ def setup_optimizations( model: StableDiffusionXLModel, config: TrainConfig, ): - if config.gradient_checkpointing.enabled(): + if config.unet.checkpointing_or_offloading_enabled(): model.unet.enable_gradient_checkpointing() - enable_checkpointing_for_basic_transformer_blocks(model.unet, config, offload_enabled=False) - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_1, config) - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_2, config) + enable_checkpointing_for_basic_transformer_blocks(model.unet, config, config.unet, offload_enabled=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) if config.force_circular_padding: apply_circular_padding_to_conv2d(model.vae) diff --git a/modules/modelSetup/BaseWuerstchenSetup.py b/modules/modelSetup/BaseWuerstchenSetup.py index e1a4c39d2..5ba478afc 100644 --- a/modules/modelSetup/BaseWuerstchenSetup.py +++ b/modules/modelSetup/BaseWuerstchenSetup.py @@ -54,9 +54,9 @@ def setup_optimizations( model: WuerstchenModel, config: TrainConfig, ): - if config.gradient_checkpointing.enabled(): + if config.prior.checkpointing_enabled(): model.prior_prior.enable_gradient_checkpointing() - enable_checkpointing_for_clip_encoder_layers(model.prior_text_encoder, config) + enable_checkpointing_for_clip_encoder_layers(model.prior_text_encoder, config, config.text_encoder) if config.force_circular_padding: apply_circular_padding_to_conv2d(model.decoder_vqgan) diff --git a/modules/modelSetup/BaseZImageSetup.py b/modules/modelSetup/BaseZImageSetup.py index b822d7304..2a0287c19 100644 --- a/modules/modelSetup/BaseZImageSetup.py +++ b/modules/modelSetup/BaseZImageSetup.py @@ -47,12 +47,9 @@ def setup_optimizations( model: ZImageModel, config: TrainConfig, ): - if config.gradient_checkpointing.enabled(): - model.transformer_offload_conductor = \ - enable_checkpointing_for_z_image_transformer(model.transformer, config) - if model.text_encoder is not None: - model.text_encoder_offload_conductor = \ - enable_checkpointing_for_qwen3_encoder_layers(model.text_encoder, config) + model.transformer_offload_conductor = enable_checkpointing_for_z_image_transformer(model.transformer, config, config.transformer) + if model.text_encoder is not None: + 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, diff --git a/modules/ui/OffloadingWindow.py b/modules/ui/OffloadingWindow.py deleted file mode 100644 index 54035e121..000000000 --- a/modules/ui/OffloadingWindow.py +++ /dev/null @@ -1,75 +0,0 @@ -from modules.util.config.TrainConfig import TrainConfig -from modules.util.enum.GradientCheckpointingMethod import ( - GradientCheckpointingMethod, -) -from modules.util.ui import components -from modules.util.ui.ui_utils import set_window_icon -from modules.util.ui.UIState import UIState - -import customtkinter as ctk - - -class OffloadingWindow(ctk.CTkToplevel): - def __init__( - self, - parent, - config: TrainConfig, - ui_state: UIState, - *args, **kwargs, - ): - super().__init__(parent, *args, **kwargs) - - self.config = config - self.ui_state = ui_state - self.image_preview_file_index = 0 - self.ax = None - self.canvas = None - - self.title("Offloading") - self.geometry("800x400") - self.resizable(True, True) - - self.grid_rowconfigure(0, weight=1) - self.grid_columnconfigure(0, weight=1) - - frame = self.__content_frame(self) - frame.grid(row=0, column=0, sticky='nsew') - components.button(self, 1, 0, "ok", self.__ok) - - self.wait_visibility() - self.grab_set() - self.focus_set() - self.after(200, lambda: set_window_icon(self)) - - - def __content_frame(self, master): - frame = ctk.CTkScrollableFrame(master, fg_color="transparent") - frame.grid_columnconfigure(0, weight=1) - frame.grid_columnconfigure(1, weight=1) - - # timestep distribution - components.label(frame, 0, 0, "Gradient checkpointing", - tooltip="Enables gradient checkpointing. This reduces memory usage, but increases training time") - components.options(frame, 0, 1, [str(x) for x in list(GradientCheckpointingMethod)], self.ui_state, - "gradient_checkpointing") - - # gradient checkpointing layer offloading - components.label(frame, 1, 0, "Async Offloading", - tooltip="Enables Asynchronous offloading.") - components.switch(frame, 1, 1, self.ui_state, "enable_async_offloading") - - # gradient checkpointing layer offloading - components.label(frame, 2, 0, "Offload Activations", - tooltip="Enables Activation Offloading") - components.switch(frame, 2, 1, self.ui_state, "enable_activation_offloading") - - # gradient checkpointing layer offloading - components.label(frame, 3, 0, "Layer offload fraction", - tooltip="Enables offloading of individual layers during training to reduce VRAM usage. Increases training time and uses more RAM. Only available if checkpointing is set to CPU_OFFLOADED. values between 0 and 1, 0=disabled") - components.entry(frame, 3, 1, self.ui_state, "layer_offload_fraction") - - frame.pack(fill="both", expand=1) - return frame - - def __ok(self): - self.destroy() diff --git a/modules/ui/TrainUI.py b/modules/ui/TrainUI.py index b9fa0c04a..1217c06b4 100644 --- a/modules/ui/TrainUI.py +++ b/modules/ui/TrainUI.py @@ -311,6 +311,10 @@ def create_general_tab(self, master): tooltip="The device used for training. Can be \"cuda\", \"cuda:0\", \"cuda:1\" etc. Default:\"cuda\". Must be \"cuda\" for multi-GPU training.") components.entry(frame, 11, 1, self.ui_state, "train_device", required=True) + components.label(frame, 11, 2, "Async Offloading", + tooltip="Overlaps CPU<->GPU transfers with computation using CUDA streams. Applies to every offloaded component") + components.switch(frame, 11, 3, self.ui_state, "async_offloading") + components.label(frame, 12, 0, "Multi-GPU", tooltip="Enable multi-GPU training") components.switch(frame, 12, 1, self.ui_state, "multi_gpu") diff --git a/modules/ui/TrainingTab.py b/modules/ui/TrainingTab.py index f897bb8ce..29ffa2fb7 100644 --- a/modules/ui/TrainingTab.py +++ b/modules/ui/TrainingTab.py @@ -1,4 +1,3 @@ -from modules.ui.OffloadingWindow import OffloadingWindow from modules.ui.OptimizerParamsWindow import OptimizerParamsWindow from modules.ui.SchedulerParamsWindow import SchedulerParamsWindow from modules.ui.TimestepDistributionWindow import TimestepDistributionWindow @@ -6,13 +5,13 @@ from modules.util.config.TrainConfig import TrainConfig from modules.util.enum.DataType import DataType from modules.util.enum.EMAMode import EMAMode -from modules.util.enum.GradientCheckpointingMethod import GradientCheckpointingMethod from modules.util.enum.LearningRateScaler import LearningRateScaler from modules.util.enum.LearningRateScheduler import LearningRateScheduler from modules.util.enum.LossScaler import LossScaler from modules.util.enum.LossWeight import LossWeight from modules.util.enum.Optimizer import Optimizer from modules.util.enum.TimestepDistribution import TimestepDistribution +from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.optimizer_util import change_optimizer from modules.util.ui import components from modules.util.ui.UIState import UIState @@ -99,6 +98,8 @@ def __setup_stable_diffusion_ui(self, column_0, column_1, column_2): self.__create_unet_frame(column_1, 1) self.__create_noise_frame(column_1, 2, supports_generalized_offset_noise=True) + if self.train_config.training_method == TrainingMethod.FINE_TUNE_VAE: + self.__create_vae_frame(column_2, 0) self.__create_masked_frame(column_2, 1) self.__create_loss_frame(column_2, 2) self.__create_layer_frame(column_2, 3) @@ -375,19 +376,6 @@ def __create_base2_frame(self, master, row, video_training_enabled: bool=False, components.entry(frame, row, 1, self.ui_state, "ema_update_step_interval") row += 1 - # gradient checkpointing - components.label(frame, row, 0, "Gradient checkpointing", - tooltip="Enables gradient checkpointing. This reduces memory usage, but increases training time") - components.options_adv(frame, row, 1, [str(x) for x in list(GradientCheckpointingMethod)], self.ui_state, - "gradient_checkpointing", adv_command=self.__open_offloading_window) - row += 1 - - # gradient checkpointing layer offloading - components.label(frame, row, 0, "Layer offload fraction", - tooltip="Enables offloading of individual layers during training to reduce VRAM usage. Increases training time and uses more RAM. Only available if checkpointing is set to CPU_OFFLOADED. values between 0 and 1, 0=disabled") - components.entry(frame, row, 1, self.ui_state, "layer_offload_fraction") - row += 1 - # train dtype components.label(frame, row, 0, "Train Data Type", tooltip="The mixed precision data type used for training. This can increase training speed, but reduces precision") @@ -434,6 +422,26 @@ def __create_base2_frame(self, master, row, video_training_enabled: bool=False, tooltip="Enables circular padding for all conv layers to better train seamless images") components.switch(frame, row, 1, self.ui_state, "force_circular_padding") + def __create_offloading_widgets(self, frame, row, part, supports_checkpointing=True, supports_activation_offloading=False): + if supports_checkpointing: + components.label(frame, row, 0, "Gradient Checkpointing", + tooltip="Enables gradient checkpointing for this component. Reduces VRAM usage at the cost of training speed") + components.switch(frame, row, 1, self.ui_state, f"{part}.gradient_checkpointing") + row += 1 + + components.label(frame, row, 0, "Layer Offload Fraction", + tooltip="Fraction of this component's layers to offload to CPU to reduce VRAM usage. Increases training time and RAM usage. 0=disabled, 1=all layers") + components.entry(frame, row, 1, self.ui_state, f"{part}.offload_fraction") + row += 1 + + if supports_activation_offloading: + components.label(frame, row, 0, "Offload Activations", + tooltip="Offloads this component's activations to CPU during training to reduce VRAM usage") + components.switch(frame, row, 1, self.ui_state, f"{part}.activation_offloading") + row += 1 + + return row + def __create_text_encoder_frame(self, master, row, supports_clip_skip=True, supports_training=True, supports_sequence_length=False): frame = ctk.CTkFrame(master=master, corner_radius=5) frame.grid(row=row, column=0, padx=5, pady=5, sticky="nsew") @@ -445,6 +453,12 @@ def __create_text_encoder_frame(self, master, row, supports_clip_skip=True, supp tooltip="Enables training the text encoder model") components.switch(frame, row, 1, self.ui_state, "text_encoder.train") row += 1 + else: + # no Train switch to act as the frame's header, so add an explicit one + components.label(frame, row, 0, "Text Encoder") + row += 1 + + row = self.__create_offloading_widgets(frame, row, "text_encoder", supports_checkpointing=supports_training) # dropout components.label(frame, row, 0, "Caption Dropout Probability", @@ -509,6 +523,8 @@ def __create_text_encoder_n_frame( components.switch(frame, row, 1, self.ui_state, f"text_encoder{suffix}.train") row += 1 + row = self.__create_offloading_widgets(frame, row, f"text_encoder{suffix}") + # train text encoder embedding components.label(frame, row, 0, f"Train Text Encoder {i} Embedding", tooltip=f"Enables training embeddings for the text encoder {i} model") @@ -566,82 +582,119 @@ def __create_unet_frame(self, master, row): frame = ctk.CTkFrame(master=master, corner_radius=5) frame.grid(row=row, column=0, padx=5, pady=5, sticky="nsew") frame.grid_columnconfigure(0, weight=1) + row = 0 # train unet - components.label(frame, 0, 0, "Train UNet", + components.label(frame, row, 0, "Train UNet", tooltip="Enables training the UNet model") - components.switch(frame, 0, 1, self.ui_state, "unet.train") + components.switch(frame, row, 1, self.ui_state, "unet.train") + row += 1 + + row = self.__create_offloading_widgets(frame, row, "unet", supports_activation_offloading=True) # train unet epochs - components.label(frame, 1, 0, "Stop Training After", + components.label(frame, row, 0, "Stop Training After", tooltip="When to stop training the UNet") - components.time_entry(frame, 1, 1, self.ui_state, "unet.stop_training_after", "unet.stop_training_after_unit", + components.time_entry(frame, row, 1, self.ui_state, "unet.stop_training_after", "unet.stop_training_after_unit", supports_time_units=False) + row += 1 # unet learning rate - components.label(frame, 2, 0, "UNet Learning Rate", + components.label(frame, row, 0, "UNet Learning Rate", tooltip="The learning rate of the UNet. Overrides the base learning rate") - components.entry(frame, 2, 1, self.ui_state, "unet.learning_rate") + components.entry(frame, row, 1, self.ui_state, "unet.learning_rate") + row += 1 # rescale noise scheduler to zero terminal SNR - rescale_label = components.label(frame, 3, 0, "Rescale Noise Scheduler + V-pred", + rescale_label = components.label(frame, row, 0, "Rescale Noise Scheduler + V-pred", tooltip="Rescales the noise scheduler to a zero terminal signal to noise ratio and switches the model to a v-prediction target") rescale_label.configure(wraplength=130, justify="left") - components.switch(frame, 3, 1, self.ui_state, "rescale_noise_scheduler_to_zero_terminal_snr") + components.switch(frame, row, 1, self.ui_state, "rescale_noise_scheduler_to_zero_terminal_snr") + row += 1 + + def __create_vae_frame(self, master, row): + frame = ctk.CTkFrame(master=master, corner_radius=5) + frame.grid(row=row, column=0, padx=5, pady=5, sticky="nsew") + frame.grid_columnconfigure(0, weight=1) + row = 0 + + components.label(frame, row, 0, "Train VAE", + tooltip="Enables training the VAE model") + components.switch(frame, row, 1, self.ui_state, "vae.train") + row += 1 + + components.label(frame, row, 0, "Gradient Checkpointing", + tooltip="Enables gradient checkpointing for the VAE. Reduces VRAM usage at the cost of training speed") + components.switch(frame, row, 1, self.ui_state, "vae.gradient_checkpointing") + row += 1 def __create_prior_frame(self, master, row): frame = ctk.CTkFrame(master=master, corner_radius=5) frame.grid(row=row, column=0, padx=5, pady=5, sticky="nsew") frame.grid_columnconfigure(0, weight=1) + row = 0 # train prior - components.label(frame, 0, 0, "Train Prior", + components.label(frame, row, 0, "Train Prior", tooltip="Enables training the Prior model") - components.switch(frame, 0, 1, self.ui_state, "prior.train") + components.switch(frame, row, 1, self.ui_state, "prior.train") + row += 1 + + row = self.__create_offloading_widgets(frame, row, "prior", supports_activation_offloading=True) # train prior epochs - components.label(frame, 1, 0, "Stop Training After", + components.label(frame, row, 0, "Stop Training After", tooltip="When to stop training the Prior") - components.time_entry(frame, 1, 1, self.ui_state, "prior.stop_training_after", "prior.stop_training_after_unit", + components.time_entry(frame, row, 1, self.ui_state, "prior.stop_training_after", "prior.stop_training_after_unit", supports_time_units=False) + row += 1 # prior learning rate - components.label(frame, 2, 0, "Prior Learning Rate", + components.label(frame, row, 0, "Prior Learning Rate", tooltip="The learning rate of the Prior. Overrides the base learning rate") - components.entry(frame, 2, 1, self.ui_state, "prior.learning_rate") + components.entry(frame, row, 1, self.ui_state, "prior.learning_rate") + row += 1 def __create_transformer_frame(self, master, row, supports_guidance_scale: bool = False, supports_force_attention_mask: bool = True): frame = ctk.CTkFrame(master=master, corner_radius=5) frame.grid(row=row, column=0, padx=5, pady=5, sticky="nsew") frame.grid_columnconfigure(0, weight=1) + row = 0 # train transformer - components.label(frame, 0, 0, "Train Transformer", + components.label(frame, row, 0, "Train Transformer", tooltip="Enables training the Transformer model") - components.switch(frame, 0, 1, self.ui_state, "transformer.train") + components.switch(frame, row, 1, self.ui_state, "transformer.train") + row += 1 + + row = self.__create_offloading_widgets(frame, row, "transformer", supports_activation_offloading=True) # train transformer epochs - components.label(frame, 1, 0, "Stop Training After", + components.label(frame, row, 0, "Stop Training After", tooltip="When to stop training the Transformer") - components.time_entry(frame, 1, 1, self.ui_state, "transformer.stop_training_after", "transformer.stop_training_after_unit", + components.time_entry(frame, row, 1, self.ui_state, "transformer.stop_training_after", "transformer.stop_training_after_unit", supports_time_units=False) + row += 1 # transformer learning rate - components.label(frame, 2, 0, "Transformer Learning Rate", + components.label(frame, row, 0, "Transformer Learning Rate", tooltip="The learning rate of the Transformer. Overrides the base learning rate") - components.entry(frame, 2, 1, self.ui_state, "transformer.learning_rate") + components.entry(frame, row, 1, self.ui_state, "transformer.learning_rate") + row += 1 if supports_force_attention_mask: # transformer learning rate - components.label(frame, 3, 0, "Force Attention Mask", + components.label(frame, row, 0, "Force Attention Mask", tooltip="Force enables passing of a text embedding attention mask to the transformer. This can improve training on shorter captions.") - components.switch(frame, 3, 1, self.ui_state, "transformer.attention_mask") + components.switch(frame, row, 1, self.ui_state, "transformer.attention_mask") + row += 1 if supports_guidance_scale: # guidance scale - components.label(frame, 4, 0, "Guidance Scale", + components.label(frame, row, 0, "Guidance Scale", tooltip="The guidance scale of guidance distilled models passed to the transformer during training.") - components.entry(frame, 4, 1, self.ui_state, "transformer.guidance_scale") + components.entry(frame, row, 1, self.ui_state, "transformer.guidance_scale") + row += 1 def __create_noise_frame(self, master, row, supports_generalized_offset_noise: bool = False, supports_dynamic_timestep_shifting: bool = False): frame = ctk.CTkFrame(master=master, corner_radius=5) @@ -844,10 +897,6 @@ def __open_timestep_distribution_window(self): window = TimestepDistributionWindow(self.master, self.train_config, self.ui_state) self.master.wait_window(window) - def __open_offloading_window(self): - window = OffloadingWindow(self.master, self.train_config, self.ui_state) - self.master.wait_window(window) - def __restore_optimizer_config(self, *args): optimizer_config = change_optimizer(self.train_config) self.ui_state.get_var("optimizer").update(optimizer_config) diff --git a/modules/util/LayerOffloadConductor.py b/modules/util/LayerOffloadConductor.py index 0b7e6c7ca..062d1391e 100644 --- a/modules/util/LayerOffloadConductor.py +++ b/modules/util/LayerOffloadConductor.py @@ -2,7 +2,7 @@ import random from typing import Any -from modules.util.config.TrainConfig import TrainConfig +from modules.util.config.TrainConfig import TrainConfig, TrainModelPartConfig from modules.util.quantization_util import get_offload_tensor_bytes, offload_quantized from modules.util.torch_util import ( create_stream_context, @@ -566,6 +566,7 @@ def __init__( self, module: nn.Module, config: TrainConfig, + part: TrainModelPartConfig, ): super().__init__() @@ -573,16 +574,16 @@ def __init__( self.__layers = [] self.__layer_device_map = [] - self.__layer_offload_fraction = config.layer_offload_fraction + self.__layer_offload_fraction = part.offload_fraction self.__layer_activations_included_offload_param_indices_map = [] self.__train_device = torch.device(config.train_device) self.__temp_device = torch.device(config.temp_device) - self.__offload_activations = config.gradient_checkpointing.offload() and config.enable_activation_offloading - self.__offload_layers = config.gradient_checkpointing.offload() and config.layer_offload_fraction > 0 - self.__async_transfer = self.__train_device.type == "cuda" and config.enable_async_offloading + self.__offload_activations = part.activation_offloading + self.__offload_layers = part.offload_fraction > 0 + self.__async_transfer = self.__train_device.type == "cuda" and config.async_offloading if self.__async_transfer: self.__train_stream = torch.cuda.default_stream(self.__train_device) diff --git a/modules/util/checkpointing_util.py b/modules/util/checkpointing_util.py index 4e6ee5529..47c84b492 100644 --- a/modules/util/checkpointing_util.py +++ b/modules/util/checkpointing_util.py @@ -3,7 +3,7 @@ from typing import Any from modules.util.compile_util import init_compile -from modules.util.config.TrainConfig import TrainConfig +from modules.util.config.TrainConfig import TrainConfig, TrainModelPartConfig from modules.util.LayerOffloadConductor import LayerOffloadConductor from modules.util.torch_util import add_dummy_grad_fn_, has_grad_fn @@ -69,21 +69,25 @@ def __init__(self, *args, **kwargs): class CheckpointLayer(BaseCheckpointLayer): - def __init__(self, orig_module: nn.Module, orig_forward, train_device: torch.device): + def __init__(self, orig_module: nn.Module, orig_forward, train_device: torch.device, checkpointing: bool = True): super().__init__() assert (orig_module is None or orig_forward is None) and not (orig_module is None and orig_forward is None) self.checkpoint = orig_module self.orig_forward = orig_forward + self.checkpointing = checkpointing # dummy tensor that requires grad is needed for checkpointing to work when training a LoRA self.dummy = torch.zeros((1,), device=train_device, requires_grad=True) - def __checkpointing_forward(self, dummy: torch.Tensor, *args, **kwargs): + def __orig(self, *args, **kwargs): return self.orig_forward(*args, **kwargs) if self.checkpoint is None else self.checkpoint(*args, **kwargs) + def __checkpointing_forward(self, dummy: torch.Tensor, *args, **kwargs): + return self.__orig(*args, **kwargs) + def forward(self, *args, **kwargs): - if torch.is_grad_enabled(): + if self.checkpointing and torch.is_grad_enabled(): return torch.utils.checkpoint.checkpoint( self.__checkpointing_forward, self.dummy, @@ -92,7 +96,7 @@ def forward(self, *args, **kwargs): use_reentrant=False ) else: - return self.orig_forward(*args, **kwargs) if self.checkpoint is None else self.checkpoint(*args, **kwargs) + return self.__orig(*args, **kwargs) class OffloadCheckpointLayer(BaseCheckpointLayer): def __init__(self, orig_module: nn.Module, orig_forward, train_device: torch.device, conductor: LayerOffloadConductor, layer_index: int): @@ -149,6 +153,7 @@ def create_checkpoint( train_device: torch.device, include_from_offload_param_names: list[str] = None, conductor: LayerOffloadConductor | None = None, + checkpointing: bool = True, layer_index: int = 0, compile: bool = False, ) -> Callable: @@ -160,6 +165,9 @@ def create_checkpoint( conductor.add_layer(orig_module, included_offload_param_indices) if conductor is not None and conductor.offload_activated(): + # offloading is structurally coupled to use_reentrant=True checkpointing during the back pass + # (the recompute is what fires before_layer/after_layer in the backward direction), so the offload + # layer always checkpoints when grad is enabled, regardless of the part's gradient_checkpointing flag. if compile: layer = OffloadCheckpointLayer(orig_module=orig_module, orig_forward=None, train_device=train_device, conductor=conductor, layer_index=layer_index) #don't compile the checkpointing layer - offloading cannot be compiled: @@ -172,12 +180,12 @@ def create_checkpoint( return orig_module else: if compile: - layer = CheckpointLayer(orig_module=orig_module, orig_forward=None, train_device=train_device) + layer = CheckpointLayer(orig_module=orig_module, orig_forward=None, train_device=train_device, checkpointing=checkpointing) #do compile the checkpointing layer - slightly faster layer.compile(fullgraph=True) return layer else: - layer = CheckpointLayer(orig_module=None, orig_forward=orig_module.forward, train_device=train_device) + layer = CheckpointLayer(orig_module=None, orig_forward=orig_module.forward, train_device=train_device, checkpointing=checkpointing) orig_module.forward = layer.forward return orig_module @@ -185,6 +193,7 @@ def _create_checkpoints_for_module_list( module_list: nn.ModuleList, include_from_offload_param_names: list[str], conductor: LayerOffloadConductor, + checkpointing: bool, train_device: torch.device, layer_index: int, compile: bool, @@ -196,7 +205,7 @@ def _create_checkpoints_for_module_list( module_list[i] = create_checkpoint( layer, train_device, include_from_offload_param_names, - conductor, layer_index, compile=compile, + conductor, checkpointing, layer_index, compile=compile, ) layer_index += 1 return layer_index @@ -209,11 +218,18 @@ def _remove_checkpoint_keys(module, state_dict, prefix, local_metadata): def enable_checkpointing( model: nn.Module, config: TrainConfig, + 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, -) -> LayerOffloadConductor: - conductor = LayerOffloadConductor(model, config) +) -> 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() + conductor = LayerOffloadConductor(model, config, part) if offload else None + checkpointing = part.checkpointing_enabled() layer_index = 0 for type_or_list, param_names in lists: @@ -224,7 +240,8 @@ def enable_checkpointing( layer_index = _create_checkpoints_for_module_list( module_list, param_names, - conductor if offload_enabled else None, + conductor, + checkpointing, torch.device(config.train_device), layer_index, compile = compile, @@ -238,7 +255,8 @@ def enable_checkpointing( layer_index = _create_checkpoints_for_module_list( module_list, param_names, - conductor if offload_enabled else None, + conductor, + checkpointing, torch.device(config.train_device), layer_index, compile = compile, @@ -249,9 +267,10 @@ def enable_checkpointing( def enable_checkpointing_for_basic_transformer_blocks( model: nn.Module, config: TrainConfig, + part: TrainModelPartConfig, offload_enabled: bool, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, config.compile, [ +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, config.compile, [ (BasicTransformerBlock , []), ], offload_enabled = offload_enabled, @@ -260,16 +279,18 @@ def enable_checkpointing_for_basic_transformer_blocks( def enable_checkpointing_for_clip_encoder_layers( model: nn.Module, config: TrainConfig, + part: TrainModelPartConfig, ): - return enable_checkpointing(model, config, False, [ + 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 ]) def enable_checkpointing_for_t5_encoder_layers( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, False, [ + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, False, [ (T5Block, []), ]) @@ -277,8 +298,9 @@ def enable_checkpointing_for_t5_encoder_layers( def enable_checkpointing_for_gemma_layers( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, False, [ + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, False, [ (Gemma2DecoderLayer, []), ]) @@ -286,17 +308,19 @@ def enable_checkpointing_for_gemma_layers( def enable_checkpointing_for_llama_encoder_layers( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, False, [ + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, False, [ (LlamaDecoderLayer, []), ]) def enable_checkpointing_for_mistral_encoder_layers( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, False, [ - (MistralDecoderLayer, []), + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, False, [ + (MistralDecoderLayer, []), # no activation offloading: this encoder is never trained ]) @@ -304,32 +328,36 @@ def enable_checkpointing_for_mistral_encoder_layers( def enable_checkpointing_for_qwen25vl_encoder_layers( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, False, [ + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, False, [ (Qwen2_5_VLDecoderLayer, []), # TODO No activation offloading for other encoders, see above. But clip skip is not implemented for QwenVL. Then do activation offloading? ]) def enable_checkpointing_for_qwen3_encoder_layers( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, False, [ - (Qwen3DecoderLayer, []), # No activation offloading, because hidden states are taken from the middle of the network by Flux2 + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, False, [ + (Qwen3DecoderLayer, []), # no activation offloading: this encoder is never trained ]) def enable_checkpointing_for_stable_diffusion_3_transformer( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, config.compile, [ + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, config.compile, [ (JointTransformerBlock, ["hidden_states", "encoder_hidden_states"]), ]) def enable_checkpointing_for_flux_transformer( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, config.compile, [ + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, config.compile, [ (model.transformer_blocks, ["hidden_states", "encoder_hidden_states"]), (model.single_transformer_blocks, ["hidden_states" ]), ]) @@ -337,8 +365,9 @@ def enable_checkpointing_for_flux_transformer( def enable_checkpointing_for_flux2_transformer( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, config.compile, [ + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, config.compile, [ (model.transformer_blocks, ["hidden_states", "encoder_hidden_states"]), (model.single_transformer_blocks, ["hidden_states" ]), ]) @@ -347,8 +376,9 @@ def enable_checkpointing_for_flux2_transformer( def enable_checkpointing_for_chroma_transformer( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, config.compile, [ + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, config.compile, [ (model.transformer_blocks, ["hidden_states", "encoder_hidden_states"]), (model.single_transformer_blocks, ["hidden_states" ]), ]) @@ -357,16 +387,18 @@ def enable_checkpointing_for_chroma_transformer( def enable_checkpointing_for_qwen_transformer( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, config.compile, [ + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, config.compile, [ (model.transformer_blocks, ["hidden_states", "encoder_hidden_states"]), ]) def enable_checkpointing_for_z_image_transformer( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, config.compile, [ + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, config.compile, [ (model.noise_refiner, ["x"]), (model.context_refiner, ["x"]), (model.layers, ["x"]), @@ -376,16 +408,18 @@ def enable_checkpointing_for_z_image_transformer( def enable_checkpointing_for_sana_transformer( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, config.compile, [ + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, config.compile, [ (SanaTransformerBlock, ["hidden_states"]), ]) def enable_checkpointing_for_hunyuan_video_transformer( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, config.compile, [ + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, config.compile, [ (HunyuanVideoIndividualTokenRefinerBlock, ["hidden_states" ]), (HunyuanVideoTransformerBlock, ["hidden_states", "encoder_hidden_states"]), (HunyuanVideoSingleTransformerBlock, ["hidden_states" ]), @@ -394,8 +428,9 @@ def enable_checkpointing_for_hunyuan_video_transformer( def enable_checkpointing_for_hi_dream_transformer( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, config.compile, [ + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, config.compile, [ (model.double_stream_blocks, ["hidden_states", "encoder_hidden_states"]), (model.single_stream_blocks, ["hidden_states" ]), ]) @@ -403,7 +438,8 @@ def enable_checkpointing_for_hi_dream_transformer( def enable_checkpointing_for_ernie_transformer( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, config.compile, [ + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, config.compile, [ (model.layers, ["x"]), ]) diff --git a/modules/util/config/TrainConfig.py b/modules/util/config/TrainConfig.py index f52988502..92528a1c2 100644 --- a/modules/util/config/TrainConfig.py +++ b/modules/util/config/TrainConfig.py @@ -13,7 +13,6 @@ from modules.util.enum.ConfigPart import ConfigPart from modules.util.enum.DataType import DataType from modules.util.enum.EMAMode import EMAMode -from modules.util.enum.GradientCheckpointingMethod import GradientCheckpointingMethod from modules.util.enum.GradientReducePrecision import GradientReducePrecision from modules.util.enum.ImageFormat import ImageFormat from modules.util.enum.LearningRateScaler import LearningRateScaler @@ -267,10 +266,27 @@ class TrainModelPartConfig(BaseConfig): train_embedding: bool attention_mask: bool guidance_scale: float + gradient_checkpointing: bool + offload_fraction: float + activation_offloading: bool def __init__(self, data: list[(str, Any, type, bool)]): super().__init__(data) + def offloading_enabled(self) -> bool: + # a conductor should exist iff this is True. Layer offloading applies even to frozen parts (to fit + # them in VRAM), but activation offloading only does work during a backward pass, so it only applies + # when the part is trained -- even if activation_offloading is True in the config. + return self.offload_fraction > 0 or (self.activation_offloading and self.train) + + def checkpointing_enabled(self) -> bool: + # the inner torch checkpoint() should run iff this is True + return self.gradient_checkpointing and self.train + + def checkpointing_or_offloading_enabled(self) -> bool: + # whether the checkpoint layer wrapper needs to be installed for this part at all + return self.checkpointing_enabled() or self.offloading_enabled() + @staticmethod def default_values(): data = [] @@ -287,6 +303,9 @@ def default_values(): data.append(("train_embedding", True, bool, False)) data.append(("attention_mask", False, bool, False)) data.append(("guidance_scale", 1.0, float, False)) + data.append(("gradient_checkpointing", True, bool, False)) + data.append(("offload_fraction", 0.0, float, False)) + data.append(("activation_offloading", True, bool, False)) return TrainModelPartConfig(data) @@ -374,10 +393,7 @@ class TrainConfig(BaseConfig): output_dtype: DataType output_model_format: ModelFormat output_model_destination: str - gradient_checkpointing: GradientCheckpointingMethod - enable_async_offloading: bool - enable_activation_offloading: bool - layer_offload_fraction: float + async_offloading: bool force_circular_padding: bool compile: bool @@ -569,7 +585,7 @@ class TrainConfig(BaseConfig): def __init__(self, data: list[(str, Any, type, bool)]): super().__init__( data, - config_version=10, + config_version=11, config_migrations={ 0: self.__migration_0, 1: self.__migration_1, @@ -581,6 +597,7 @@ def __init__(self, data: list[(str, Any, type, bool)]): 7: self.__migration_7, 8: self.__migration_8, 9: self.__migration_9, + 10: self.__migration_10, } ) @@ -727,12 +744,14 @@ def __migration_3(self, data: dict) -> dict: def __migration_4(self, data: dict) -> dict: migrated_data = data.copy() + # Translate the old bool form of gradient_checkpointing into the v5..v10 + # string/enum form. __migration_10 later fans this out per-component. gradient_checkpointing = migrated_data.pop("gradient_checkpointing", True) if gradient_checkpointing: - migrated_data["gradient_checkpointing"] = GradientCheckpointingMethod.ON + migrated_data["gradient_checkpointing"] = "ON" else: - migrated_data["gradient_checkpointing"] = GradientCheckpointingMethod.OFF + migrated_data["gradient_checkpointing"] = "OFF" return migrated_data @@ -800,6 +819,42 @@ def replace_dtype(part: str): return migrated_data + def __migration_10(self, data: dict) -> dict: + migrated_data = data.copy() + + # Fan the four old global offload/checkpointing settings out per-component. + # After __migration_4 gradient_checkpointing is a string "OFF"/"ON"/"CPU_OFFLOADED". + gc = migrated_data.pop("gradient_checkpointing", "ON") + act = migrated_data.pop("enable_activation_offloading", True) + frac = migrated_data.pop("layer_offload_fraction", 0.0) + migrated_data["async_offloading"] = migrated_data.pop("enable_async_offloading", True) + + def fan_out(part: str): + if part in migrated_data: + migrated_data[part]["gradient_checkpointing"] = gc != "OFF" + migrated_data[part]["activation_offloading"] = (gc == "CPU_OFFLOADED") and act + migrated_data[part]["offload_fraction"] = frac if gc == "CPU_OFFLOADED" else 0.0 + + fan_out("unet") + fan_out("prior") + fan_out("transformer") + fan_out("text_encoder") + fan_out("text_encoder_2") + fan_out("text_encoder_3") + fan_out("text_encoder_4") + fan_out("vae") + fan_out("effnet_encoder") + fan_out("decoder") + fan_out("decoder_text_encoder") + fan_out("decoder_vqgan") + + return migrated_data + + def model_part_configs(self) -> list[TrainModelPartConfig]: + # the per-part configs for the components this model_type actually has. Avoids "phantom" parts whose + # fields keep their defaults (train=True) or migrated offload values but don't exist in the model. + return [getattr(self, name) for name in self.model_type.model_parts()] + def weight_dtypes(self) -> ModelWeightDtypes: return ModelWeightDtypes( self.train_dtype, @@ -969,10 +1024,7 @@ def default_values() -> 'TrainConfig': data.append(("output_dtype", DataType.FLOAT_32, DataType, False)) data.append(("output_model_format", ModelFormat.SAFETENSORS, ModelFormat, False)) data.append(("output_model_destination", "models/model.safetensors", str, False)) - data.append(("gradient_checkpointing", GradientCheckpointingMethod.ON, GradientCheckpointingMethod, False)) - data.append(("enable_async_offloading", True, bool, False)) - data.append(("enable_activation_offloading", True, bool, False)) - data.append(("layer_offload_fraction", 0.0, float, False)) + data.append(("async_offloading", True, bool, False)) data.append(("force_circular_padding", False, bool, False)) data.append(("compile", False, bool, False)) diff --git a/modules/util/create.py b/modules/util/create.py index 7c0194da8..0c0adba16 100644 --- a/modules/util/create.py +++ b/modules/util/create.py @@ -110,7 +110,11 @@ def create_data_loader( train_progress: TrainProgress | None = None, is_validation: bool = False ) -> BaseDataLoader | None: - if config.gradient_checkpointing.offload() and config.layer_offload_fraction > 0 and config.dataloader_threads > 1: + # Layer offloading uses a non-thread-safe conductor. This check is too broad: it trips whenever any model + # part does layer offloading, even though only a component that is actually cached really runs in the + # dataloader worker threads. + # TODO: narrow this to the cached components only. + if config.dataloader_threads > 1 and any(part.offload_fraction > 0 for part in config.model_part_configs()): raise RuntimeError('layer offloading can not be activated if "dataloader_threads" > 1') if train_progress is None: @@ -133,7 +137,8 @@ def create_optimizer( if optimizer_config.optimizer is None: return None - if config.gradient_checkpointing.offload() and config.layer_offload_fraction > 0: + # a trained, layer-offloaded part has its params evicted during the back pass, so it needs fused_back_pass + if any(part.offload_fraction > 0 and part.train for part in config.model_part_configs()): if (not optimizer_config.optimizer.supports_fused_back_pass() or not optimizer_config.fused_back_pass) \ and config.training_method == TrainingMethod.FINE_TUNE: raise RuntimeError('layer offloading can only be used for fine tuning when using an optimizer that supports "fused_back_pass"') diff --git a/modules/util/enum/GradientCheckpointingMethod.py b/modules/util/enum/GradientCheckpointingMethod.py deleted file mode 100644 index d3f05666a..000000000 --- a/modules/util/enum/GradientCheckpointingMethod.py +++ /dev/null @@ -1,17 +0,0 @@ -from enum import Enum - - -class GradientCheckpointingMethod(Enum): - OFF = 'OFF' - ON = 'ON' - CPU_OFFLOADED = 'CPU_OFFLOADED' - - def __str__(self): - return self.value - - def enabled(self): - return self == GradientCheckpointingMethod.ON \ - or self == GradientCheckpointingMethod.CPU_OFFLOADED - - def offload(self): - return self == GradientCheckpointingMethod.CPU_OFFLOADED diff --git a/training_presets/#chroma Finetune 16GB.json b/training_presets/#chroma Finetune 16GB.json index 2dacbee20..145ab5b04 100644 --- a/training_presets/#chroma Finetune 16GB.json +++ b/training_presets/#chroma Finetune 16GB.json @@ -4,16 +4,15 @@ "learning_rate": 1e-5, "model_type": "CHROMA_1", "resolution": "512", - "gradient_checkpointing": "CPU_OFFLOADED", - "layer_offload_fraction": 0.4, "dataloader_threads": 1, "transformer": { "train": true, - "weight_dtype": "BFLOAT_16" + "weight_dtype": "BFLOAT_16", + "offload_fraction": 0.4 }, "text_encoder": { "train": false, - "weight_dtype": "BFLOAT_16" + "weight_dtype": "FLOAT_8" }, "training_method": "FINE_TUNE", "vae": { diff --git a/training_presets/#chroma Finetune 8GB.json b/training_presets/#chroma Finetune 8GB.json index 508410995..29b36c84b 100644 --- a/training_presets/#chroma Finetune 8GB.json +++ b/training_presets/#chroma Finetune 8GB.json @@ -4,16 +4,15 @@ "learning_rate": 1e-5, "model_type": "CHROMA_1", "resolution": "512", - "gradient_checkpointing": "CPU_OFFLOADED", - "layer_offload_fraction": 0.85, "dataloader_threads": 1, "transformer": { "train": true, - "weight_dtype": "BFLOAT_16" + "weight_dtype": "BFLOAT_16", + "offload_fraction": 0.85 }, "text_encoder": { "train": false, - "weight_dtype": "BFLOAT_16" + "weight_dtype": "FLOAT_8" }, "training_method": "FINE_TUNE", "vae": { diff --git a/training_presets/#chroma LoRA 8GB.json b/training_presets/#chroma LoRA 8GB.json index 78027aac4..437ac71d7 100644 --- a/training_presets/#chroma LoRA 8GB.json +++ b/training_presets/#chroma LoRA 8GB.json @@ -4,16 +4,15 @@ "learning_rate": 0.0003, "model_type": "CHROMA_1", "resolution": "512", - "gradient_checkpointing": "CPU_OFFLOADED", - "layer_offload_fraction": 0.6, "dataloader_threads": 1, "transformer": { "train": true, - "weight_dtype": "FLOAT_8" + "weight_dtype": "FLOAT_8", + "offload_fraction": 0.6 }, "text_encoder": { "train": false, - "weight_dtype": "BFLOAT_16" + "weight_dtype": "FLOAT_8" }, "training_method": "LORA", "vae": { diff --git a/training_presets/#ernie LoRA 8GB.json b/training_presets/#ernie LoRA 8GB.json index bc1e4a82b..a63d34ece 100644 --- a/training_presets/#ernie LoRA 8GB.json +++ b/training_presets/#ernie LoRA 8GB.json @@ -7,7 +7,8 @@ "compile": true, "transformer": { "train": true, - "weight_dtype": "INT_W8A8" + "weight_dtype": "INT_W8A8", + "offload_fraction": 0.7 }, "text_encoder": { "train": false, @@ -26,7 +27,5 @@ "layer_filter_preset": "blocks" }, "timestep_distribution": "LOGIT_NORMAL", - "dataloader_threads": 1, - "gradient_checkpointing": "CPU_OFFLOADED", - "layer_offload_fraction": 0.7 + "dataloader_threads": 1 } diff --git a/training_presets/#flux2 Finetune 16GB.json b/training_presets/#flux2 Finetune 16GB.json index ac07a501e..00094da12 100644 --- a/training_presets/#flux2 Finetune 16GB.json +++ b/training_presets/#flux2 Finetune 16GB.json @@ -7,11 +7,12 @@ "compile": true, "transformer": { "train": true, - "weight_dtype": "BFLOAT_16" + "weight_dtype": "BFLOAT_16", + "offload_fraction": 0.6 }, "text_encoder": { "train": false, - "weight_dtype": "BFLOAT_16" + "weight_dtype": "FLOAT_8" }, "training_method": "FINE_TUNE", "vae": { @@ -26,8 +27,6 @@ "timestep_distribution": "LOGIT_NORMAL", "dynamic_timestep_shifting": true, "dataloader_threads": 1, - "gradient_checkpointing": "CPU_OFFLOADED", - "layer_offload_fraction": 0.6, "optimizer": { "optimizer": "ADAFACTOR" }, diff --git a/training_presets/#flux2 LoRA 8GB.json b/training_presets/#flux2 LoRA 8GB.json index 0bd10d116..160a999ef 100644 --- a/training_presets/#flux2 LoRA 8GB.json +++ b/training_presets/#flux2 LoRA 8GB.json @@ -7,11 +7,13 @@ "compile": true, "transformer": { "train": true, - "weight_dtype": "INT_W8A8" + "weight_dtype": "INT_W8A8", + "offload_fraction": 0.7 }, "text_encoder": { "train": false, - "weight_dtype": "FLOAT_8" + "weight_dtype": "FLOAT_8", + "offload_fraction": 0.7 }, "training_method": "LORA", "vae": { @@ -27,7 +29,5 @@ }, "timestep_distribution": "LOGIT_NORMAL", "dynamic_timestep_shifting": true, - "dataloader_threads": 1, - "gradient_checkpointing": "CPU_OFFLOADED", - "layer_offload_fraction": 0.7 + "dataloader_threads": 1 } diff --git a/training_presets/#hidream LoRA.json b/training_presets/#hidream LoRA.json index 2eb588763..8b025a628 100644 --- a/training_presets/#hidream LoRA.json +++ b/training_presets/#hidream LoRA.json @@ -2,8 +2,6 @@ "backup_after": 10, "base_model_name": "HiDream-ai/HiDream-I1-Full", "batch_size": 4, - "gradient_checkpointing": "CPU_OFFLOADED", - "layer_offload_fraction": 0.5, "dataloader_threads": 1, "learning_rate": 0.0003, "model_type": "HI_DREAM_FULL", @@ -16,7 +14,8 @@ "training_method": "LORA", "transformer": { "train": true, - "weight_dtype": "FLOAT_8" + "weight_dtype": "FLOAT_8", + "offload_fraction": 0.5 }, "text_encoder": { "train": false, @@ -28,11 +27,13 @@ }, "text_encoder_3": { "train": false, - "weight_dtype": "FLOAT_8" + "weight_dtype": "FLOAT_8", + "offload_fraction": 0.5 }, "text_encoder_4": { "model_name": "meta-llama/Llama-3.1-8B-Instruct", "train": false, - "weight_dtype": "FLOAT_8" + "weight_dtype": "FLOAT_8", + "offload_fraction": 0.5 } } diff --git a/training_presets/#hunyuan video LoRA.json b/training_presets/#hunyuan video LoRA.json index 754550155..96bbb80e7 100644 --- a/training_presets/#hunyuan video LoRA.json +++ b/training_presets/#hunyuan video LoRA.json @@ -2,8 +2,6 @@ "backup_after": 10, "base_model_name": "hunyuanvideo-community/HunyuanVideo", "batch_size": 4, - "gradient_checkpointing": "CPU_OFFLOADED", - "layer_offload_fraction": 0.5, "dataloader_threads": 1, "learning_rate": 0.0003, "model_type": "HUNYUAN_VIDEO", @@ -16,11 +14,13 @@ "training_method": "LORA", "transformer": { "train": true, - "weight_dtype": "FLOAT_8" + "weight_dtype": "FLOAT_8", + "offload_fraction": 0.5 }, "text_encoder": { "train": false, - "weight_dtype": "FLOAT_8" + "weight_dtype": "FLOAT_8", + "offload_fraction": 0.5 }, "text_encoder_2": { "train": false, diff --git a/training_presets/#qwen Finetune 16GB.json b/training_presets/#qwen Finetune 16GB.json index 811d7e0b1..2224993f4 100644 --- a/training_presets/#qwen Finetune 16GB.json +++ b/training_presets/#qwen Finetune 16GB.json @@ -4,12 +4,11 @@ "learning_rate": 1e-5, "model_type": "QWEN", "resolution": "512", - "gradient_checkpointing": "CPU_OFFLOADED", - "layer_offload_fraction": 0.75, "dataloader_threads": 1, "transformer": { "train": true, - "weight_dtype": "BFLOAT_16" + "weight_dtype": "BFLOAT_16", + "offload_fraction": 0.75 }, "text_encoder": { "train": false, diff --git a/training_presets/#qwen Finetune 24GB.json b/training_presets/#qwen Finetune 24GB.json index 8bee3cd3f..1cf9dc09f 100644 --- a/training_presets/#qwen Finetune 24GB.json +++ b/training_presets/#qwen Finetune 24GB.json @@ -4,12 +4,11 @@ "learning_rate": 1e-5, "model_type": "QWEN", "resolution": "512", - "gradient_checkpointing": "CPU_OFFLOADED", - "layer_offload_fraction": 0.55, "dataloader_threads": 1, "transformer": { "train": true, - "weight_dtype": "BFLOAT_16" + "weight_dtype": "BFLOAT_16", + "offload_fraction": 0.55 }, "text_encoder": { "train": false, diff --git a/training_presets/#qwen LoRA 16GB.json b/training_presets/#qwen LoRA 16GB.json index b4e0d7e88..0eda34d5d 100644 --- a/training_presets/#qwen LoRA 16GB.json +++ b/training_presets/#qwen LoRA 16GB.json @@ -4,12 +4,11 @@ "learning_rate": 0.0003, "model_type": "QWEN", "resolution": "512", - "gradient_checkpointing": "CPU_OFFLOADED", - "layer_offload_fraction": 0.5, "dataloader_threads": 1, "transformer": { "train": true, - "weight_dtype": "FLOAT_8" + "weight_dtype": "FLOAT_8", + "offload_fraction": 0.5 }, "text_encoder": { "train": false, diff --git a/training_presets/#qwen LoRA 24GB.json b/training_presets/#qwen LoRA 24GB.json index 696648a42..cd7b7216e 100644 --- a/training_presets/#qwen LoRA 24GB.json +++ b/training_presets/#qwen LoRA 24GB.json @@ -4,12 +4,11 @@ "learning_rate": 0.0003, "model_type": "QWEN", "resolution": "512", - "gradient_checkpointing": "CPU_OFFLOADED", - "layer_offload_fraction": 0.1, "dataloader_threads": 1, "transformer": { "train": true, - "weight_dtype": "FLOAT_8" + "weight_dtype": "FLOAT_8", + "offload_fraction": 0.1 }, "text_encoder": { "train": false, diff --git a/training_presets/#z-image DeTurbo LoRA 8GB.json b/training_presets/#z-image DeTurbo LoRA 8GB.json index cc38e60eb..957f08b80 100644 --- a/training_presets/#z-image DeTurbo LoRA 8GB.json +++ b/training_presets/#z-image DeTurbo LoRA 8GB.json @@ -4,13 +4,12 @@ "learning_rate": 0.0003, "model_type": "Z_IMAGE", "resolution": "512", - "gradient_checkpointing": "CPU_OFFLOADED", - "layer_offload_fraction": 0.6, "compile": true, "transformer": { "train": true, "weight_dtype": "INT_W8A8", - "model_name": "https://huggingface.co/ostris/Z-Image-De-Turbo/blob/main/z_image_de_turbo_v1_bf16.safetensors" + "model_name": "https://huggingface.co/ostris/Z-Image-De-Turbo/blob/main/z_image_de_turbo_v1_bf16.safetensors", + "offload_fraction": 0.6 }, "text_encoder": { "train": false, diff --git a/training_presets/#z-image Finetune 16GB.json b/training_presets/#z-image Finetune 16GB.json index 0d23d3992..3911c866c 100644 --- a/training_presets/#z-image Finetune 16GB.json +++ b/training_presets/#z-image Finetune 16GB.json @@ -4,12 +4,11 @@ "learning_rate": 1e-5, "model_type": "Z_IMAGE", "resolution": "512", - "gradient_checkpointing": "CPU_OFFLOADED", - "layer_offload_fraction": 0.1, "compile": true, "transformer": { "train": true, - "weight_dtype": "BFLOAT_16" + "weight_dtype": "BFLOAT_16", + "offload_fraction": 0.1 }, "text_encoder": { "train": false, diff --git a/training_presets/#z-image LoRA 8GB.json b/training_presets/#z-image LoRA 8GB.json index 78b4b05cc..8dbce7b21 100644 --- a/training_presets/#z-image LoRA 8GB.json +++ b/training_presets/#z-image LoRA 8GB.json @@ -4,12 +4,11 @@ "learning_rate": 0.0003, "model_type": "Z_IMAGE", "resolution": "512", - "gradient_checkpointing": "CPU_OFFLOADED", - "layer_offload_fraction": 0.6, "compile": true, "transformer": { "train": true, - "weight_dtype": "FLOAT_8" + "weight_dtype": "FLOAT_8", + "offload_fraction": 0.6 }, "text_encoder": { "train": false, From d680a1986c7bd1f511a84a2c3479c9d43d67dc8d Mon Sep 17 00:00:00 2001 From: dxqb <183307934+dxqb@users.noreply.github.com> Date: Wed, 1 Jul 2026 09:13:27 +0200 Subject: [PATCH 3/8] Centralize model composition and training methods in ModelType Adds ModelType.model_parts()/has_multiple_text_encoders()/supported_training_methods() as the single source of truth for which components a model has and which training methods it supports. BaseModelTabView.build_content and TopBarController now derive the Model tab layout and training-method dropdown generically from these, replacing one hand-written per-model method / if-chain each. --- modules/ui/BaseModelTabView.py | 377 ++++----------------------------- modules/ui/TopBarController.py | 38 +--- modules/util/enum/ModelType.py | 57 +++++ 3 files changed, 111 insertions(+), 361 deletions(-) diff --git a/modules/ui/BaseModelTabView.py b/modules/ui/BaseModelTabView.py index 061207f4a..ec9d03d10 100644 --- a/modules/ui/BaseModelTabView.py +++ b/modules/ui/BaseModelTabView.py @@ -19,172 +19,21 @@ def _make_svd_frames(self, parent, row: int): pass def build_content(self, frame, controller, ui_state): - if controller.train_config.model_type.is_stable_diffusion(): # TODO simplify - self.__setup_stable_diffusion_ui(frame, controller, ui_state) - if controller.train_config.model_type.is_stable_diffusion_3(): - self.__setup_stable_diffusion_3_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_stable_diffusion_xl(): - self.__setup_stable_diffusion_xl_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_wuerstchen(): - self.__setup_wuerstchen_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_pixart(): - self.__setup_pixart_alpha_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_flux_1(): - self.__setup_flux_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_flux_2(): - self.__setup_flux_2_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_z_image(): - self.__setup_z_image_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_chroma(): - self.__setup_chroma_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_qwen(): - self.__setup_qwen_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_sana(): - self.__setup_sana_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_hunyuan_video(): - self.__setup_hunyuan_video_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_hi_dream(): - self.__setup_hi_dream_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_ernie(): - self.__setup_ernie_ui(frame, controller, ui_state) - - def __setup_stable_diffusion_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_unet=True, - has_text_encoder=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - ui_state, - allow_safetensors=True, - allow_diffusers=controller.train_config.training_method in [ - TrainingMethod.FINE_TUNE, - TrainingMethod.FINE_TUNE_VAE, - ], - allow_legacy_safetensors=controller.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_stable_diffusion_3_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_text_encoder_3=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - ui_state, - allow_safetensors=True, - allow_diffusers=controller.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=controller.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_flux_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - ui_state, - allow_safetensors=True, - allow_diffusers=controller.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=controller.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_flux_2_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - ui_state, - allow_safetensors=True, - allow_diffusers=controller.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=controller.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_z_image_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - ui_state, - allow_safetensors=True, - allow_diffusers=controller.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=controller.train_config.training_method == TrainingMethod.LORA, - ) + model_type = controller.train_config.model_type + training_method = controller.train_config.training_method + parts = model_type.model_parts() - def __setup_ernie_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - ui_state, - allow_safetensors=True, - allow_diffusers=controller.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=controller.train_config.training_method == TrainingMethod.LORA, + # The transformer override path exists only for these architectures; SD3, PixArt, Sana + # and HiDream have a transformer but expose no override field. + allow_override_transformer = ( + model_type.is_flux() + or model_type.is_z_image() + or model_type.is_ernie() + or model_type.is_chroma() + or model_type.is_qwen() + or model_type.is_hunyuan_video() ) - def __setup_chroma_ui(self, frame, controller, ui_state): row = 0 row = self.__create_base_dtype_components(frame, row, ui_state) row = self.__create_base_components( @@ -192,176 +41,44 @@ def __setup_chroma_ui(self, frame, controller, ui_state): row, controller, ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - ui_state, - allow_safetensors=True, - allow_diffusers=controller.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=controller.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_qwen_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - ui_state, - allow_safetensors=True, - allow_diffusers=controller.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=controller.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_stable_diffusion_xl_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_unet=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - ui_state, - allow_safetensors=True, - allow_diffusers=controller.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=controller.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_wuerstchen_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_prior=True, - allow_override_prior=controller.train_config.model_type.is_stable_cascade(), - has_text_encoder=True, - ) - row = self.__create_effnet_encoder_components(frame, row, ui_state) - row = self.__create_decoder_components(frame, row, ui_state, controller.train_config.model_type.is_wuerstchen_v2()) - row = self.__create_output_components( - frame, - row, - ui_state, - allow_safetensors=controller.train_config.training_method != TrainingMethod.FINE_TUNE - or controller.train_config.model_type.is_stable_cascade(), - allow_diffusers=controller.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=controller.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_pixart_alpha_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - has_text_encoder=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - ui_state, - allow_safetensors=True, - allow_diffusers=controller.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=controller.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_sana_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - has_text_encoder=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - ui_state, - allow_safetensors=controller.train_config.training_method != TrainingMethod.FINE_TUNE, - allow_diffusers=controller.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=controller.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_hunyuan_video_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - ui_state, - allow_safetensors=True, - allow_diffusers=controller.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=controller.train_config.training_method == TrainingMethod.LORA, - ) - - def __setup_hi_dream_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_text_encoder_3=True, - has_text_encoder_4=True, - allow_override_text_encoder_4=True, - has_vae=True, - ) - row = self.__create_output_components( + has_unet="unet" in parts, + has_prior="prior" in parts, + allow_override_prior=model_type.is_stable_cascade(), + has_transformer="transformer" in parts, + allow_override_transformer=allow_override_transformer, + has_text_encoder=not model_type.has_multiple_text_encoders(), + has_text_encoder_1=model_type.has_multiple_text_encoders(), + has_text_encoder_2="text_encoder_2" in parts, + has_text_encoder_3="text_encoder_3" in parts, + has_text_encoder_4="text_encoder_4" in parts, + allow_override_text_encoder_4="text_encoder_4" in parts, + has_vae="vae" in parts, + ) + if "effnet_encoder" in parts: + row = self.__create_effnet_encoder_components(frame, row, ui_state) + if "decoder" in parts: + row = self.__create_decoder_components(frame, row, ui_state, "decoder_text_encoder" in parts) + + if model_type.is_sana(): + allow_safetensors = training_method != TrainingMethod.FINE_TUNE + elif model_type.is_wuerstchen(): + allow_safetensors = training_method != TrainingMethod.FINE_TUNE \ + or model_type.is_stable_cascade() + else: + allow_safetensors = True + + if model_type.is_stable_diffusion(): + allow_diffusers = training_method in [TrainingMethod.FINE_TUNE, TrainingMethod.FINE_TUNE_VAE] + else: + allow_diffusers = training_method == TrainingMethod.FINE_TUNE + + self.__create_output_components( frame, row, ui_state, - allow_safetensors=True, - allow_diffusers=controller.train_config.training_method == TrainingMethod.FINE_TUNE, - allow_legacy_safetensors=controller.train_config.training_method == TrainingMethod.LORA, + allow_safetensors=allow_safetensors, + allow_diffusers=allow_diffusers, + allow_legacy_safetensors=training_method == TrainingMethod.LORA, ) def __create_dtype_options(self, include_gguf: bool = False, include_a8: bool = False) -> list[tuple[str, DataType]]: diff --git a/modules/ui/TopBarController.py b/modules/ui/TopBarController.py index 31d3fc1be..693e7718c 100644 --- a/modules/ui/TopBarController.py +++ b/modules/ui/TopBarController.py @@ -40,37 +40,13 @@ def get_model_types(self) -> list[tuple[str, ModelType]]: ] def get_training_methods(self, model_type: ModelType) -> list[tuple[str, TrainingMethod]]: - #TODO simplify - if model_type.is_stable_diffusion(): - return [ - ("Fine Tune", TrainingMethod.FINE_TUNE), - ("LoRA", TrainingMethod.LORA), - ("Embedding", TrainingMethod.EMBEDDING), - ("Fine Tune VAE", TrainingMethod.FINE_TUNE_VAE), - ] - elif model_type.is_stable_diffusion_3() \ - or model_type.is_stable_diffusion_xl() \ - or model_type.is_wuerstchen() \ - or model_type.is_pixart() \ - or model_type.is_flux_1() \ - or model_type.is_sana() \ - or model_type.is_hunyuan_video() \ - or model_type.is_hi_dream() \ - or model_type.is_chroma(): - return [ - ("Fine Tune", TrainingMethod.FINE_TUNE), - ("LoRA", TrainingMethod.LORA), - ("Embedding", TrainingMethod.EMBEDDING), - ] - elif model_type.is_qwen() \ - or model_type.is_z_image() \ - or model_type.is_flux_2() \ - or model_type.is_ernie(): - return [ - ("Fine Tune", TrainingMethod.FINE_TUNE), - ("LoRA", TrainingMethod.LORA), - ] - return [] + labels = { + TrainingMethod.FINE_TUNE: "Fine Tune", + TrainingMethod.LORA: "LoRA", + TrainingMethod.EMBEDDING: "Embedding", + TrainingMethod.FINE_TUNE_VAE: "Fine Tune VAE", + } + return [(labels[m], m) for m in model_type.supported_training_methods()] def load_available_config_names(self, dir: str) -> list[tuple[str, str]]: configs = [("", path_util.canonical_join(dir, "#.json"))] diff --git a/modules/util/enum/ModelType.py b/modules/util/enum/ModelType.py index a3ad940ec..aecb3cf8f 100644 --- a/modules/util/enum/ModelType.py +++ b/modules/util/enum/ModelType.py @@ -1,5 +1,7 @@ from enum import Enum +from modules.util.enum.TrainingMethod import TrainingMethod + class ModelType(Enum): STABLE_DIFFUSION_15 = 'STABLE_DIFFUSION_15' @@ -166,6 +168,61 @@ def is_flow_matching(self) -> bool: def is_video_model(self) -> bool: return self.is_hunyuan_video() #incase we add more video models in the future + def model_parts(self) -> tuple[str, ...]: + return _MODEL_PARTS[self] + + def supported_training_methods(self) -> tuple[TrainingMethod, ...]: + if self.is_stable_diffusion(): + return (TrainingMethod.FINE_TUNE, TrainingMethod.LORA, TrainingMethod.EMBEDDING, TrainingMethod.FINE_TUNE_VAE) + if self.is_stable_diffusion_3() \ + or self.is_stable_diffusion_xl() \ + or self.is_wuerstchen() \ + or self.is_pixart() \ + or self.is_flux_1() \ + or self.is_sana() \ + or self.is_hunyuan_video() \ + or self.is_hi_dream() \ + or self.is_chroma(): + return (TrainingMethod.FINE_TUNE, TrainingMethod.LORA, TrainingMethod.EMBEDDING) + if self.is_qwen() or self.is_z_image() or self.is_flux_2() or self.is_ernie(): + return (TrainingMethod.FINE_TUNE, TrainingMethod.LORA) + raise ValueError(f"No supported training methods defined for model type {self}") + + +# The components each model type has, keyed by TrainConfig field names, as the single source of truth. +# The diffusion model (unet / transformer / prior) is always listed first; the first text encoder is +# "text_encoder" (matching the config field), even for multi-encoder models that refer to it as +# "text_encoder_1" elsewhere in the code. +_MODEL_PARTS: dict[ModelType, tuple[str, ...]] = { + ModelType.STABLE_DIFFUSION_15: ("unet", "text_encoder", "vae"), + ModelType.STABLE_DIFFUSION_15_INPAINTING: ("unet", "text_encoder", "vae"), + ModelType.STABLE_DIFFUSION_20: ("unet", "text_encoder", "vae"), + ModelType.STABLE_DIFFUSION_20_BASE: ("unet", "text_encoder", "vae"), + ModelType.STABLE_DIFFUSION_20_INPAINTING: ("unet", "text_encoder", "vae"), + ModelType.STABLE_DIFFUSION_20_DEPTH: ("unet", "text_encoder", "vae"), + ModelType.STABLE_DIFFUSION_21: ("unet", "text_encoder", "vae"), + ModelType.STABLE_DIFFUSION_21_BASE: ("unet", "text_encoder", "vae"), + ModelType.STABLE_DIFFUSION_3: ("transformer", "text_encoder", "text_encoder_2", "text_encoder_3", "vae"), + ModelType.STABLE_DIFFUSION_35: ("transformer", "text_encoder", "text_encoder_2", "text_encoder_3", "vae"), + ModelType.STABLE_DIFFUSION_XL_10_BASE: ("unet", "text_encoder", "text_encoder_2", "vae"), + ModelType.STABLE_DIFFUSION_XL_10_BASE_INPAINTING: ("unet", "text_encoder", "text_encoder_2", "vae"), + # Only Würstchen v2's decoder has its own text encoder; Stable Cascade's decoder does not. + ModelType.WUERSTCHEN_2: ("prior", "text_encoder", "effnet_encoder", "decoder", "decoder_text_encoder", "decoder_vqgan"), + ModelType.STABLE_CASCADE_1: ("prior", "text_encoder", "effnet_encoder", "decoder", "decoder_vqgan"), + ModelType.PIXART_ALPHA: ("transformer", "text_encoder", "vae"), + ModelType.PIXART_SIGMA: ("transformer", "text_encoder", "vae"), + ModelType.FLUX_DEV_1: ("transformer", "text_encoder", "text_encoder_2", "vae"), + ModelType.FLUX_FILL_DEV_1: ("transformer", "text_encoder", "text_encoder_2", "vae"), + ModelType.FLUX_2: ("transformer", "text_encoder", "vae"), + ModelType.SANA: ("transformer", "text_encoder", "vae"), + ModelType.HUNYUAN_VIDEO: ("transformer", "text_encoder", "text_encoder_2", "vae"), + ModelType.HI_DREAM_FULL: ("transformer", "text_encoder", "text_encoder_2", "text_encoder_3", "text_encoder_4", "vae"), + ModelType.CHROMA_1: ("transformer", "text_encoder", "vae"), + ModelType.QWEN: ("transformer", "text_encoder", "vae"), + ModelType.Z_IMAGE: ("transformer", "text_encoder", "vae"), + ModelType.ERNIE: ("transformer", "text_encoder", "vae"), +} + class PeftType(Enum): LORA = 'LORA' From c9c88e7e6415bd2c5628aa384ae363bd144acef9 Mon Sep 17 00:00:00 2001 From: dxqb <183307934+dxqb@users.noreply.github.com> Date: Sat, 4 Jul 2026 08:07:55 +0200 Subject: [PATCH 4/8] Centralize model composition and training methods in ModelType Add a ModelType.supported_training_methods() that enumerates every model type explicitly, raising on an unknown type rather than defaulting. Collapse ModelTab's per-type __setup_*_ui methods into one build_content that derives the has_* widget flags from ModelType.model_parts(), and collapse TopBar's per-type training-method dispatch to build its dropdown from supported_training_methods(). Co-Authored-By: Claude Sonnet 5 --- modules/ui/BaseModelTabView.py | 360 ++----------------------------- modules/ui/ModelTabController.py | 15 ++ modules/ui/TopBarController.py | 40 +--- modules/util/enum/ModelType.py | 18 ++ 4 files changed, 59 insertions(+), 374 deletions(-) diff --git a/modules/ui/BaseModelTabView.py b/modules/ui/BaseModelTabView.py index 08d873d3b..0757a71c2 100644 --- a/modules/ui/BaseModelTabView.py +++ b/modules/ui/BaseModelTabView.py @@ -17,320 +17,9 @@ def _make_svd_frames(self, parent, row: int): pass def build_content(self, frame, controller, ui_state): - if controller.train_config.model_type.is_stable_diffusion(): # TODO simplify - self.__setup_stable_diffusion_ui(frame, controller, ui_state) - if controller.train_config.model_type.is_stable_diffusion_3(): - self.__setup_stable_diffusion_3_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_stable_diffusion_xl(): - self.__setup_stable_diffusion_xl_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_wuerstchen(): - self.__setup_wuerstchen_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_pixart(): - self.__setup_pixart_alpha_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_flux_1(): - self.__setup_flux_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_flux_2(): - self.__setup_flux_2_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_z_image(): - self.__setup_z_image_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_chroma(): - self.__setup_chroma_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_qwen(): - self.__setup_qwen_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_anima(): - self.__setup_anima_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_krea2(): - self.__setup_krea2_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_sana(): - self.__setup_sana_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_hunyuan_video(): - self.__setup_hunyuan_video_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_hi_dream(): - self.__setup_hi_dream_ui(frame, controller, ui_state) - elif controller.train_config.model_type.is_ernie(): - self.__setup_ernie_ui(frame, controller, ui_state) - - def __setup_stable_diffusion_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_unet=True, - has_text_encoder=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_stable_diffusion_3_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_text_encoder_3=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_flux_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_flux_2_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_z_image_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_ernie_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_chroma_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_qwen_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_anima_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_krea2_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_stable_diffusion_xl_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_unet=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_wuerstchen_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_prior=True, - allow_override_prior=controller.train_config.model_type.is_stable_cascade(), - has_text_encoder=True, - ) - row = self.__create_effnet_encoder_components(frame, row, ui_state) - row = self.__create_decoder_components(frame, row, ui_state, controller.train_config.model_type.is_wuerstchen_v2()) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_pixart_alpha_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - has_text_encoder=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) - - def __setup_sana_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - has_text_encoder=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, - ) + model_type = controller.train_config.model_type + parts = model_type.model_parts() - def __setup_hunyuan_video_ui(self, frame, controller, ui_state): row = 0 row = self.__create_base_dtype_components(frame, row, ui_state) row = self.__create_base_components( @@ -338,36 +27,25 @@ def __setup_hunyuan_video_ui(self, frame, controller, ui_state): row, controller, ui_state, - has_transformer=True, - allow_override_transformer=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_vae=True, - ) - row = self.__create_output_components( - frame, - row, - controller, - ui_state, + has_unet="unet" in parts, + has_prior="prior" in parts, + allow_override_prior=model_type.is_stable_cascade(), + has_transformer="transformer" in parts, + allow_override_transformer=controller.supports_override_transformer(), + has_text_encoder=not model_type.has_multiple_text_encoders(), + has_text_encoder_1=model_type.has_multiple_text_encoders(), + has_text_encoder_2="text_encoder_2" in parts, + has_text_encoder_3="text_encoder_3" in parts, + has_text_encoder_4="text_encoder_4" in parts, + allow_override_text_encoder_4="text_encoder_4" in parts, + has_vae="vae" in parts, ) + if "effnet_encoder" in parts: + row = self.__create_effnet_encoder_components(frame, row, ui_state) + if "decoder" in parts: + row = self.__create_decoder_components(frame, row, ui_state, "decoder_text_encoder" in parts) - def __setup_hi_dream_ui(self, frame, controller, ui_state): - row = 0 - row = self.__create_base_dtype_components(frame, row, ui_state) - row = self.__create_base_components( - frame, - row, - controller, - ui_state, - has_transformer=True, - has_text_encoder_1=True, - has_text_encoder_2=True, - has_text_encoder_3=True, - has_text_encoder_4=True, - allow_override_text_encoder_4=True, - has_vae=True, - ) - row = self.__create_output_components( + self.__create_output_components( frame, row, controller, diff --git a/modules/ui/ModelTabController.py b/modules/ui/ModelTabController.py index 7e5522d6e..70835d87a 100644 --- a/modules/ui/ModelTabController.py +++ b/modules/ui/ModelTabController.py @@ -13,6 +13,21 @@ def get_presets(self) -> dict: cls = create.get_model_setup_class(self.train_config.model_type, self.train_config.training_method) return cls.LAYER_PRESETS if cls is not None else {"full": []} + def supports_override_transformer(self) -> bool: + model_type = self.train_config.model_type + # The transformer override path exists only for these architectures; SD3, PixArt, Sana + # and HiDream have a transformer but expose no override field. + return ( + model_type.is_flux() + or model_type.is_z_image() + or model_type.is_ernie() + or model_type.is_chroma() + or model_type.is_qwen() + or model_type.is_anima() + or model_type.is_krea2() + or model_type.is_hunyuan_video() + ) + def get_output_formats(self) -> list[tuple[str, ModelFormat]]: labels = { ModelFormat.SAFETENSORS: "Safetensors", diff --git a/modules/ui/TopBarController.py b/modules/ui/TopBarController.py index 25593a80e..7cdd1b3bb 100644 --- a/modules/ui/TopBarController.py +++ b/modules/ui/TopBarController.py @@ -42,39 +42,13 @@ def get_model_types(self) -> list[tuple[str, ModelType]]: ] def get_training_methods(self, model_type: ModelType) -> list[tuple[str, TrainingMethod]]: - #TODO simplify - if model_type.is_stable_diffusion(): - return [ - ("Fine Tune", TrainingMethod.FINE_TUNE), - ("LoRA", TrainingMethod.LORA), - ("Embedding", TrainingMethod.EMBEDDING), - ("Fine Tune VAE", TrainingMethod.FINE_TUNE_VAE), - ] - elif model_type.is_stable_diffusion_3() \ - or model_type.is_stable_diffusion_xl() \ - or model_type.is_wuerstchen() \ - or model_type.is_pixart() \ - or model_type.is_flux_1() \ - or model_type.is_sana() \ - or model_type.is_hunyuan_video() \ - or model_type.is_hi_dream() \ - or model_type.is_chroma(): - return [ - ("Fine Tune", TrainingMethod.FINE_TUNE), - ("LoRA", TrainingMethod.LORA), - ("Embedding", TrainingMethod.EMBEDDING), - ] - elif model_type.is_qwen() \ - or model_type.is_anima() \ - or model_type.is_krea2() \ - or model_type.is_z_image() \ - or model_type.is_flux_2() \ - or model_type.is_ernie(): - return [ - ("Fine Tune", TrainingMethod.FINE_TUNE), - ("LoRA", TrainingMethod.LORA), - ] - return [] + labels = { + TrainingMethod.FINE_TUNE: "Fine Tune", + TrainingMethod.LORA: "LoRA", + TrainingMethod.EMBEDDING: "Embedding", + TrainingMethod.FINE_TUNE_VAE: "Fine Tune VAE", + } + return [(labels[m], m) for m in model_type.supported_training_methods()] def load_available_config_names(self, dir: str) -> list[tuple[str, str]]: configs = [("", path_util.canonical_join(dir, "#.json"))] diff --git a/modules/util/enum/ModelType.py b/modules/util/enum/ModelType.py index 55377ba3e..5638fd261 100644 --- a/modules/util/enum/ModelType.py +++ b/modules/util/enum/ModelType.py @@ -183,6 +183,24 @@ def is_video_model(self) -> bool: def model_parts(self) -> tuple[str, ...]: return _MODEL_PARTS[self] + def supported_training_methods(self) -> tuple[TrainingMethod, ...]: + if self.is_stable_diffusion(): + return (TrainingMethod.FINE_TUNE, TrainingMethod.LORA, TrainingMethod.EMBEDDING, TrainingMethod.FINE_TUNE_VAE) + if self.is_stable_diffusion_3() \ + or self.is_stable_diffusion_xl() \ + or self.is_wuerstchen() \ + or self.is_pixart() \ + or self.is_flux_1() \ + or self.is_sana() \ + or self.is_hunyuan_video() \ + or self.is_hi_dream() \ + or self.is_chroma(): + return (TrainingMethod.FINE_TUNE, TrainingMethod.LORA, TrainingMethod.EMBEDDING) + if self.is_qwen() or self.is_z_image() or self.is_flux_2() or self.is_ernie() \ + or self.is_anima() or self.is_krea2(): + return (TrainingMethod.FINE_TUNE, TrainingMethod.LORA) + raise ValueError(f"No supported training methods defined for model type {self}") + def denoising_model_part(self) -> str: # the denoising model component (unet / transformer / prior), always listed first in model_parts(). return _MODEL_PARTS[self][0] From baab41cd6e451b6c15b680597990414e10fe244f Mon Sep 17 00:00:00 2001 From: dxqb <183307934+dxqb@users.noreply.github.com> Date: Sat, 4 Jul 2026 10:31:53 +0200 Subject: [PATCH 5/8] Fix Anima/Krea2 checkpointing calls to match part-based conductor API BaseAnimaSetup/BaseKrea2Setup were added wholesale by the centralize-model-type merge and kept the old enabled()-guarded call pattern instead of the new part-based one applied to Chroma/Qwen during conflict resolution. --- modules/modelSetup/BaseAnimaSetup.py | 7 ++----- modules/modelSetup/BaseKrea2Setup.py | 7 ++----- modules/util/checkpointing_util.py | 12 +++++++----- 3 files changed, 11 insertions(+), 15 deletions(-) diff --git a/modules/modelSetup/BaseAnimaSetup.py b/modules/modelSetup/BaseAnimaSetup.py index dd8b7bf99..79179cdfb 100644 --- a/modules/modelSetup/BaseAnimaSetup.py +++ b/modules/modelSetup/BaseAnimaSetup.py @@ -47,11 +47,8 @@ def setup_optimizations( model: AnimaModel, config: TrainConfig, ): - if config.gradient_checkpointing.enabled(): - model.transformer_offload_conductor = \ - enable_checkpointing_for_qwen_transformer(model.transformer, config) - model.text_encoder_offload_conductor = \ - enable_checkpointing_for_qwen3_encoder_layers(model.text_encoder, config) + 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, diff --git a/modules/modelSetup/BaseKrea2Setup.py b/modules/modelSetup/BaseKrea2Setup.py index 66ed79d7d..441dc21ee 100644 --- a/modules/modelSetup/BaseKrea2Setup.py +++ b/modules/modelSetup/BaseKrea2Setup.py @@ -47,11 +47,8 @@ def setup_optimizations( model: Krea2Model, config: TrainConfig, ): - if config.gradient_checkpointing.enabled(): - model.transformer_offload_conductor = \ - enable_checkpointing_for_krea2_transformer(model.transformer, config) - model.text_encoder_offload_conductor = \ - enable_checkpointing_for_qwen3vl_encoder_layers(model.text_encoder, config) + 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, diff --git a/modules/util/checkpointing_util.py b/modules/util/checkpointing_util.py index e8fc7e7fd..e728fd632 100644 --- a/modules/util/checkpointing_util.py +++ b/modules/util/checkpointing_util.py @@ -460,9 +460,10 @@ def enable_checkpointing_for_ernie_transformer( def enable_checkpointing_for_krea2_transformer( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: # Krea2TransformerBlock takes (hidden_states, temb, image_rotary_emb, attention_mask). - return enable_checkpointing(model, config, config.compile, [ + return enable_checkpointing(model, config, part, config.compile, [ (model.text_fusion.layerwise_blocks, ["hidden_states"]), (model.text_fusion.refiner_blocks, ["hidden_states"]), (model.transformer_blocks, ["hidden_states"]), @@ -471,7 +472,8 @@ def enable_checkpointing_for_krea2_transformer( def enable_checkpointing_for_qwen3vl_encoder_layers( model: nn.Module, config: TrainConfig, -) -> LayerOffloadConductor: - return enable_checkpointing(model, config, False, [ - (Qwen3VLTextDecoderLayer, []), + part: TrainModelPartConfig, +) -> LayerOffloadConductor | None: + return enable_checkpointing(model, config, part, False, [ + (Qwen3VLTextDecoderLayer, []), # no activation offloading: this encoder is never trained ]) From 6908a6a7a091428a275f6794edfd1b99405372a4 Mon Sep 17 00:00:00 2001 From: dxqb <183307934+dxqb@users.noreply.github.com> Date: Sat, 4 Jul 2026 12:27:01 +0200 Subject: [PATCH 6/8] Address PR #1476 review comments - default activation_offloading to False for fresh configs (it previously only took effect combined with layer offloading, so a fresh fine-tune now offloads activations out of the box, which enable_activation_offloading never did) - hide the no-op "Layer Offload Fraction" control for CLIP text encoders, whose conductor is discarded in setup_optimizations - narrow the caching_threads/layer-offloading guard in create_data_loader to the text-encoder parts that actually run in the caching dataloader's worker threads, via a new ModelType.text_encoder_parts() helper - drop the dead model.vae.enable_gradient_checkpointing() call: the base SD setup only ever decodes the VAE under torch.no_grad(), so checkpointing it never had any effect; VAEs are small enough that FINE_TUNE_VAE doesn't need it either --- .../modelSetup/BaseStableDiffusionSetup.py | 2 - modules/ui/BaseTrainingTabView.py | 57 +++++++------------ modules/util/config/TrainConfig.py | 2 +- modules/util/create.py | 12 ++-- modules/util/enum/ModelType.py | 4 ++ 5 files changed, 33 insertions(+), 44 deletions(-) diff --git a/modules/modelSetup/BaseStableDiffusionSetup.py b/modules/modelSetup/BaseStableDiffusionSetup.py index fc7e5960a..dae2f8ee9 100644 --- a/modules/modelSetup/BaseStableDiffusionSetup.py +++ b/modules/modelSetup/BaseStableDiffusionSetup.py @@ -54,8 +54,6 @@ def setup_optimizations( if config.unet.checkpointing_or_offloading_enabled(): model.unet.enable_gradient_checkpointing() enable_checkpointing_for_basic_transformer_blocks(model.unet, config, config.unet, offload_enabled=False) - if config.vae.checkpointing_enabled(): - model.vae.enable_gradient_checkpointing() enable_checkpointing_for_clip_encoder_layers(model.text_encoder, config, config.text_encoder) if config.force_circular_padding: diff --git a/modules/ui/BaseTrainingTabView.py b/modules/ui/BaseTrainingTabView.py index 3dc8ce381..58664ade8 100644 --- a/modules/ui/BaseTrainingTabView.py +++ b/modules/ui/BaseTrainingTabView.py @@ -8,7 +8,6 @@ from modules.util.enum.LossWeight import LossWeight from modules.util.enum.Optimizer import Optimizer from modules.util.enum.TimestepDistribution import TimestepDistribution -from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.ui.validation_helpers import check_range, validate_resolution @@ -68,23 +67,21 @@ def build(self, column_0, column_1, column_2, controller, ui_state): def __setup_stable_diffusion_ui(self, column_0, column_1, column_2, controller, ui_state): self.__create_base_frame(column_0, 0, controller, ui_state) - self.__create_text_encoder_frame(column_0, 1, ui_state) + self.__create_text_encoder_frame(column_0, 1, ui_state, supports_layer_offloading=False) self.__create_embedding_frame(column_0, 2, ui_state) self.__create_base2_frame(column_1, 0, controller, ui_state, supports_circular_padding=True) self.__create_unet_frame(column_1, 1, ui_state) self.__create_noise_frame(column_1, 2, ui_state, supports_generalized_offset_noise=True) - if controller.config.training_method == TrainingMethod.FINE_TUNE_VAE: - self.__create_vae_frame(column_2, 0, ui_state) self.__create_masked_frame(column_2, 1, ui_state) self.__create_loss_frame(column_2, 2, controller, ui_state) self.__create_layer_frame(column_2, 3, controller, ui_state) def __setup_stable_diffusion_3_ui(self, column_0, column_1, column_2, controller, ui_state): self.__create_base_frame(column_0, 0, controller, ui_state) - self.__create_text_encoder_n_frame(column_0, 1, ui_state, i=1, supports_include=True) - self.__create_text_encoder_n_frame(column_0, 2, ui_state, i=2, supports_include=True) + self.__create_text_encoder_n_frame(column_0, 1, ui_state, i=1, supports_include=True, supports_layer_offloading=False) + self.__create_text_encoder_n_frame(column_0, 2, ui_state, i=2, supports_include=True, supports_layer_offloading=False) self.__create_text_encoder_n_frame(column_0, 3, ui_state, i=3, supports_include=True) self.__create_embedding_frame(column_0, 4, ui_state) @@ -98,8 +95,8 @@ def __setup_stable_diffusion_3_ui(self, column_0, column_1, column_2, controller def __setup_stable_diffusion_xl_ui(self, column_0, column_1, column_2, controller, ui_state): self.__create_base_frame(column_0, 0, controller, ui_state) - self.__create_text_encoder_n_frame(column_0, 1, ui_state, i=1) - self.__create_text_encoder_n_frame(column_0, 2, ui_state, i=2) + self.__create_text_encoder_n_frame(column_0, 1, ui_state, i=1, supports_layer_offloading=False) + self.__create_text_encoder_n_frame(column_0, 2, ui_state, i=2, supports_layer_offloading=False) self.__create_embedding_frame(column_0, 3, ui_state) self.__create_base2_frame(column_1, 0, controller, ui_state, supports_circular_padding=True) @@ -112,7 +109,7 @@ def __setup_stable_diffusion_xl_ui(self, column_0, column_1, column_2, controlle def __setup_wuerstchen_ui(self, column_0, column_1, column_2, controller, ui_state): self.__create_base_frame(column_0, 0, controller, ui_state) - self.__create_text_encoder_frame(column_0, 1, ui_state) + self.__create_text_encoder_frame(column_0, 1, ui_state, supports_layer_offloading=False) self.__create_embedding_frame(column_0, 2, ui_state) self.__create_base2_frame(column_1, 0, controller, ui_state, supports_circular_padding=True) @@ -138,7 +135,7 @@ def __setup_pixart_alpha_ui(self, column_0, column_1, column_2, controller, ui_s def __setup_flux_ui(self, column_0, column_1, column_2, controller, ui_state): self.__create_base_frame(column_0, 0, controller, ui_state) - self.__create_text_encoder_n_frame(column_0, 1, ui_state, i=1, supports_include=True) + self.__create_text_encoder_n_frame(column_0, 1, ui_state, i=1, supports_include=True, supports_layer_offloading=False) self.__create_text_encoder_n_frame(column_0, 2, ui_state, i=2, supports_include=True, supports_sequence_length=True) self.__create_embedding_frame(column_0, 4, ui_state) @@ -251,7 +248,7 @@ def __setup_sana_ui(self, column_0, column_1, column_2, controller, ui_state): def __setup_hunyuan_video_ui(self, column_0, column_1, column_2, controller, ui_state): self.__create_base_frame(column_0, 0, controller, ui_state) self.__create_text_encoder_n_frame(column_0, 1, ui_state, i=1, supports_include=True) - self.__create_text_encoder_n_frame(column_0, 2, ui_state, i=2, supports_include=True) + self.__create_text_encoder_n_frame(column_0, 2, ui_state, i=2, supports_include=True, supports_layer_offloading=False) self.__create_embedding_frame(column_0, 4, ui_state) self.__create_base2_frame(column_1, 0, controller, ui_state, video_training_enabled=True) @@ -264,8 +261,8 @@ def __setup_hunyuan_video_ui(self, column_0, column_1, column_2, controller, ui_ def __setup_hi_dream_ui(self, column_0, column_1, column_2, controller, ui_state): self.__create_base_frame(column_0, 0, controller, ui_state) - self.__create_text_encoder_n_frame(column_0, 1, ui_state, i=1, supports_include=True) - self.__create_text_encoder_n_frame(column_0, 2, ui_state, i=2, supports_include=True) + self.__create_text_encoder_n_frame(column_0, 1, ui_state, i=1, supports_include=True, supports_layer_offloading=False) + self.__create_text_encoder_n_frame(column_0, 2, ui_state, i=2, supports_include=True, supports_layer_offloading=False) self.__create_text_encoder_n_frame(column_0, 3, ui_state, i=3, supports_include=True) self.__create_text_encoder_n_frame(column_0, 4, ui_state, i=4, supports_include=True, supports_layer_skip=False) self.__create_embedding_frame(column_0, 5, ui_state) @@ -430,17 +427,18 @@ def __create_base2_frame(self, master, row, controller, ui_state, video_training self.components.switch(frame, row, 1, ui_state, "force_circular_padding") def __create_offloading_widgets(self, frame, row, ui_state, part, supports_checkpointing=True, - supports_activation_offloading=False): + supports_activation_offloading=False, supports_layer_offloading=True): if supports_checkpointing: self.components.label(frame, row, 0, "Gradient Checkpointing", tooltip="Enables gradient checkpointing for this component. Reduces VRAM usage at the cost of training speed") self.components.switch(frame, row, 1, ui_state, f"{part}.gradient_checkpointing") row += 1 - self.components.label(frame, row, 0, "Layer Offload Fraction", - tooltip="Fraction of this component's layers to offload to CPU to reduce VRAM usage. Increases training time and RAM usage. 0=disabled, 1=all layers") - self.components.entry(frame, row, 1, ui_state, f"{part}.offload_fraction") - row += 1 + if supports_layer_offloading: + self.components.label(frame, row, 0, "Layer Offload Fraction", + tooltip="Fraction of this component's layers to offload to CPU to reduce VRAM usage. Increases training time and RAM usage. 0=disabled, 1=all layers") + self.components.entry(frame, row, 1, ui_state, f"{part}.offload_fraction") + row += 1 if supports_activation_offloading: self.components.label(frame, row, 0, "Offload Activations", @@ -451,7 +449,7 @@ def __create_offloading_widgets(self, frame, row, ui_state, part, supports_check return row def __create_text_encoder_frame(self, master, row, ui_state, supports_clip_skip=True, supports_training=True, - supports_sequence_length=False): + supports_sequence_length=False, supports_layer_offloading=True): frame = self.components.section_frame(master, row) row = 0 @@ -465,7 +463,8 @@ def __create_text_encoder_frame(self, master, row, ui_state, supports_clip_skip= self.components.label(frame, row, 0, "Text Encoder") row += 1 - row = self.__create_offloading_widgets(frame, row, ui_state, "text_encoder", supports_checkpointing=supports_training) + row = self.__create_offloading_widgets(frame, row, ui_state, "text_encoder", supports_checkpointing=supports_training, + supports_layer_offloading=supports_layer_offloading) # dropout self.components.label(frame, row, 0, "Caption Dropout Probability", @@ -510,6 +509,7 @@ def __create_text_encoder_n_frame( supports_include: bool = False, supports_layer_skip: bool = True, supports_sequence_length: bool = False, + supports_layer_offloading: bool = True, ): frame = self.components.section_frame(master, row) row = 0 @@ -529,7 +529,8 @@ def __create_text_encoder_n_frame( self.components.switch(frame, row, 1, ui_state, f"text_encoder{suffix}.train") row += 1 - row = self.__create_offloading_widgets(frame, row, ui_state, f"text_encoder{suffix}") + row = self.__create_offloading_widgets(frame, row, ui_state, f"text_encoder{suffix}", + supports_layer_offloading=supports_layer_offloading) # train text encoder embedding self.components.label(frame, row, 0, f"Train Text Encoder {i} Embedding", @@ -615,20 +616,6 @@ def __create_unet_frame(self, master, row, ui_state): self.components.switch(frame, row, 1, ui_state, "rescale_noise_scheduler_to_zero_terminal_snr") row += 1 - def __create_vae_frame(self, master, row, ui_state): - frame = self.components.section_frame(master, row) - row = 0 - - self.components.label(frame, row, 0, "Train VAE", - tooltip="Enables training the VAE model") - self.components.switch(frame, row, 1, ui_state, "vae.train") - row += 1 - - self.components.label(frame, row, 0, "Gradient Checkpointing", - tooltip="Enables gradient checkpointing for the VAE. Reduces VRAM usage at the cost of training speed") - self.components.switch(frame, row, 1, ui_state, "vae.gradient_checkpointing") - row += 1 - def __create_prior_frame(self, master, row, ui_state): frame = self.components.section_frame(master, row) row = 0 diff --git a/modules/util/config/TrainConfig.py b/modules/util/config/TrainConfig.py index baffec45e..8d77f9c6b 100644 --- a/modules/util/config/TrainConfig.py +++ b/modules/util/config/TrainConfig.py @@ -306,7 +306,7 @@ def default_values(): data.append(("guidance_scale", 1.0, float, False)) data.append(("gradient_checkpointing", True, bool, False)) data.append(("offload_fraction", 0.0, float, False)) - data.append(("activation_offloading", True, bool, False)) + data.append(("activation_offloading", False, bool, False)) return TrainModelPartConfig(data) diff --git a/modules/util/create.py b/modules/util/create.py index 0c0adba16..5cda076fc 100644 --- a/modules/util/create.py +++ b/modules/util/create.py @@ -110,12 +110,12 @@ def create_data_loader( train_progress: TrainProgress | None = None, is_validation: bool = False ) -> BaseDataLoader | None: - # Layer offloading uses a non-thread-safe conductor. This check is too broad: it trips whenever any model - # part does layer offloading, even though only a component that is actually cached really runs in the - # dataloader worker threads. - # TODO: narrow this to the cached components only. - if config.dataloader_threads > 1 and any(part.offload_fraction > 0 for part in config.model_part_configs()): - raise RuntimeError('layer offloading can not be activated if "dataloader_threads" > 1') + # Layer offloading uses a non-thread-safe conductor. Only text encoders run inside the caching dataloader's + # worker threads (to produce the text cache), so only their offload_fraction can conflict with threading. + if config.dataloader_threads > 1 and any( + getattr(config, name).offload_fraction > 0 for name in model_type.text_encoder_parts() + ): + raise RuntimeError('layer offloading can not be activated for a text encoder if "dataloader_threads" > 1') if train_progress is None: train_progress = TrainProgress() diff --git a/modules/util/enum/ModelType.py b/modules/util/enum/ModelType.py index 242b0ff4e..430c7204d 100644 --- a/modules/util/enum/ModelType.py +++ b/modules/util/enum/ModelType.py @@ -205,6 +205,10 @@ def denoising_model_part(self) -> str: # the denoising model component (unet / transformer / prior), always listed first in model_parts(). return _MODEL_PARTS[self][0] + def text_encoder_parts(self) -> tuple[str, ...]: + # the text encoder components, named "text_encoder"/"text_encoder_2"/... by convention (see below). + return tuple(part for part in _MODEL_PARTS[self] if part.startswith("text_encoder")) + def supported_lora_formats(self) -> list[ModelFormat]: formats = [ ModelFormat.DIFFUSERS_LORA, From a30f033bb14e519e3599b9c449de262b347831f3 Mon Sep 17 00:00:00 2001 From: dxqb <183307934+dxqb@users.noreply.github.com> Date: Sat, 4 Jul 2026 14:16:56 +0200 Subject: [PATCH 7/8] Address further PR #1476 review: decouple checkpointing/offloading from train - do not gate checkpointing or offloading on a part's train flag: grad can flow through a frozen part (embedding training runs grad through the frozen denoiser and text encoders), where checkpointing and offloading still save VRAM. checkpointing_enabled() -> gradient_checkpointing; offloading_enabled() -> offload_fraction > 0 or activation_offloading. Restores pre-branch behaviour for embedding training. - reject offloading with checkpointing disabled (raise NotImplementedError) instead of silently forcing checkpointing on: the current offload modes need the use_reentrant recompute to service the backward pass. - CLIP encoder checkpointing passes offload_enabled=False so a migrated offload_fraction can never build a self-activating conductor for a non-offloadable text encoder. - hide the dead Layer Offload Fraction / Offload Activations controls for the SD/SDXL unet and Wuerstchen prior (no conductor), and gate the SD/SDXL native gradient checkpointing on checkpointing_enabled() so a dead offload toggle can no longer flip it on. - derive ModelType.has_multiple_text_encoders() from model_parts() instead of a parallel is_*() chain. - remove the dead LayerOffloadConductor.layer_offload_activated(). Co-Authored-By: Claude Opus 4.8 --- modules/modelSetup/BaseStableDiffusionSetup.py | 2 +- modules/modelSetup/BaseStableDiffusionXLSetup.py | 2 +- modules/ui/BaseTrainingTabView.py | 4 ++-- modules/util/LayerOffloadConductor.py | 3 --- modules/util/checkpointing_util.py | 11 +++++++---- modules/util/config/TrainConfig.py | 15 +++++++++------ modules/util/enum/ModelType.py | 6 +----- 7 files changed, 21 insertions(+), 22 deletions(-) diff --git a/modules/modelSetup/BaseStableDiffusionSetup.py b/modules/modelSetup/BaseStableDiffusionSetup.py index dae2f8ee9..f8aa5dec3 100644 --- a/modules/modelSetup/BaseStableDiffusionSetup.py +++ b/modules/modelSetup/BaseStableDiffusionSetup.py @@ -51,7 +51,7 @@ def setup_optimizations( model: StableDiffusionModel, config: TrainConfig, ): - if config.unet.checkpointing_or_offloading_enabled(): + 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_clip_encoder_layers(model.text_encoder, config, config.text_encoder) diff --git a/modules/modelSetup/BaseStableDiffusionXLSetup.py b/modules/modelSetup/BaseStableDiffusionXLSetup.py index 98d52cbec..17d9cf4e0 100644 --- a/modules/modelSetup/BaseStableDiffusionXLSetup.py +++ b/modules/modelSetup/BaseStableDiffusionXLSetup.py @@ -48,7 +48,7 @@ def setup_optimizations( model: StableDiffusionXLModel, config: TrainConfig, ): - if config.unet.checkpointing_or_offloading_enabled(): + 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_clip_encoder_layers(model.text_encoder_1, config, config.text_encoder) diff --git a/modules/ui/BaseTrainingTabView.py b/modules/ui/BaseTrainingTabView.py index 58664ade8..eef9c82e3 100644 --- a/modules/ui/BaseTrainingTabView.py +++ b/modules/ui/BaseTrainingTabView.py @@ -594,7 +594,7 @@ def __create_unet_frame(self, master, row, ui_state): self.components.switch(frame, row, 1, ui_state, "unet.train") row += 1 - row = self.__create_offloading_widgets(frame, row, ui_state, "unet", supports_activation_offloading=True) + row = self.__create_offloading_widgets(frame, row, ui_state, "unet", supports_layer_offloading=False) # train unet epochs self.components.label(frame, row, 0, "Stop Training After", @@ -626,7 +626,7 @@ def __create_prior_frame(self, master, row, ui_state): self.components.switch(frame, row, 1, ui_state, "prior.train") row += 1 - row = self.__create_offloading_widgets(frame, row, ui_state, "prior", supports_activation_offloading=True) + row = self.__create_offloading_widgets(frame, row, ui_state, "prior", supports_layer_offloading=False) # train prior epochs self.components.label(frame, row, 0, "Stop Training After", diff --git a/modules/util/LayerOffloadConductor.py b/modules/util/LayerOffloadConductor.py index 062d1391e..a69094c75 100644 --- a/modules/util/LayerOffloadConductor.py +++ b/modules/util/LayerOffloadConductor.py @@ -618,9 +618,6 @@ def __init__( def offload_activated(self) -> bool: return self.__offload_activations or self.__offload_layers - def layer_offload_activated(self) -> bool: - return self.__offload_layers - def to(self, device: torch.device): torch_gc() diff --git a/modules/util/checkpointing_util.py b/modules/util/checkpointing_util.py index e728fd632..31b3819ea 100644 --- a/modules/util/checkpointing_util.py +++ b/modules/util/checkpointing_util.py @@ -178,9 +178,12 @@ def create_checkpoint( conductor.add_layer(orig_module, included_offload_param_indices) if conductor is not None and conductor.offload_activated(): - # offloading is structurally coupled to use_reentrant=True checkpointing during the back pass - # (the recompute is what fires before_layer/after_layer in the backward direction), so the offload - # layer always checkpoints when grad is enabled, regardless of the part's gradient_checkpointing flag. + # offloading is structurally coupled to use_reentrant=True checkpointing during the back pass: + # the recompute is the only thing firing before_layer/after_layer in the backward direction, so + # both layer and activation offloading need checkpointing to move tensors back for backward. + # Rather than silently forcing checkpointing on when the part disabled it, reject the combination. + if not checkpointing: + raise NotImplementedError("offloading currently requires gradient checkpointing") if compile: layer = OffloadCheckpointLayer(orig_module=orig_module, orig_forward=None, train_device=train_device, conductor=conductor, layer_index=layer_index) #don't compile the checkpointing layer - offloading cannot be compiled: @@ -296,7 +299,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 def enable_checkpointing_for_t5_encoder_layers( model: nn.Module, diff --git a/modules/util/config/TrainConfig.py b/modules/util/config/TrainConfig.py index 8d77f9c6b..da7c936cc 100644 --- a/modules/util/config/TrainConfig.py +++ b/modules/util/config/TrainConfig.py @@ -275,14 +275,17 @@ def __init__(self, data: list[(str, Any, type, bool)]): super().__init__(data) def offloading_enabled(self) -> bool: - # a conductor should exist iff this is True. Layer offloading applies even to frozen parts (to fit - # them in VRAM), but activation offloading only does work during a backward pass, so it only applies - # when the part is trained -- even if activation_offloading is True in the config. - return self.offload_fraction > 0 or (self.activation_offloading and self.train) + # a conductor should exist iff this is True. Not gated on train: activations are built and cost + # VRAM whenever grad flows through a part, not only when it is trained (e.g. embedding training + # runs grad through the frozen denoiser and TE), so both layer and activation offloading apply to + # frozen grad-carrying parts too. + return self.offload_fraction > 0 or self.activation_offloading def checkpointing_enabled(self) -> bool: - # the inner torch checkpoint() should run iff this is True - return self.gradient_checkpointing and self.train + # the inner torch checkpoint() should run iff this is True. Not gated on train: grad can flow + # through a frozen part (e.g. embedding training runs grad through the frozen denoiser and TE), + # and checkpointing then still saves VRAM. On a part with no grad flow the wrapper is a no-op. + return self.gradient_checkpointing def checkpointing_or_offloading_enabled(self) -> bool: # whether the checkpoint layer wrapper needs to be installed for this part at all diff --git a/modules/util/enum/ModelType.py b/modules/util/enum/ModelType.py index 430c7204d..1e927a69f 100644 --- a/modules/util/enum/ModelType.py +++ b/modules/util/enum/ModelType.py @@ -140,11 +140,7 @@ def has_depth_input(self): return self == ModelType.STABLE_DIFFUSION_20_DEPTH def has_multiple_text_encoders(self): - return self.is_stable_diffusion_3() \ - or self.is_stable_diffusion_xl() \ - or self.is_flux_1() \ - or self.is_hunyuan_video() \ - or self.is_hi_dream() \ + return "text_encoder_2" in self.model_parts() def is_sd_v1(self): return self == ModelType.STABLE_DIFFUSION_15 \ From f40ba013dc6b18ea90d9bc8dd7d01603ca0028ac Mon Sep 17 00:00:00 2001 From: dxqb <183307934+dxqb@users.noreply.github.com> Date: Sun, 5 Jul 2026 08:34:44 +0200 Subject: [PATCH 8/8] Remove stale text_encoder None-checks in Chroma/Qwen/ZImage/Flux2 setups These models always load their text encoder, but the checkpointing-refactor merge on this branch didn't pick up the None-check removal from #1579 because the surrounding lines had already diverged. Co-Authored-By: Claude Sonnet 5 --- modules/modelSetup/BaseChromaSetup.py | 3 +-- modules/modelSetup/BaseFlux2Setup.py | 9 ++++----- modules/modelSetup/BaseQwenSetup.py | 3 +-- modules/modelSetup/BaseZImageSetup.py | 3 +-- 4 files changed, 7 insertions(+), 11 deletions(-) diff --git a/modules/modelSetup/BaseChromaSetup.py b/modules/modelSetup/BaseChromaSetup.py index 952a1f3f6..006587b71 100644 --- a/modules/modelSetup/BaseChromaSetup.py +++ b/modules/modelSetup/BaseChromaSetup.py @@ -50,8 +50,7 @@ def setup_optimizations( config: TrainConfig, ): model.transformer_offload_conductor = enable_checkpointing_for_chroma_transformer(model.transformer, config, config.transformer) - if model.text_encoder is not None: - model.text_encoder_offload_conductor = enable_checkpointing_for_t5_encoder_layers(model.text_encoder, config, config.text_encoder) + 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, diff --git a/modules/modelSetup/BaseFlux2Setup.py b/modules/modelSetup/BaseFlux2Setup.py index c9d331ded..ef22c6ef6 100644 --- a/modules/modelSetup/BaseFlux2Setup.py +++ b/modules/modelSetup/BaseFlux2Setup.py @@ -46,11 +46,10 @@ def setup_optimizations( config: TrainConfig, ): model.transformer_offload_conductor = enable_checkpointing_for_flux2_transformer(model.transformer, config, config.transformer) - if model.text_encoder is not None: - 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) + 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.weight_dtypes().transformer, diff --git a/modules/modelSetup/BaseQwenSetup.py b/modules/modelSetup/BaseQwenSetup.py index fe718f343..c229fd695 100644 --- a/modules/modelSetup/BaseQwenSetup.py +++ b/modules/modelSetup/BaseQwenSetup.py @@ -47,8 +47,7 @@ def setup_optimizations( config: TrainConfig, ): model.transformer_offload_conductor = enable_checkpointing_for_qwen_transformer(model.transformer, config, config.transformer) - if model.text_encoder is not None: - model.text_encoder_offload_conductor = enable_checkpointing_for_qwen25vl_encoder_layers(model.text_encoder, config, config.text_encoder) + 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, diff --git a/modules/modelSetup/BaseZImageSetup.py b/modules/modelSetup/BaseZImageSetup.py index bbf8c6073..3ea818952 100644 --- a/modules/modelSetup/BaseZImageSetup.py +++ b/modules/modelSetup/BaseZImageSetup.py @@ -48,8 +48,7 @@ def setup_optimizations( config: TrainConfig, ): model.transformer_offload_conductor = enable_checkpointing_for_z_image_transformer(model.transformer, config, config.transformer) - if model.text_encoder is not None: - model.text_encoder_offload_conductor = enable_checkpointing_for_qwen3_encoder_layers(model.text_encoder, config, config.text_encoder) + 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,