diff --git a/modules/dataLoader/AnimaBaseDataLoader.py b/modules/dataLoader/AnimaBaseDataLoader.py index 9dba5d433..9acdbcb6a 100644 --- a/modules/dataLoader/AnimaBaseDataLoader.py +++ b/modules/dataLoader/AnimaBaseDataLoader.py @@ -109,7 +109,7 @@ def _debug_modules(self, config: TrainConfig, model: AnimaModel): #TODO clean up debug_dir = os.path.join(config.debug_dir, "dataloader") def before_save_fun(): - model.vae_to(self.train_device) + model.materialize_only("vae") decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) upscale_mask = ScaleImage(in_name='latent_mask', out_name='decoded_mask', factor=8) diff --git a/modules/dataLoader/ChromaBaseDataLoader.py b/modules/dataLoader/ChromaBaseDataLoader.py index 6b89457e6..b83beec1c 100644 --- a/modules/dataLoader/ChromaBaseDataLoader.py +++ b/modules/dataLoader/ChromaBaseDataLoader.py @@ -120,7 +120,7 @@ def _debug_modules(self, config: TrainConfig, model: ChromaModel): #TODO clean u debug_dir = os.path.join(config.debug_dir, "dataloader") def before_save_fun(): - model.vae_to(self.train_device) + model.materialize_only("vae") decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) upscale_mask = ScaleImage(in_name='latent_mask', out_name='decoded_mask', factor=8) diff --git a/modules/dataLoader/ErnieBaseDataLoader.py b/modules/dataLoader/ErnieBaseDataLoader.py index 032180bc3..dacedd807 100644 --- a/modules/dataLoader/ErnieBaseDataLoader.py +++ b/modules/dataLoader/ErnieBaseDataLoader.py @@ -110,7 +110,7 @@ def _debug_modules(self, config: TrainConfig, model: ErnieModel): debug_dir = os.path.join(config.debug_dir, "dataloader") def before_save_fun(): - model.vae_to(self.train_device) + model.materialize_only("vae") decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) upscale_mask = ScaleImage(in_name='latent_mask', out_name='decoded_mask', factor=8) diff --git a/modules/dataLoader/Flux2BaseDataLoader.py b/modules/dataLoader/Flux2BaseDataLoader.py index a6bc3a05d..bfe191457 100644 --- a/modules/dataLoader/Flux2BaseDataLoader.py +++ b/modules/dataLoader/Flux2BaseDataLoader.py @@ -117,7 +117,7 @@ def _debug_modules(self, config: TrainConfig, model: Flux2Model): debug_dir = os.path.join(config.debug_dir, "dataloader") def before_save_fun(): - model.vae_to(self.train_device) + model.materialize_only("vae") decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) upscale_mask = ScaleImage(in_name='latent_mask', out_name='decoded_mask', factor=8) diff --git a/modules/dataLoader/FluxBaseDataLoader.py b/modules/dataLoader/FluxBaseDataLoader.py index d23dee3a8..e18663311 100644 --- a/modules/dataLoader/FluxBaseDataLoader.py +++ b/modules/dataLoader/FluxBaseDataLoader.py @@ -143,7 +143,7 @@ def _debug_modules(self, config: TrainConfig, model: FluxModel): debug_dir = os.path.join(config.debug_dir, "dataloader") def before_save_fun(): - model.vae_to(self.train_device) + model.materialize_only("vae") decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) decode_conditioning_image = DecodeVAE(in_name='latent_conditioning_image', out_name='decoded_conditioning_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) diff --git a/modules/dataLoader/HiDreamBaseDataLoader.py b/modules/dataLoader/HiDreamBaseDataLoader.py index f080943bf..9ed754bec 100644 --- a/modules/dataLoader/HiDreamBaseDataLoader.py +++ b/modules/dataLoader/HiDreamBaseDataLoader.py @@ -180,7 +180,7 @@ def _debug_modules(self, config: TrainConfig, model: HiDreamModel): debug_dir = os.path.join(config.debug_dir, "dataloader") def before_save_fun(): - model.vae_to(self.train_device) + model.materialize_only("vae") decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) decode_conditioning_image = DecodeVAE(in_name='latent_conditioning_image', out_name='decoded_conditioning_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) diff --git a/modules/dataLoader/HunyuanVideoBaseDataLoader.py b/modules/dataLoader/HunyuanVideoBaseDataLoader.py index 38f7e6a8a..308864efa 100644 --- a/modules/dataLoader/HunyuanVideoBaseDataLoader.py +++ b/modules/dataLoader/HunyuanVideoBaseDataLoader.py @@ -136,7 +136,7 @@ def _debug_modules(self, config: TrainConfig, model: HunyuanVideoModel): debug_dir = os.path.join(config.debug_dir, "dataloader") def before_save_fun(): - model.vae_to(self.train_device) + model.materialize_only("vae") decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) upscale_mask = ScaleImage(in_name='latent_mask', out_name='decoded_mask', factor=8) diff --git a/modules/dataLoader/IdeogramBaseDataLoader.py b/modules/dataLoader/IdeogramBaseDataLoader.py index 28de23e40..4f987dac4 100644 --- a/modules/dataLoader/IdeogramBaseDataLoader.py +++ b/modules/dataLoader/IdeogramBaseDataLoader.py @@ -122,7 +122,7 @@ def _debug_modules(self, config: TrainConfig, model: IdeogramModel): debug_dir = os.path.join(config.debug_dir, "dataloader") def before_save_fun(): - model.vae_to(self.train_device) + model.materialize_only("vae") decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) upscale_mask = ScaleImage(in_name='latent_mask', out_name='decoded_mask', factor=8) diff --git a/modules/dataLoader/Krea2BaseDataLoader.py b/modules/dataLoader/Krea2BaseDataLoader.py index 5d243f637..c1a3dd083 100644 --- a/modules/dataLoader/Krea2BaseDataLoader.py +++ b/modules/dataLoader/Krea2BaseDataLoader.py @@ -131,7 +131,7 @@ def _debug_modules(self, config: TrainConfig, model: Krea2Model): #TODO clean up debug_dir = os.path.join(config.debug_dir, "dataloader") def before_save_fun(): - model.vae_to(self.train_device) + model.materialize_only("vae") decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) upscale_mask = ScaleImage(in_name='latent_mask', out_name='decoded_mask', factor=8) diff --git a/modules/dataLoader/PixArtAlphaBaseDataLoader.py b/modules/dataLoader/PixArtAlphaBaseDataLoader.py index bb110dc2e..5e56496d4 100644 --- a/modules/dataLoader/PixArtAlphaBaseDataLoader.py +++ b/modules/dataLoader/PixArtAlphaBaseDataLoader.py @@ -121,7 +121,7 @@ def _debug_modules(self, config: TrainConfig, model: PixArtAlphaModel): debug_dir = os.path.join(config.debug_dir, "dataloader") def before_save_fun(): - model.vae_to(self.train_device) + model.materialize_only("vae") decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) decode_conditioning_image = DecodeVAE(in_name='latent_conditioning_image', out_name='decoded_conditioning_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) diff --git a/modules/dataLoader/QwenBaseDataLoader.py b/modules/dataLoader/QwenBaseDataLoader.py index 9a4a962a9..8e99f7cab 100644 --- a/modules/dataLoader/QwenBaseDataLoader.py +++ b/modules/dataLoader/QwenBaseDataLoader.py @@ -124,7 +124,7 @@ def _debug_modules(self, config: TrainConfig, model: QwenModel): #TODO clean up debug_dir = os.path.join(config.debug_dir, "dataloader") def before_save_fun(): - model.vae_to(self.train_device) + model.materialize_only("vae") decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) upscale_mask = ScaleImage(in_name='latent_mask', out_name='decoded_mask', factor=8) diff --git a/modules/dataLoader/SanaBaseDataLoader.py b/modules/dataLoader/SanaBaseDataLoader.py index a44ff8130..983d8d439 100644 --- a/modules/dataLoader/SanaBaseDataLoader.py +++ b/modules/dataLoader/SanaBaseDataLoader.py @@ -113,7 +113,7 @@ def _debug_modules(self, config: TrainConfig, model: SanaModel): debug_dir = os.path.join(config.debug_dir, "dataloader") def before_save_fun(): - model.vae_to(self.train_device) + model.materialize_only("vae") decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae, autocast_contexts=[model.autocast_context, model.vae_autocast_context], dtype=model.train_dtype.torch_dtype()) decode_conditioning_image = DecodeVAE(in_name='latent_conditioning_image', out_name='decoded_conditioning_image', vae=model.vae, autocast_contexts=[model.autocast_context, model.vae_autocast_context], dtype=model.train_dtype.torch_dtype()) diff --git a/modules/dataLoader/StableDiffusion3BaseDataLoader.py b/modules/dataLoader/StableDiffusion3BaseDataLoader.py index 55a0d9001..b9f261ba3 100644 --- a/modules/dataLoader/StableDiffusion3BaseDataLoader.py +++ b/modules/dataLoader/StableDiffusion3BaseDataLoader.py @@ -160,7 +160,7 @@ def _debug_modules(self, config: TrainConfig, model: StableDiffusion3Model): debug_dir = os.path.join(config.debug_dir, "dataloader") def before_save_fun(): - model.vae_to(self.train_device) + model.materialize_only("vae") decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) decode_conditioning_image = DecodeVAE(in_name='latent_conditioning_image', out_name='decoded_conditioning_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) diff --git a/modules/dataLoader/StableDiffusionBaseDataLoader.py b/modules/dataLoader/StableDiffusionBaseDataLoader.py index 63ad57cac..f5f45c274 100644 --- a/modules/dataLoader/StableDiffusionBaseDataLoader.py +++ b/modules/dataLoader/StableDiffusionBaseDataLoader.py @@ -130,7 +130,7 @@ def _debug_modules(self, config: TrainConfig, model: StableDiffusionModel): debug_dir = os.path.join(config.debug_dir, "dataloader") def before_save_fun(): - model.vae_to(self.train_device) + model.materialize_only("vae") decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) decode_conditioning_image = DecodeVAE(in_name='latent_conditioning_image', out_name='decoded_conditioning_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) diff --git a/modules/dataLoader/StableDiffusionFineTuneVaeDataLoader.py b/modules/dataLoader/StableDiffusionFineTuneVaeDataLoader.py index ad0c890b0..a0eda2913 100644 --- a/modules/dataLoader/StableDiffusionFineTuneVaeDataLoader.py +++ b/modules/dataLoader/StableDiffusionFineTuneVaeDataLoader.py @@ -8,7 +8,6 @@ from modules.util.config.TrainConfig import TrainConfig from modules.util.enum.ModelType import ModelType from modules.util.enum.TrainingMethod import TrainingMethod -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress from mgds.OutputPipelineModule import OutputPipelineModule @@ -55,12 +54,9 @@ def _setup_cache_device( temp_device: torch.device, config: TrainConfig, ): - model.to(self.temp_device) - - model.vae_to(train_device) + model.materialize_only("vae") model.eval() - torch_gc() def __enumerate_input_modules(self, config: TrainConfig) -> list: supported_extensions = path_util.supported_image_extensions() @@ -248,7 +244,7 @@ def __debug_modules(self, config: TrainConfig, model: StableDiffusionModel): debug_dir = os.path.join(config.debug_dir, "dataloader") def before_save_fun(): - model.vae_to(self.train_device) + model.materialize_only("vae") decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae) upscale_mask = ScaleImage(in_name='latent_mask', out_name='decoded_mask', factor=8) diff --git a/modules/dataLoader/StableDiffusionXLBaseDataLoader.py b/modules/dataLoader/StableDiffusionXLBaseDataLoader.py index 12739adf8..3a4fc9517 100644 --- a/modules/dataLoader/StableDiffusionXLBaseDataLoader.py +++ b/modules/dataLoader/StableDiffusionXLBaseDataLoader.py @@ -138,7 +138,7 @@ def _debug_modules(self, config: TrainConfig, model: StableDiffusionXLModel): debug_dir = os.path.join(config.debug_dir, "dataloader") def before_save_fun(): - model.vae_to(self.train_device) + model.materialize_only("vae") decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae, autocast_contexts=[model.autocast_context, model.vae_autocast_context], dtype=model.vae_train_dtype.torch_dtype()) decode_conditioning_image = DecodeVAE(in_name='latent_conditioning_image', out_name='decoded_conditioning_image', vae=model.vae, autocast_contexts=[model.autocast_context, model.vae_autocast_context], dtype=model.vae_train_dtype.torch_dtype()) diff --git a/modules/dataLoader/WuerstchenBaseDataLoader.py b/modules/dataLoader/WuerstchenBaseDataLoader.py index 4689b09a4..e6f27d621 100644 --- a/modules/dataLoader/WuerstchenBaseDataLoader.py +++ b/modules/dataLoader/WuerstchenBaseDataLoader.py @@ -10,7 +10,6 @@ from modules.util import factory from modules.util.config.TrainConfig import TrainConfig from modules.util.enum.ModelType import ModelType -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress from mgds.pipelineModules.DecodeTokens import DecodeTokens @@ -74,10 +73,8 @@ def _cache_modules(self, config: TrainConfig, model: WuerstchenModel, model_setu ] def before_cache_image_fun(): - model.to(self.temp_device) - model.effnet_encoder_to(self.train_device) + model.materialize_only("effnet_encoder") model.eval() - torch_gc() return self._cache_modules_from_names( model, model_setup, @@ -109,10 +106,8 @@ def _output_modules(self, config: TrainConfig, model: WuerstchenModel, model_set output_names.append('pooled_text_encoder_output') def before_cache_image_fun(): - model.to(self.temp_device) - model.effnet_encoder_to(self.train_device) + model.materialize_only("effnet_encoder") model.eval() - torch_gc() return self._output_modules_from_out_names( model, model_setup, diff --git a/modules/dataLoader/ZImageBaseDataLoader.py b/modules/dataLoader/ZImageBaseDataLoader.py index 23863d625..2617e8561 100644 --- a/modules/dataLoader/ZImageBaseDataLoader.py +++ b/modules/dataLoader/ZImageBaseDataLoader.py @@ -116,7 +116,7 @@ def _debug_modules(self, config: TrainConfig, model: ZImageModel): debug_dir = os.path.join(config.debug_dir, "dataloader") def before_save_fun(): - model.vae_to(self.train_device) + model.materialize_only("vae") decode_image = DecodeVAE(in_name='latent_image', out_name='decoded_image', vae=model.vae, autocast_contexts=[model.autocast_context], dtype=model.train_dtype.torch_dtype()) upscale_mask = ScaleImage(in_name='latent_mask', out_name='decoded_mask', factor=8) diff --git a/modules/dataLoader/mixin/DataLoaderText2ImageMixin.py b/modules/dataLoader/mixin/DataLoaderText2ImageMixin.py index 2654bdd19..a9ed0970d 100644 --- a/modules/dataLoader/mixin/DataLoaderText2ImageMixin.py +++ b/modules/dataLoader/mixin/DataLoaderText2ImageMixin.py @@ -10,7 +10,6 @@ from modules.util import path_util from modules.util.config.TrainConfig import TrainConfig from modules.util.enum.DataType import DataType -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress from mgds.OutputPipelineModule import OutputPipelineModule @@ -274,10 +273,8 @@ def _output_modules_from_out_names( ): if before_cache_image_fun is None: def prepare_vae(): - model.to(self.temp_device) - model.vae_to(self.train_device) + model.materialize_only("vae") model.eval() - torch_gc() before_cache_image_fun = prepare_vae sort_names = output_names + ['concept'] @@ -340,10 +337,8 @@ def _cache_modules_from_names( if before_cache_image_fun is None: def prepare_vae(): - model.to(self.temp_device) - model.vae_to(self.train_device) + model.materialize_only("vae") model.eval() - torch_gc() before_cache_image_fun = prepare_vae def before_cache_text_fun(): diff --git a/modules/model/AnimaModel.py b/modules/model/AnimaModel.py index a23d2998d..9c365273d 100644 --- a/modules/model/AnimaModel.py +++ b/modules/model/AnimaModel.py @@ -72,11 +72,6 @@ def __init__( self.transformer_lora = None self.lora_state_dict = None - def adapters(self) -> list[LoRAModuleWrapper]: - return [a for a in [ - self.transformer_lora, - ] if a is not None] - def _diffusers_to_dit(self) -> list: # the netless diffusers CosmosTransformer3DModel -> Anima DiT rename (the inverse of diffusers' # scripts/convert_anima_to_diffusers.py transformer rename). These are the bare module names kohya-ss @@ -130,35 +125,22 @@ def lora_diffusers_to_kohya(self) -> list | None: # kohya-ss loads the DiT with the net. wrapper stripped -> the netless body. return self._diffusers_to_dit() - 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: - self.text_encoder_offload_conductor.to(device) - else: - self.text_encoder.to(device=device) - self.text_conditioner.to(device=device) - - def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None: - self.transformer_offload_conductor.to(device) - else: - self.transformer.to(device=device) - - if self.transformer_lora is not None: - self.transformer_lora.to(device) + def materialize(self, *parts: str): + super().materialize(*parts) + # text_conditioner isn't in ModelType.model_parts(); it always travels with text_encoder. + if "text_encoder" in parts: + self.text_conditioner.to(device=self.train_device) - def to(self, device: torch.device): - self.vae_to(device) - self.text_encoder_to(device) - self.transformer_to(device) + def evict(self, *parts: str): + super().evict(*parts) + # evict() with no parts means "evict all", which includes text_encoder. + if not parts or "text_encoder" in parts: + self.text_conditioner.to(device=self.temp_device) def eval(self): - self.vae.eval() - self.text_encoder.eval() + super().eval() + # text_conditioner isn't in ModelType.model_parts(); it always travels with text_encoder. self.text_conditioner.eval() - self.transformer.eval() def create_pipeline(self): pipe = AnimaAutoBlocks().init_pipeline() diff --git a/modules/model/BaseModel.py b/modules/model/BaseModel.py index c346a60fe..2f0106dc2 100644 --- a/modules/model/BaseModel.py +++ b/modules/model/BaseModel.py @@ -1,4 +1,4 @@ -from abc import ABCMeta, abstractmethod +from abc import ABCMeta from contextlib import nullcontext from uuid import uuid4 @@ -11,6 +11,7 @@ from modules.util.enum.ModelType import ModelType from modules.util.modelSpec.ModelSpec import ModelSpec from modules.util.NamedParameterGroup import NamedParameterGroupCollection +from modules.util.torch_util import device_equals, torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -96,17 +97,82 @@ def __init__( self.autocast_context = nullcontext() self.train_dtype = DataType.FLOAT_32 - @abstractmethod - def to(self, device: torch.device): - pass + @property + def train_device(self) -> torch.device: + return torch.device(self.train_config.train_device) + + @property + def temp_device(self) -> torch.device: + return torch.device(self.train_config.temp_device) + + def materialize(self, *parts: str): + # Move `parts` onto train_device. + for part in parts: + self._move_part(part, self.train_device) + + def evict(self, *parts: str): + # Move `parts` onto temp_device. No parts given -> every component in ModelType.model_parts(). + for part in parts or self.model_type.model_parts(): + self._move_part(part, self.temp_device) + torch_gc() + + def materialize_only(self, *parts: str): + # Materialize exactly `parts` on train_device; evict every other component in ModelType.model_parts() + # to temp_device. Lets a caller state what it needs now without tracking what to evict first. + # Evicts before materializing, so the two sets are never resident on train_device at once. + # Skipped (rather than passed as evict()) when empty, since evict() with no parts means "evict all". + to_evict = [part for part in self.model_type.model_parts() if part not in parts] + if to_evict: + self.evict(*to_evict) + self.materialize(*parts) + + def materialize_only_text_encoders(self): + # Materialize all of this model's text encoders on train_device, evicting everything else. Samplers + # call this before encode_text, which reads every text encoder the model has. + self.materialize_only(*self.model_type.text_encoder_parts()) + + def _move_part(self, part: str, device: torch.device): + # The generic per-component move: `part` (or `part_1` for the first of several split text encoders), + # its LoRA (`{part}_lora`), and its layer-offload conductor (`{part}_offload_conductor`), if present. + stem = f"{part}_1" if hasattr(self, f"{part}_1") else part + + conductor = getattr(self, f"{stem}_offload_conductor", None) + if conductor is not None: + if device_equals(device, self.train_device): + conductor.materialize() + else: + conductor.evict() + else: + component = getattr(self, stem) # raises if `part` doesn't name a real attribute + # None when the part is excluded from training (e.g. a text encoder with include_text_encoder off): + # it stays in model_parts() but the loader never populated it, so there is nothing to move. + if component is not None: + component.to(device=device) + + lora = getattr(self, f"{stem}_lora", None) + if lora is not None: + lora.to(device) - @abstractmethod def eval(self): - pass + # Put every present component on eval(); driven by the same part registry as materialize()/evict(). + # A model whose component names diverge (Wuerstchen) or that has a component outside model_parts() + # (SD's depth_estimator, Anima's text_conditioner) overrides this. + for part in self.model_type.model_parts(): + stem = f"{part}_1" if hasattr(self, f"{part}_1") else part + component = getattr(self, stem) + if component is not None: + component.eval() - @abstractmethod def adapters(self) -> list[LoRAModuleWrapper]: - pass + # Every LoRA adapter present on a model part, in model_parts() order. Parts without a LoRA + # (e.g. the vae, or an untrained component) contribute nothing. + result = [] + for part in self.model_type.model_parts(): + stem = f"{part}_1" if hasattr(self, f"{part}_1") else part + lora = getattr(self, f"{stem}_lora", None) + if lora is not None: + result.append(lora) + return result def diffusers_to_original(self) -> list | None: # the canonical(diffusers) -> native key-conversion BODY (rename only) for this model's denoising diff --git a/modules/model/ChromaModel.py b/modules/model/ChromaModel.py index 64279f982..4fbcf430b 100644 --- a/modules/model/ChromaModel.py +++ b/modules/model/ChromaModel.py @@ -95,12 +95,6 @@ def __init__( self.transformer_lora = None self.lora_state_dict = None - def adapters(self) -> list[LoRAModuleWrapper]: - return [a for a in [ - self.text_encoder_lora, - self.transformer_lora, - ] if a is not None] - def fusion_groups(self) -> list | None: # Chroma is Flux-structured: double blocks fuse img q/k/v and txt q/k/v separately; single blocks # fuse q/k/v + mlp into one linear1. @@ -164,37 +158,6 @@ def all_text_encoder_embeddings(self) -> list[BaseModelEmbedding]: return [embedding.text_encoder_embedding for embedding in self.additional_embeddings] \ + ([self.embedding.text_encoder_embedding] if self.embedding is not None else []) - 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: - self.text_encoder_offload_conductor.to(device) - else: - self.text_encoder.to(device=device) - - if self.text_encoder_lora is not None: - self.text_encoder_lora.to(device) - - def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None: - self.transformer_offload_conductor.to(device) - else: - self.transformer.to(device=device) - - if self.transformer_lora is not None: - self.transformer_lora.to(device) - - def to(self, device: torch.device): - self.vae_to(device) - self.text_encoder_to(device) - self.transformer_to(device) - - def eval(self): - self.vae.eval() - self.text_encoder.eval() - self.transformer.eval() - def create_pipeline(self, use_original_tokenizers: bool = False) -> DiffusionPipeline: return ChromaPipeline( transformer=self.transformer, diff --git a/modules/model/ErnieModel.py b/modules/model/ErnieModel.py index 3d15e5385..42bb1be74 100644 --- a/modules/model/ErnieModel.py +++ b/modules/model/ErnieModel.py @@ -62,41 +62,6 @@ def __init__( self.transformer_lora = None self.lora_state_dict = None - def adapters(self) -> list[LoRAModuleWrapper]: - return [a for a in [ - self.transformer_lora, - ] if a is not None] - - def vae_to(self, device: torch.device): - self.vae.to(device=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: - 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: - self.transformer_offload_conductor.to(device) - else: - self.transformer.to(device=device) - - if self.transformer_lora is not None: - self.transformer_lora.to(device) - - def to(self, device: torch.device): - self.vae_to(device) - self.text_encoder_to(device) - self.transformer_to(device) - - def eval(self): - self.vae.eval() - if self.text_encoder is not None: - self.text_encoder.eval() - self.transformer.eval() - def create_pipeline(self) -> DiffusionPipeline: return ErnieImagePipeline( transformer=self.transformer, diff --git a/modules/model/Flux2Model.py b/modules/model/Flux2Model.py index 15c64bcc3..79eb02006 100644 --- a/modules/model/Flux2Model.py +++ b/modules/model/Flux2Model.py @@ -71,11 +71,6 @@ def __init__( self.transformer_lora = None self.lora_state_dict = None - def adapters(self) -> list[LoRAModuleWrapper]: - return [a for a in [ - self.transformer_lora, - ] if a is not None] - def fusion_groups(self) -> list | None: # Only the two double-block qkv groups -- Flux2's single block is the already-fused attn.to_qkv_mlp_proj. # NOTE: the fused suffix (attn.qkv / attn.added_qkv) is OneTrainer's name for the fused module on the @@ -126,36 +121,6 @@ def diffusers_to_original(self) -> list | None: ]), ] - def vae_to(self, device: torch.device): - self.vae.to(device=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: - 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: - self.transformer_offload_conductor.to(device) - else: - self.transformer.to(device=device) - - if self.transformer_lora is not None: - self.transformer_lora.to(device) - - def to(self, device: torch.device): - self.vae_to(device) - self.text_encoder_to(device) - self.transformer_to(device) - - def eval(self): - self.vae.eval() - if self.text_encoder is not None: - self.text_encoder.eval() - self.transformer.eval() - def create_pipeline(self) -> DiffusionPipeline: klass = Flux2Pipeline if self.is_dev() else Flux2KleinPipeline return klass( diff --git a/modules/model/FluxModel.py b/modules/model/FluxModel.py index 02631c3a3..5236b3498 100644 --- a/modules/model/FluxModel.py +++ b/modules/model/FluxModel.py @@ -116,13 +116,6 @@ def __init__( self.transformer_lora = None self.lora_state_dict = None - def adapters(self) -> list[LoRAModuleWrapper]: - return [a for a in [ - self.text_encoder_1_lora, - self.text_encoder_2_lora, - self.transformer_lora, - ] if a is not None] - def fusion_groups(self) -> list | None: # Flux fuses img qkv, txt qkv, and the single-block qkv+mlp (4 leaves -> linear1). return [ @@ -206,52 +199,6 @@ def all_text_encoder_2_embeddings(self) -> list[BaseModelEmbedding]: return [embedding.text_encoder_2_embedding for embedding in self.additional_embeddings] \ + ([self.embedding.text_encoder_2_embedding] if self.embedding is not None else []) - def vae_to(self, device: torch.device): - self.vae.to(device=device) - - def text_encoder_to(self, device: torch.device): - self.text_encoder_1_to(device=device) - self.text_encoder_2_to(device=device) - - def text_encoder_1_to(self, device: torch.device): - if self.text_encoder_1 is not None: - self.text_encoder_1.to(device=device) - - if self.text_encoder_1_lora is not None: - self.text_encoder_1_lora.to(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: - self.text_encoder_2_offload_conductor.to(device) - else: - self.text_encoder_2.to(device=device) - - if self.text_encoder_2_lora is not None: - self.text_encoder_2_lora.to(device) - - def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None: - self.transformer_offload_conductor.to(device) - else: - self.transformer.to(device=device) - - if self.transformer_lora is not None: - self.transformer_lora.to(device) - - def to(self, device: torch.device): - self.vae_to(device) - self.text_encoder_to(device) - self.transformer_to(device) - - def eval(self): - self.vae.eval() - if self.text_encoder_1 is not None: - self.text_encoder_1.eval() - if self.text_encoder_2 is not None: - self.text_encoder_2.eval() - self.transformer.eval() - def create_pipeline(self, use_original_tokenizers: bool = False) -> DiffusionPipeline: return FluxPipeline( transformer=self.transformer, diff --git a/modules/model/HiDreamModel.py b/modules/model/HiDreamModel.py index a68feaf9d..f99b5116b 100644 --- a/modules/model/HiDreamModel.py +++ b/modules/model/HiDreamModel.py @@ -167,15 +167,6 @@ def __init__( self.transformer_lora = None self.lora_state_dict = None - def adapters(self) -> list[LoRAModuleWrapper]: - return [a for a in [ - self.text_encoder_1_lora, - self.text_encoder_2_lora, - self.text_encoder_3_lora, - self.text_encoder_4_lora, - self.transformer_lora, - ] if a is not None] - def lora_text_encoders(self) -> list[tuple[torch.nn.Module | None, dict[ModelFormat, str]]]: # HiDream's four TEs: clip_l + clip_g + t5xxl + llama (Comfy's HiDreamTEModel). Any can be absent, so # only the TEs actually present are declared. @@ -226,75 +217,6 @@ def all_text_encoder_4_embeddings(self) -> list[BaseModelEmbedding]: return [embedding.text_encoder_4_embedding for embedding in self.additional_embeddings] \ + ([self.embedding.text_encoder_4_embedding] if self.embedding is not None else []) - def vae_to(self, device: torch.device): - self.vae.to(device=device) - - def text_encoder_to(self, device: torch.device): - self.text_encoder_1_to(device=device) - self.text_encoder_2_to(device=device) - self.text_encoder_3_to(device=device) - self.text_encoder_4_to(device=device) - - def text_encoder_1_to(self, device: torch.device): - if self.text_encoder_1 is not None: - self.text_encoder_1.to(device=device) - - if self.text_encoder_1_lora is not None: - self.text_encoder_1_lora.to(device) - - def text_encoder_2_to(self, device: torch.device): - if self.text_encoder_2 is not None: - self.text_encoder_2.to(device=device) - - if self.text_encoder_2_lora is not None: - self.text_encoder_2_lora.to(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: - self.text_encoder_3_offload_conductor.to(device) - else: - self.text_encoder_3.to(device=device) - - if self.text_encoder_3_lora is not None: - self.text_encoder_3_lora.to(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: - self.text_encoder_4_offload_conductor.to(device) - else: - self.text_encoder_4.to(device=device) - - if self.text_encoder_4_lora is not None: - self.text_encoder_4_lora.to(device) - - def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None: - self.transformer_offload_conductor.to(device) - else: - self.transformer.to(device=device) - - if self.transformer_lora is not None: - self.transformer_lora.to(device) - - def to(self, device: torch.device): - self.vae_to(device) - self.text_encoder_to(device) - self.transformer_to(device) - - def eval(self): - self.vae.eval() - if self.text_encoder_1 is not None: - self.text_encoder_1.eval() - if self.text_encoder_2 is not None: - self.text_encoder_2.eval() - if self.text_encoder_3 is not None: - self.text_encoder_3.eval() - if self.text_encoder_4 is not None: - self.text_encoder_4.eval() - self.transformer.eval() - def create_pipeline(self, use_original_tokenizers: bool = False) -> DiffusionPipeline: return HiDreamImagePipeline( transformer=self.transformer, diff --git a/modules/model/HunyuanVideoModel.py b/modules/model/HunyuanVideoModel.py index eb3e8271b..9ec6b4516 100644 --- a/modules/model/HunyuanVideoModel.py +++ b/modules/model/HunyuanVideoModel.py @@ -131,13 +131,6 @@ def __init__( self.transformer_lora = None self.lora_state_dict = None - def adapters(self) -> list[LoRAModuleWrapper]: - return [a for a in [ - self.text_encoder_1_lora, - self.text_encoder_2_lora, - self.transformer_lora, - ] if a is not None] - def fusion_groups(self) -> list | None: # HunyuanVideo fuses qkv in three places: the context-embedder token-refiner blocks, the double # blocks (img + txt), and the single blocks (qkv+mlp). @@ -230,52 +223,6 @@ def all_text_encoder_2_embeddings(self) -> list[BaseModelEmbedding]: return [embedding.text_encoder_2_embedding for embedding in self.additional_embeddings] \ + ([self.embedding.text_encoder_2_embedding] if self.embedding is not None else []) - def vae_to(self, device: torch.device): - self.vae.to(device=device) - - def text_encoder_to(self, device: torch.device): - self.text_encoder_1_to(device=device) - self.text_encoder_2_to(device=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: - self.text_encoder_1_offload_conductor.to(device) - else: - self.text_encoder_1.to(device=device) - - if self.text_encoder_1_lora is not None: - self.text_encoder_1_lora.to(device) - - def text_encoder_2_to(self, device: torch.device): - if self.text_encoder_2 is not None: - self.text_encoder_2.to(device=device) - - if self.text_encoder_2_lora is not None: - self.text_encoder_2_lora.to(device) - - def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None: - self.transformer_offload_conductor.to(device) - else: - self.transformer.to(device=device) - - if self.transformer_lora is not None: - self.transformer_lora.to(device) - - def to(self, device: torch.device): - self.vae_to(device) - self.text_encoder_to(device) - self.transformer_to(device) - - def eval(self): - self.vae.eval() - if self.text_encoder_1 is not None: - self.text_encoder_1.eval() - if self.text_encoder_2 is not None: - self.text_encoder_2.eval() - self.transformer.eval() - def create_pipeline(self, use_original_tokenizers: bool = False) -> DiffusionPipeline: return HunyuanVideoPipeline( transformer=self.transformer, diff --git a/modules/model/IdeogramModel.py b/modules/model/IdeogramModel.py index d388615f8..70fae8444 100644 --- a/modules/model/IdeogramModel.py +++ b/modules/model/IdeogramModel.py @@ -67,12 +67,6 @@ def __init__( self.transformer_lora = None self.lora_state_dict = None - def adapters(self) -> list[LoRAModuleWrapper]: - # only the conditional transformer is trainable; the unconditional transformer never sees the concept - return [a for a in [ - self.transformer_lora, - ] if a is not None] - def fusion_groups(self) -> list | None: # Ideogram4 fuses q/k/v into one qkv Linear per block; everything else in the transformer -- including # the output projection (to_out.0 -> o, see diffusers_to_original) -- already matches the original @@ -91,44 +85,6 @@ def diffusers_to_original(self) -> list | None: ("{path}", "{path}"), ] - 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: - 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: - self.transformer_offload_conductor.to(device) - else: - self.transformer.to(device=device) - - if self.transformer_lora is not None: - self.transformer_lora.to(device) - - def unconditional_transformer_to(self, device: torch.device): - if self.unconditional_transformer is not None: - if self.unconditional_transformer_offload_conductor is not None: - self.unconditional_transformer_offload_conductor.to(device) - else: - self.unconditional_transformer.to(device=device) - - def to(self, device: torch.device): - self.vae_to(device) - self.text_encoder_to(device) - self.transformer_to(device) - self.unconditional_transformer_to(device) - - def eval(self): - self.vae.eval() - self.text_encoder.eval() - self.transformer.eval() - if self.unconditional_transformer is not None: - self.unconditional_transformer.eval() - def create_pipeline(self) -> DiffusionPipeline: return Ideogram4Pipeline( transformer=self.transformer, diff --git a/modules/model/Krea2Model.py b/modules/model/Krea2Model.py index 8bd614848..cf5ca0d69 100644 --- a/modules/model/Krea2Model.py +++ b/modules/model/Krea2Model.py @@ -80,11 +80,6 @@ def __init__( self.transformer_lora = None self.lora_state_dict = None - def adapters(self) -> list[LoRAModuleWrapper]: - return [a for a in [ - self.transformer_lora, - ] if a is not None] - def diffusers_to_original(self) -> list | None: # Krea 2's native checkpoint (krea/Krea-2-Raw's raw.safetensors) is a pure rename of the diffusers # Krea2Transformer2DModel state dict -- q/k/v are already split in both namespaces, so no qkv fusion @@ -127,34 +122,6 @@ def table_mod(t): return t.reshape(6, -1) ]), ] - def vae_to(self, device: torch.device): - self.vae.to(device=device) - - def text_encoder_to(self, device: torch.device): #TODO share more code between models - 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: - self.transformer_offload_conductor.to(device) - else: - self.transformer.to(device=device) - - if self.transformer_lora is not None: - self.transformer_lora.to(device) - - def to(self, device: torch.device): - self.vae_to(device) - self.text_encoder_to(device) - self.transformer_to(device) - - def eval(self): - self.vae.eval() - self.text_encoder.eval() - self.transformer.eval() - def create_pipeline(self) -> DiffusionPipeline: return Krea2Pipeline( transformer=self.transformer, diff --git a/modules/model/PixArtAlphaModel.py b/modules/model/PixArtAlphaModel.py index d53989da9..bf36d1a21 100644 --- a/modules/model/PixArtAlphaModel.py +++ b/modules/model/PixArtAlphaModel.py @@ -97,12 +97,6 @@ def __init__( self.transformer_lora = None self.lora_state_dict = None - def adapters(self) -> list[LoRAModuleWrapper]: - return [a for a in [ - self.text_encoder_lora, - self.transformer_lora, - ] if a is not None] - def fusion_groups(self) -> list | None: # PixArt fuses TWO attentions differently: self-attention fuses q/k/v (3 leaves), while cross-attention # fuses ONLY k/v (2 leaves) into kv_linear. @@ -158,37 +152,6 @@ def all_text_encoder_embeddings(self) -> list[BaseModelEmbedding]: return [embedding.text_encoder_embedding for embedding in self.additional_embeddings] \ + ([self.embedding.text_encoder_embedding] if self.embedding is not None else []) - 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: - self.text_encoder_offload_conductor.to(device) - else: - self.text_encoder.to(device=device) - - if self.text_encoder_lora is not None: - self.text_encoder_lora.to(device) - - def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None: - self.transformer_offload_conductor.to(device) - else: - self.transformer.to(device=device) - - if self.transformer_lora is not None: - self.transformer_lora.to(device) - - def to(self, device: torch.device): - self.vae_to(device) - self.text_encoder_to(device) - self.transformer_to(device) - - def eval(self): - self.vae.eval() - self.text_encoder.eval() - self.transformer.eval() - def create_pipeline(self, use_original_tokenizers: bool = False) -> DiffusionPipeline: tokenizer = self.orig_tokenizer if use_original_tokenizers else self.tokenizer match self.model_type: diff --git a/modules/model/QwenModel.py b/modules/model/QwenModel.py index 775143063..3b69d8f8b 100644 --- a/modules/model/QwenModel.py +++ b/modules/model/QwenModel.py @@ -71,12 +71,6 @@ def __init__( self.transformer_lora = None self.lora_state_dict = None - def adapters(self) -> list[LoRAModuleWrapper]: - return [a for a in [ - self.text_encoder_lora, - self.transformer_lora, - ] if a is not None] - def lora_text_encoders(self) -> list[tuple[torch.nn.Module | None, dict[ModelFormat, str]]]: # Single Qwen2.5-VL TE (Comfy's QwenImageTEModel is a single qwen25_7b). return [ @@ -87,37 +81,6 @@ def lora_text_encoders(self) -> list[tuple[torch.nn.Module | None, dict[ModelFor }), ] - def vae_to(self, device: torch.device): - self.vae.to(device=device) - - def text_encoder_to(self, device: torch.device): #TODO share more code between models - if self.text_encoder_offload_conductor is not None: - self.text_encoder_offload_conductor.to(device) - else: - self.text_encoder.to(device=device) - - if self.text_encoder_lora is not None: - self.text_encoder_lora.to(device) - - def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None: - self.transformer_offload_conductor.to(device) - else: - self.transformer.to(device=device) - - if self.transformer_lora is not None: - self.transformer_lora.to(device) - - def to(self, device: torch.device): - self.vae_to(device) - self.text_encoder_to(device) - self.transformer_to(device) - - def eval(self): - self.vae.eval() - self.text_encoder.eval() - self.transformer.eval() - def create_pipeline(self) -> DiffusionPipeline: return QwenImagePipeline( transformer=self.transformer, diff --git a/modules/model/SanaModel.py b/modules/model/SanaModel.py index 75dae101c..f43e98ed5 100644 --- a/modules/model/SanaModel.py +++ b/modules/model/SanaModel.py @@ -99,12 +99,6 @@ def __init__( self.transformer_lora = None self.lora_state_dict = None - def adapters(self) -> list[LoRAModuleWrapper]: - return [a for a in [ - self.text_encoder_lora, - self.transformer_lora, - ] if a is not None] - def lora_text_encoders(self) -> list[tuple[torch.nn.Module | None, dict[ModelFormat, str]]]: # Single Gemma2 TE. No COMFY_LORA name -- ComfyUI cannot load Sana, so the COMFY format refuses to # write its TE keys. @@ -123,37 +117,6 @@ def all_text_encoder_embeddings(self) -> list[BaseModelEmbedding]: return [embedding.text_encoder_embedding for embedding in self.additional_embeddings] \ + ([self.embedding.text_encoder_embedding] if self.embedding is not None else []) - 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: - self.text_encoder_offload_conductor.to(device) - else: - self.text_encoder.to(device=device) - - if self.text_encoder_lora is not None: - self.text_encoder_lora.to(device) - - def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None: - self.transformer_offload_conductor.to(device) - else: - self.transformer.to(device=device) - - if self.transformer_lora is not None: - self.transformer_lora.to(device) - - def to(self, device: torch.device): - self.vae_to(device) - self.text_encoder_to(device) - self.transformer_to(device) - - def eval(self): - self.vae.eval() - self.text_encoder.eval() - self.transformer.eval() - def create_pipeline(self, use_original_tokenizers: bool = False) -> DiffusionPipeline: return SanaPipeline( tokenizer=self.orig_tokenizer if use_original_tokenizers else self.tokenizer, diff --git a/modules/model/StableDiffusion3Model.py b/modules/model/StableDiffusion3Model.py index 7de0c4125..e9324bb1b 100644 --- a/modules/model/StableDiffusion3Model.py +++ b/modules/model/StableDiffusion3Model.py @@ -134,14 +134,6 @@ def __init__( self.transformer_lora = None self.lora_state_dict = None - def adapters(self) -> list[LoRAModuleWrapper]: - return [a for a in [ - self.text_encoder_1_lora, - self.text_encoder_2_lora, - self.text_encoder_3_lora, - self.transformer_lora, - ] if a is not None] - def fusion_groups(self) -> list | None: # SD3 fuses TWO joint streams (x_block + context_block) plus a dual-attention attn2 that only exists # in some SD3.5 blocks (the group fires per-block only where all its leaves are present). @@ -231,62 +223,6 @@ def all_text_encoder_3_embeddings(self) -> list[BaseModelEmbedding]: return [embedding.text_encoder_3_embedding for embedding in self.additional_embeddings] \ + ([self.embedding.text_encoder_3_embedding] if self.embedding is not None else []) - def vae_to(self, device: torch.device): - self.vae.to(device=device) - - def text_encoder_to(self, device: torch.device): - self.text_encoder_1_to(device=device) - self.text_encoder_2_to(device=device) - self.text_encoder_3_to(device=device) - - def text_encoder_1_to(self, device: torch.device): - if self.text_encoder_1 is not None: - self.text_encoder_1.to(device=device) - - if self.text_encoder_1_lora is not None: - self.text_encoder_1_lora.to(device) - - def text_encoder_2_to(self, device: torch.device): - if self.text_encoder_2 is not None: - self.text_encoder_2.to(device=device) - - if self.text_encoder_2_lora is not None: - self.text_encoder_2_lora.to(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: - self.text_encoder_3_offload_conductor.to(device) - else: - self.text_encoder_3.to(device=device) - - if self.text_encoder_3_lora is not None: - self.text_encoder_3_lora.to(device) - - def transformer_to(self, device: torch.device): - if self.transformer_offload_conductor is not None: - self.transformer_offload_conductor.to(device) - else: - self.transformer.to(device=device) - - if self.transformer_lora is not None: - self.transformer_lora.to(device) - - def to(self, device: torch.device): - self.vae_to(device) - self.text_encoder_to(device) - self.transformer_to(device) - - def eval(self): - self.vae.eval() - if self.text_encoder_1 is not None: - self.text_encoder_1.eval() - if self.text_encoder_2 is not None: - self.text_encoder_2.eval() - if self.text_encoder_3 is not None: - self.text_encoder_3.eval() - self.transformer.eval() - def create_pipeline(self, use_original_tokenizers: bool = False) -> DiffusionPipeline: return StableDiffusion3Pipeline( transformer=self.transformer, diff --git a/modules/model/StableDiffusionModel.py b/modules/model/StableDiffusionModel.py index be3a8af8b..697d2420a 100644 --- a/modules/model/StableDiffusionModel.py +++ b/modules/model/StableDiffusionModel.py @@ -104,12 +104,6 @@ def __init__( self.sd_config = None self.sd_config_filename = None - def adapters(self) -> list[LoRAModuleWrapper]: - return [a for a in [ - self.text_encoder_lora, - self.unet_lora, - ] if a is not None] - def diffusers_to_original(self) -> list | None: # SD1.5/2.x UNet diffusers -> original/sgm key map, convert()-native (bare sgm names, no top prefix). # SD has NO qkv fusion and NO add_embedding, so this is a pure key rename. Spatial-transformer @@ -181,37 +175,11 @@ def all_text_encoder_embeddings(self) -> list[BaseModelEmbedding]: return [embedding.text_encoder_embedding for embedding in self.additional_embeddings] \ + ([self.embedding.text_encoder_embedding] if self.embedding is not None else []) - def vae_to(self, device: torch.device): - self.vae.to(device=device) - - def depth_estimator_to(self, device: torch.device): - if self.depth_estimator is not None: - self.depth_estimator.to(device=device) - - def text_encoder_to(self, device: torch.device): - self.text_encoder.to(device=device) - - if self.text_encoder_lora is not None: - self.text_encoder_lora.to(device) - - def unet_to(self, device: torch.device): - self.unet.to(device=device) - - if self.unet_lora is not None: - self.unet_lora.to(device) - - def to(self, device: torch.device): - self.vae_to(device) - self.depth_estimator_to(device) - self.text_encoder_to(device) - self.unet_to(device) - def eval(self): - self.vae.eval() + super().eval() + # depth_estimator isn't in ModelType.model_parts(); only the depth model variants have it. if self.depth_estimator is not None: self.depth_estimator.eval() - self.text_encoder.eval() - self.unet.eval() def create_pipeline(self, use_original_tokenizers: bool = False) -> DiffusionPipeline: tokenizer = self.orig_tokenizer if use_original_tokenizers else self.tokenizer diff --git a/modules/model/StableDiffusionXLModel.py b/modules/model/StableDiffusionXLModel.py index 79af4e55f..7f2881d70 100644 --- a/modules/model/StableDiffusionXLModel.py +++ b/modules/model/StableDiffusionXLModel.py @@ -121,13 +121,6 @@ def __init__( self.sd_config = None self.sd_config_filename = None - def adapters(self) -> list[LoRAModuleWrapper]: - return [a for a in [ - self.text_encoder_1_lora, - self.text_encoder_2_lora, - self.unet_lora, - ] if a is not None] - def diffusers_to_original(self) -> list | None: # SDXL UNet diffusers -> original/sgm key map, convert()-native (bare sgm names, no top prefix). # SDXL has NO qkv fusion, so this is a pure key rename. Spatial-transformer (attention) blocks have @@ -199,48 +192,6 @@ def all_text_encoder_2_embeddings(self) -> list[BaseModelEmbedding]: return [embedding.text_encoder_2_embedding for embedding in self.additional_embeddings] \ + ([self.embedding.text_encoder_2_embedding] if self.embedding is not None else []) - def vae_to(self, device: torch.device): - self.vae.to(device=device) - - def text_encoder_to(self, device: torch.device): - self.text_encoder_1.to(device=device) - self.text_encoder_2.to(device=device) - - if self.text_encoder_1_lora is not None: - self.text_encoder_1_lora.to(device) - - if self.text_encoder_2_lora is not None: - self.text_encoder_2_lora.to(device) - - def text_encoder_1_to(self, device: torch.device): - self.text_encoder_1.to(device=device) - - if self.text_encoder_1_lora is not None: - self.text_encoder_1_lora.to(device) - - def text_encoder_2_to(self, device: torch.device): - self.text_encoder_2.to(device=device) - - if self.text_encoder_2_lora is not None: - self.text_encoder_2_lora.to(device) - - def unet_to(self, device: torch.device): - self.unet.to(device=device) - - if self.unet_lora is not None: - self.unet_lora.to(device) - - def to(self, device: torch.device): - self.vae_to(device) - self.text_encoder_to(device) - self.unet_to(device) - - def eval(self): - self.vae.eval() - self.text_encoder_1.eval() - self.text_encoder_2.eval() - self.unet.eval() - def create_pipeline(self, use_original_tokenizers: bool = False) -> DiffusionPipeline: return StableDiffusionXLPipeline( vae=self.vae, diff --git a/modules/model/WuerstchenModel.py b/modules/model/WuerstchenModel.py index 5268a9c16..f04c5b8c7 100644 --- a/modules/model/WuerstchenModel.py +++ b/modules/model/WuerstchenModel.py @@ -174,38 +174,27 @@ def all_prior_text_encoder_embeddings(self) -> list[BaseModelEmbedding]: return [embedding.text_encoder_embedding for embedding in self.additional_embeddings] \ + ([self.embedding.prior_text_encoder_embedding] if self.embedding is not None else []) - def decoder_text_encoder_to(self, device: torch.device): - self.decoder_text_encoder.to(device=device) - - def decoder_decoder_to(self, device: torch.device): - self.decoder_decoder.to(device=device) - - def decoder_vqgan_to(self, device: torch.device): - self.decoder_vqgan.to(device=device) - - def effnet_encoder_to(self, device: torch.device): - self.effnet_encoder.to(device=device) - - def prior_text_encoder_to(self, device: torch.device): - self.prior_text_encoder.to(device=device) - - if self.prior_text_encoder_lora is not None: - self.prior_text_encoder_lora.to(device) - - def prior_prior_to(self, device: torch.device): - self.prior_prior.to(device=device) - - if self.prior_prior_lora is not None: - self.prior_prior_lora.to(device) - - def to(self, device: torch.device): - if self.model_type.is_wuerstchen_v2(): - self.decoder_text_encoder_to(device) - self.decoder_decoder_to(device) - self.decoder_vqgan_to(device) - self.effnet_encoder_to(device) - self.prior_text_encoder_to(device) - self.prior_prior_to(device) + def materialize(self, *parts: str): + super().materialize(*self._translate_parts(parts)) + + def evict(self, *parts: str): + # evict() with no parts means "evict all"; translate against the model's own full part list. + super().evict(*self._translate_parts(parts or self.model_type.model_parts())) + + def _translate_parts(self, parts: tuple[str, ...]) -> tuple[str, ...]: + # The prior stage's own diffusion module and the decoder stage's own diffusion module are each + # named after their stage, and the main text encoder belongs to the prior stage. + translated = [] + for part in parts: + if part == "prior": + translated.append("prior_prior") + elif part == "decoder": + translated.append("decoder_decoder") + elif part == "text_encoder": + translated.append("prior_text_encoder") + else: + translated.append(part) + return tuple(translated) def eval(self): if self.model_type.is_wuerstchen_v2(): diff --git a/modules/model/ZImageModel.py b/modules/model/ZImageModel.py index ba8511dce..8b60f2c28 100644 --- a/modules/model/ZImageModel.py +++ b/modules/model/ZImageModel.py @@ -74,11 +74,6 @@ def __init__( self.transformer_lora = None self.lora_state_dict = None - def adapters(self) -> list[LoRAModuleWrapper]: - return [a for a in [ - self.transformer_lora, - ] if a is not None] - def checkpoint_diffusers_to_comfy(self) -> list | None: # Full-model COMFY_TRANSFORMER conversion: Z-Image is the one model whose Comfy checkpoint layout # diverges from diffusers/original (ComfyUI #12303). Only these keys change -- everything else passes @@ -97,36 +92,6 @@ def checkpoint_diffusers_to_comfy(self) -> list | None: ("{p}.attention.to_out.0", "{p}.attention.out"), ] - def vae_to(self, device: torch.device): - self.vae.to(device=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: - 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: - self.transformer_offload_conductor.to(device) - else: - self.transformer.to(device=device) - - if self.transformer_lora is not None: - self.transformer_lora.to(device) - - def to(self, device: torch.device): - self.vae_to(device) - self.text_encoder_to(device) - self.transformer_to(device) - - def eval(self): - self.vae.eval() - if self.text_encoder is not None: - self.text_encoder.eval() - self.transformer.eval() - def create_pipeline(self) -> DiffusionPipeline: return ZImagePipeline( transformer=self.transformer, diff --git a/modules/modelSampler/AnimaSampler.py b/modules/modelSampler/AnimaSampler.py index 2c7322816..d5de18ac1 100644 --- a/modules/modelSampler/AnimaSampler.py +++ b/modules/modelSampler/AnimaSampler.py @@ -12,7 +12,6 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat -from modules.util.torch_util import torch_gc import torch @@ -66,7 +65,7 @@ def __sample_base( num_latent_channels = 16 # prepare prompt - self.model.text_encoder_to(self.train_device) + self.model.materialize_only("text_encoder") batch_size = 2 if cfg_scale > 1.0 else 1 combined_prompt_embedding = self.model.encode_text( @@ -75,9 +74,6 @@ def __sample_base( train_device=self.train_device, ) - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare latent image latent_image = torch.randn( size=(1, num_latent_channels, 1, height // vae_scale_factor, width // vae_scale_factor), @@ -99,7 +95,7 @@ def __sample_base( if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): extra_step_kwargs["generator"] = generator - self.model.transformer_to(self.train_device) + self.model.materialize_only("transformer") for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): latent_model_input = torch.cat([latent_image] * batch_size) expanded_timestep = timestep.expand(batch_size) / noise_scheduler.config.num_train_timesteps @@ -121,11 +117,8 @@ def __sample_base( on_update_progress(i + 1, len(timesteps)) - self.model.transformer_to(self.temp_device) - torch_gc() - # decode - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") latents = self.model.unscale_latents(latent_image) image = vae.decode(latents, return_dict=False)[0][:, :, 0] @@ -133,9 +126,6 @@ def __sample_base( do_denormalize = [True] * image.shape[0] image = self.image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], diff --git a/modules/modelSampler/ChromaSampler.py b/modules/modelSampler/ChromaSampler.py index bc89e5cab..23ef4aee5 100644 --- a/modules/modelSampler/ChromaSampler.py +++ b/modules/modelSampler/ChromaSampler.py @@ -12,7 +12,6 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat -from modules.util.torch_util import torch_gc import torch @@ -64,7 +63,7 @@ def __sample_base( num_latent_channels = 16 # prepare prompt - self.model.text_encoder_to(self.train_device) + self.model.materialize_only("text_encoder") combined_prompt_embedding, text_attention_mask = self.model.encode_text( text=[prompt, negative_prompt], @@ -73,9 +72,6 @@ def __sample_base( text_encoder_layer_skip=text_encoder_layer_skip, ) - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare latent image latent_image = torch.randn( size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), @@ -109,7 +105,7 @@ def __sample_base( image_attention_mask = torch.full((2, image_seq_len), True, dtype=torch.bool, device=text_attention_mask.device) attention_mask = torch.cat([text_attention_mask, image_attention_mask], dim=1) - self.model.transformer_to(self.train_device) + self.model.materialize_only("transformer") for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): latent_model_input = torch.cat([latent_image] * 2) expanded_timestep = timestep.expand(2) @@ -134,9 +130,6 @@ def __sample_base( on_update_progress(i + 1, len(timesteps)) - self.model.transformer_to(self.temp_device) - torch_gc() - latent_image = self.model.unpack_latents( latent_image, height // vae_scale_factor, @@ -144,7 +137,7 @@ def __sample_base( ) # decode - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") latents = (latent_image / vae.config.scaling_factor) + vae.config.shift_factor image = vae.decode(latents, return_dict=False)[0] @@ -152,9 +145,6 @@ def __sample_base( do_denormalize = [True] * image.shape[0] image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], diff --git a/modules/modelSampler/ErnieSampler.py b/modules/modelSampler/ErnieSampler.py index 2122e5343..a9cb57e0a 100644 --- a/modules/modelSampler/ErnieSampler.py +++ b/modules/modelSampler/ErnieSampler.py @@ -11,7 +11,6 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat -from modules.util.torch_util import torch_gc import torch @@ -63,7 +62,7 @@ def __sample_base( num_latent_channels = 32 # encode text - self.model.text_encoder_to(self.train_device) + self.model.materialize_only("text_encoder") batch_size = 2 if cfg_scale > 1.0 else 1 text_bth, text_lens = self.model.encode_text( @@ -72,9 +71,6 @@ def __sample_base( ) dtype = self.model.train_dtype.torch_dtype() - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare latents latent_image = torch.randn( size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), @@ -88,7 +84,7 @@ def __sample_base( noise_scheduler.set_timesteps(sigmas=sigmas, device=self.train_device) timesteps = noise_scheduler.timesteps - self.model.transformer_to(self.train_device) + self.model.materialize_only("transformer") transformer = self.pipeline.transformer for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): @@ -112,9 +108,7 @@ def __sample_base( on_update_progress(i + 1, len(timesteps)) - self.model.transformer_to(self.temp_device) - torch_gc() - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") # unscale and unpatchify latents = self.model.unscale_latents(latent_image) @@ -126,9 +120,6 @@ def __sample_base( image = image.cpu().permute(0, 2, 3, 1).float().numpy() image = [PILImage.fromarray((img * 255).astype(np.uint8)) for img in image] - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], diff --git a/modules/modelSampler/Flux2Sampler.py b/modules/modelSampler/Flux2Sampler.py index 0a4cca9b9..7ecbd5c83 100644 --- a/modules/modelSampler/Flux2Sampler.py +++ b/modules/modelSampler/Flux2Sampler.py @@ -12,7 +12,6 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat -from modules.util.torch_util import torch_gc import torch @@ -69,7 +68,7 @@ def __sample_base( patch_size = 2 # prepare prompt - self.model.text_encoder_to(self.train_device) + self.model.materialize_only("text_encoder") batch_size = 2 if cfg_scale > 1.0 and not transformer.config.guidance_embeds else 1 prompt_embedding = self.model.encode_text( @@ -78,9 +77,6 @@ def __sample_base( text_encoder_sequence_length=text_encoder_sequence_length, ) - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare latent image latent_image = torch.randn( size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), @@ -109,15 +105,13 @@ def __sample_base( text_ids = self.model.prepare_text_ids(prompt_embedding) - - self.model.transformer_to(self.train_device) + self.model.materialize_only("transformer") guidance = (torch.tensor([cfg_scale], device=self.train_device, dtype=self.model.train_dtype.torch_dtype()) if transformer.config.guidance_embeds else None) for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): latent_model_input = torch.cat([latent_image] * batch_size) expanded_timestep = timestep.expand(latent_model_input.shape[0]) - noise_pred = transformer( hidden_states=latent_model_input.to(dtype=self.model.train_dtype.torch_dtype()), timestep=expanded_timestep / 1000, @@ -137,9 +131,7 @@ def __sample_base( on_update_progress(i + 1, len(timesteps)) - self.model.transformer_to(self.temp_device) - torch_gc() - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") latent_image = self.model.unpack_latents( latent_image, @@ -153,9 +145,6 @@ def __sample_base( image = image_processor.postprocess(image, output_type='pil') - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], diff --git a/modules/modelSampler/FluxSampler.py b/modules/modelSampler/FluxSampler.py index fbe532053..fc2b2e0c8 100644 --- a/modules/modelSampler/FluxSampler.py +++ b/modules/modelSampler/FluxSampler.py @@ -14,7 +14,6 @@ from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat from modules.util.image_util import load_image -from modules.util.torch_util import torch_gc import torch from torch import nn @@ -72,7 +71,7 @@ def __sample_base( num_latent_channels = 16 # prepare prompt - self.model.text_encoder_to(self.train_device) + self.model.materialize_only_text_encoders() prompt_embedding, pooled_prompt_embedding = self.model.encode_text( text=prompt, @@ -83,9 +82,6 @@ def __sample_base( apply_attention_mask=transformer_attention_mask, ) - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare latent image latent_image = torch.randn( size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), @@ -115,7 +111,7 @@ def __sample_base( text_ids = torch.zeros(prompt_embedding.shape[1], 3, device=self.train_device) - self.model.transformer_to(self.train_device) + self.model.materialize_only("transformer") for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): latent_model_input = torch.cat([latent_image]) expanded_timestep = timestep.expand(latent_model_input.shape[0]) @@ -147,8 +143,6 @@ def __sample_base( on_update_progress(i + 1, len(timesteps)) - self.model.transformer_to(self.temp_device) - torch_gc() latent_image = self.model.unpack_latents( latent_image, height // vae_scale_factor, @@ -156,7 +150,7 @@ def __sample_base( ) # decode - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") latents = (latent_image / vae.config.scaling_factor) + vae.config.shift_factor image = vae.decode(latents, return_dict=False)[0] @@ -164,9 +158,6 @@ def __sample_base( do_denormalize = [True] * image.shape[0] #TODO remove and test, from Flux and other models. True is the default image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], @@ -222,7 +213,7 @@ def __sample_inpainting( num_latent_channels = 16 # prepare conditioning image - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") if sample_inpainting: t = transforms.Compose([ @@ -296,7 +287,7 @@ def __sample_inpainting( ) # prepare prompt - self.model.text_encoder_to(self.train_device) + self.model.materialize_only_text_encoders() prompt_embedding, pooled_prompt_embedding = self.model.encode_text( text=prompt, @@ -307,9 +298,6 @@ def __sample_inpainting( apply_attention_mask=transformer_attention_mask, ) - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare latent image latent_image = torch.randn( size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), @@ -337,7 +325,7 @@ def __sample_inpainting( text_ids = torch.zeros(prompt_embedding.shape[1], 3, device=self.train_device) - self.model.transformer_to(self.train_device) + self.model.materialize_only("transformer") for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): latent_model_input = torch.cat([latent_image]) latent_model_input = torch.concat( @@ -372,9 +360,6 @@ def __sample_inpainting( on_update_progress(i + 1, len(timesteps)) - self.model.transformer_to(self.temp_device) - torch_gc() - latent_image = self.model.unpack_latents( latent_image, height // vae_scale_factor, @@ -382,7 +367,7 @@ def __sample_inpainting( ) # decode - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") latents = (latent_image / vae.config.scaling_factor) + vae.config.shift_factor image = vae.decode(latents, return_dict=False)[0] @@ -390,9 +375,6 @@ def __sample_inpainting( do_denormalize = [True] * image.shape[0] image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], diff --git a/modules/modelSampler/HiDreamSampler.py b/modules/modelSampler/HiDreamSampler.py index ee7569ccf..c5b30723b 100644 --- a/modules/modelSampler/HiDreamSampler.py +++ b/modules/modelSampler/HiDreamSampler.py @@ -12,7 +12,6 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat -from modules.util.torch_util import torch_gc import torch @@ -65,7 +64,7 @@ def __sample_base( num_latent_channels = 16 # prepare prompt - self.model.text_encoder_to(self.train_device) + self.model.materialize_only_text_encoders() text_encoder_3_prompt_embedding, text_encoder_4_prompt_embedding, pooled_prompt_embedding = \ self.model.combine_text_encoder_output( @@ -92,9 +91,6 @@ def __sample_base( combined_pooled_prompt_embedding = torch.cat( [negative_pooled_prompt_embedding, pooled_prompt_embedding], dim=0) - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare latent image latent_image = torch.randn( size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), @@ -111,7 +107,7 @@ def __sample_base( if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): extra_step_kwargs["generator"] = generator - self.model.transformer_to(self.train_device) + self.model.materialize_only("transformer") for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): latent_model_input = torch.cat([latent_image] * 2) expanded_timestep = timestep.expand(latent_model_input.shape[0]) @@ -142,11 +138,8 @@ def __sample_base( on_update_progress(i + 1, len(timesteps)) - self.model.transformer_to(self.temp_device) - torch_gc() - # decode - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") latents = (latent_image / vae.config.scaling_factor) + vae.config.shift_factor image = vae.decode(latents, return_dict=False)[0] @@ -154,9 +147,6 @@ def __sample_base( do_denormalize = [True] * image.shape[0] image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], diff --git a/modules/modelSampler/HunyuanVideoSampler.py b/modules/modelSampler/HunyuanVideoSampler.py index c12056c92..10b22bfc9 100644 --- a/modules/modelSampler/HunyuanVideoSampler.py +++ b/modules/modelSampler/HunyuanVideoSampler.py @@ -12,7 +12,6 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat -from modules.util.torch_util import torch_gc import torch @@ -69,7 +68,7 @@ def __sample_base( num_latent_channels = 16 # prepare prompt - self.model.text_encoder_to(self.train_device) + self.model.materialize_only_text_encoders() prompt_embedding, pooled_prompt_embedding, prompt_attention_mask = self.model.encode_text( text=prompt, @@ -78,9 +77,6 @@ def __sample_base( text_encoder_2_layer_skip=text_encoder_2_layer_skip, ) - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare latent image num_latent_frames = (num_frames - 1) // vae_temporal_scale_factor + 1 latent_image = torch.randn( @@ -108,7 +104,7 @@ def __sample_base( if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): extra_step_kwargs["generator"] = generator - self.model.transformer_to(self.train_device) + self.model.materialize_only("transformer") for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): latent_model_input = torch.cat([latent_image]) expanded_timestep = timestep.expand(latent_model_input.shape[0]) @@ -139,20 +135,14 @@ def __sample_base( on_update_progress(i + 1, len(timesteps)) - self.model.transformer_to(self.temp_device) - torch_gc() - # decode - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") latents = latent_image / vae.config.scaling_factor image = vae.decode(latents, return_dict=False)[0] image = video_processor.postprocess(image, output_type='pt') - self.model.vae_to(self.temp_device) - torch_gc() - is_image = image.shape[2] == 1 if is_image: image = image.view((image.shape[0], image.shape[1], image.shape[3], image.shape[4])) diff --git a/modules/modelSampler/IdeogramSampler.py b/modules/modelSampler/IdeogramSampler.py index 0a5dd420b..cb253db0e 100644 --- a/modules/modelSampler/IdeogramSampler.py +++ b/modules/modelSampler/IdeogramSampler.py @@ -11,7 +11,6 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat -from modules.util.torch_util import torch_gc import torch @@ -92,7 +91,7 @@ def pack_conditioning(text_features: torch.Tensor, text_lengths: torch.Tensor) - return max_text_tokens, position_ids, segment_ids, indicator, llm_features, text_z_padding # encode text (conditional branch, and the empty-prompt negative branch if needed) - self.model.text_encoder_to(self.train_device) + self.model.materialize_only("text_encoder") text_features, text_lengths = self.model.encode_text( train_device=self.train_device, text=prompt, @@ -111,8 +110,6 @@ def pack_conditioning(text_features: torch.Tensor, text_lengths: torch.Tensor) - neg_text_z_padding, ) = pack_conditioning(neg_text_features, neg_text_lengths) del neg_text_features - self.model.text_encoder_to(self.temp_device) - torch_gc() if use_unconditional_transformer: # unconditional (image-only) branch: zeroed text features over the image-region slices of the layout @@ -141,9 +138,8 @@ def pack_conditioning(text_features: torch.Tensor, text_lengths: torch.Tensor) - timesteps = noise_scheduler.timesteps num_train_timesteps = noise_scheduler.config.num_train_timesteps - self.model.transformer_to(self.train_device) - if use_unconditional_transformer: - self.model.unconditional_transformer_to(self.train_device) + transformer_parts = ("transformer", "unconditional_transformer") if use_unconditional_transformer else ("transformer",) + self.model.materialize_only(*transformer_parts) for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): # scheduler stores num_train_timesteps-scaled timesteps; convert back to model time (0=noise, 1=data) @@ -190,10 +186,7 @@ def pack_conditioning(text_features: torch.Tensor, text_lengths: torch.Tensor) - on_update_progress(i + 1, len(timesteps)) - self.model.transformer_to(self.temp_device) - self.model.unconditional_transformer_to(self.temp_device) - torch_gc() - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") # bn-denormalize the packed latents and unpatchify back to (B, C, H, W) before VAE decode latents = self.model.unscale_latents(latent_image) @@ -205,9 +198,6 @@ def pack_conditioning(text_features: torch.Tensor, text_lengths: torch.Tensor) - image = image.cpu().permute(0, 2, 3, 1).float().numpy() image = [PILImage.fromarray((img * 255).astype(np.uint8)) for img in image] - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], diff --git a/modules/modelSampler/Krea2Sampler.py b/modules/modelSampler/Krea2Sampler.py index 20aecd5b1..b83205a93 100644 --- a/modules/modelSampler/Krea2Sampler.py +++ b/modules/modelSampler/Krea2Sampler.py @@ -13,7 +13,6 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat -from modules.util.torch_util import torch_gc import torch @@ -67,7 +66,7 @@ def __sample_base( num_latent_channels = 16 # prepare prompt - self.model.text_encoder_to(self.train_device) + self.model.materialize_only("text_encoder") batch_size = 2 if cfg_scale > 1.0 else 1 combined_prompt_embedding, text_attention_mask = self.model.encode_text( @@ -76,9 +75,6 @@ def __sample_base( train_device=self.train_device, ) - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare latent image latent_image = torch.randn( size=(1, num_latent_channels, 1, height // vae_scale_factor, width // vae_scale_factor), @@ -104,7 +100,7 @@ def __sample_base( if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): extra_step_kwargs["generator"] = generator - self.model.transformer_to(self.train_device) + self.model.materialize_only("transformer") for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): latent_model_input = torch.cat([latent_image] * batch_size) expanded_timestep = timestep.expand(batch_size) @@ -125,16 +121,13 @@ def __sample_base( on_update_progress(i + 1, len(timesteps)) - self.model.transformer_to(self.temp_device) - - torch_gc() latent_image = self.model.unpack_latents( latent_image, height // vae_scale_factor, width // vae_scale_factor, ) - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") latents = self.model.unscale_latents(latent_image) image = vae.decode(latents, return_dict=False)[0].squeeze(-3) @@ -142,9 +135,6 @@ def __sample_base( do_denormalize = [True] * image.shape[0] image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], diff --git a/modules/modelSampler/PixArtAlphaSampler.py b/modules/modelSampler/PixArtAlphaSampler.py index 9adddf3f2..f8c38f135 100644 --- a/modules/modelSampler/PixArtAlphaSampler.py +++ b/modules/modelSampler/PixArtAlphaSampler.py @@ -11,7 +11,6 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat -from modules.util.torch_util import torch_gc import torch @@ -63,7 +62,7 @@ def __sample_base( vae_scale_factor = self.pipeline.vae_scale_factor # prepare prompt - self.model.text_encoder_to(self.train_device) + self.model.materialize_only("text_encoder") prompt_embedding, tokens_attention_mask = self.model.encode_text( text=prompt, @@ -80,9 +79,6 @@ def __sample_base( combined_prompt_embedding = torch.cat([negative_prompt_embedding, prompt_embedding]) combined_prompt_attention_mask = torch.cat([negative_tokens_attention_mask, tokens_attention_mask]) - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare timesteps noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) timesteps = noise_scheduler.timesteps @@ -113,7 +109,7 @@ def __sample_base( added_cond_kwargs = {"resolution": resolution, "aspect_ratio": aspect_ratio} # denoising loop - self.model.transformer_to(self.train_device) + self.model.materialize_only("transformer") for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): latent_model_input = torch.cat([latent_image] * 2) latent_model_input = noise_scheduler.scale_model_input(latent_model_input, timestep) @@ -143,11 +139,8 @@ def __sample_base( on_update_progress(i + 1, len(timesteps)) - self.model.transformer_to(self.temp_device) - torch_gc() - # decode - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") latent_image = latent_image.to(dtype=vae.dtype) image = vae.decode(latent_image / vae.config.scaling_factor, return_dict=False)[0] @@ -155,9 +148,6 @@ def __sample_base( do_denormalize = [True] * image.shape[0] image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], diff --git a/modules/modelSampler/QwenSampler.py b/modules/modelSampler/QwenSampler.py index 4ca604102..c18eece7e 100644 --- a/modules/modelSampler/QwenSampler.py +++ b/modules/modelSampler/QwenSampler.py @@ -13,7 +13,6 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat -from modules.util.torch_util import torch_gc import torch @@ -65,7 +64,7 @@ def __sample_base( num_latent_channels = 16 # prepare prompt - self.model.text_encoder_to(self.train_device) + self.model.materialize_only("text_encoder") #unlike other models, Qwen benefits from CFG but is still quite good at CFG 1. Optimize for that: batch_size = 2 if cfg_scale > 1.0 else 1 @@ -75,9 +74,6 @@ def __sample_base( train_device=self.train_device, ) - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare latent image latent_image = torch.randn( size=(1, num_latent_channels, 1, height // vae_scale_factor, width // vae_scale_factor), @@ -110,7 +106,7 @@ def __sample_base( if torch.all(text_attention_mask): text_attention_mask = None - self.model.transformer_to(self.train_device) + self.model.materialize_only("transformer") for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): latent_model_input = torch.cat([latent_image] * batch_size) expanded_timestep = timestep.expand(batch_size) @@ -134,9 +130,6 @@ def __sample_base( on_update_progress(i + 1, len(timesteps)) - self.model.transformer_to(self.temp_device) - - torch_gc() latent_image = self.model.unpack_latents( latent_image, height // vae_scale_factor, @@ -144,7 +137,7 @@ def __sample_base( ) # decode - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") latents = self.model.unscale_latents(latent_image) image = vae.decode(latents, return_dict=False)[0].squeeze(-3) @@ -152,9 +145,6 @@ def __sample_base( do_denormalize = [True] * image.shape[0] image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], diff --git a/modules/modelSampler/SanaSampler.py b/modules/modelSampler/SanaSampler.py index 6251ab87e..f4089a222 100644 --- a/modules/modelSampler/SanaSampler.py +++ b/modules/modelSampler/SanaSampler.py @@ -12,7 +12,6 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat -from modules.util.torch_util import torch_gc import torch @@ -63,7 +62,7 @@ def __sample_base( vae_scale_factor = self.pipeline.vae_scale_factor # prepare prompt - self.model.text_encoder_to(self.train_device) + self.model.materialize_only("text_encoder") prompt_embedding, tokens_attention_mask = self.model.encode_text( text=prompt, @@ -80,9 +79,6 @@ def __sample_base( combined_prompt_embedding = torch.cat([negative_prompt_embedding, prompt_embedding]) combined_prompt_attention_mask = torch.cat([negative_tokens_attention_mask, tokens_attention_mask]) - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare timesteps noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) timesteps = noise_scheduler.timesteps @@ -102,7 +98,7 @@ def __sample_base( extra_step_kwargs["generator"] = generator # denoising loop - self.model.transformer_to(self.train_device) + self.model.materialize_only("transformer") for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): latent_model_input = torch.cat([latent_image] * 2) @@ -127,11 +123,8 @@ def __sample_base( on_update_progress(i + 1, len(timesteps)) - self.model.transformer_to(self.temp_device) - torch_gc() - # decode - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") latent_image = latent_image.to(dtype=vae.dtype) with self.model.vae_autocast_context: @@ -140,9 +133,6 @@ def __sample_base( do_denormalize = [True] * image.shape[0] image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], diff --git a/modules/modelSampler/StableDiffusion3Sampler.py b/modules/modelSampler/StableDiffusion3Sampler.py index fe159e914..f21f34627 100644 --- a/modules/modelSampler/StableDiffusion3Sampler.py +++ b/modules/modelSampler/StableDiffusion3Sampler.py @@ -12,7 +12,6 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat -from modules.util.torch_util import torch_gc import torch @@ -67,7 +66,7 @@ def __sample_base( vae_scale_factor = self.pipeline.vae_scale_factor # prepare prompt - self.model.text_encoder_to(self.train_device) + self.model.materialize_only_text_encoders() prompt_embedding, pooled_prompt_embedding = self.model.combine_text_encoder_output( *self.model.encode_text( @@ -93,9 +92,6 @@ def __sample_base( combined_pooled_prompt_embedding = torch.cat( [negative_pooled_prompt_embedding, pooled_prompt_embedding], dim=0) - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare timesteps noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) timesteps = noise_scheduler.timesteps @@ -114,7 +110,7 @@ def __sample_base( if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): extra_step_kwargs["generator"] = generator - self.model.transformer_to(self.train_device) + self.model.materialize_only("transformer") for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): latent_model_input = torch.cat([latent_image] * 2) expanded_timestep = timestep.expand(latent_model_input.shape[0]) @@ -140,11 +136,8 @@ def __sample_base( on_update_progress(i + 1, len(timesteps)) - self.model.transformer_to(self.temp_device) - torch_gc() - # decode - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") latents = (latent_image / vae.config.scaling_factor) + vae.config.shift_factor image = vae.decode(latents, return_dict=False)[0] @@ -152,9 +145,6 @@ def __sample_base( do_denormalize = [True] * image.shape[0] image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], diff --git a/modules/modelSampler/StableDiffusionSampler.py b/modules/modelSampler/StableDiffusionSampler.py index 92bfcd759..791290edc 100644 --- a/modules/modelSampler/StableDiffusionSampler.py +++ b/modules/modelSampler/StableDiffusionSampler.py @@ -12,7 +12,6 @@ from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat from modules.util.image_util import load_image -from modules.util.torch_util import torch_gc import torch from torch import nn @@ -74,7 +73,7 @@ def __sample_base( vae_scale_factor = self.pipeline.vae_scale_factor # prepare prompt - self.model.text_encoder_to(self.train_device) + self.model.materialize_only("text_encoder") prompt_embedding = self.model.encode_text( text=prompt, @@ -90,9 +89,6 @@ def __sample_base( combined_prompt_embedding = torch.cat([negative_prompt_embedding, prompt_embedding]) \ .to(dtype=self.model.train_dtype.torch_dtype()) - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare timesteps noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) timesteps = noise_scheduler.timesteps @@ -121,7 +117,7 @@ def __sample_base( extra_step_kwargs["generator"] = generator # denoising loop - self.model.unet_to(self.train_device) + self.model.materialize_only("unet") for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): latent_model_input = torch.cat([latent_image] * 2) latent_model_input = noise_scheduler.scale_model_input(latent_model_input, timestep) @@ -155,11 +151,8 @@ def __sample_base( on_update_progress(i + 1, len(timesteps)) - self.model.unet_to(self.temp_device) - torch_gc() - # decode - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") latent_image = latent_image.to(dtype=vae.dtype) image = vae.decode(latent_image / vae.config.scaling_factor, return_dict=False)[0] @@ -167,9 +160,6 @@ def __sample_base( do_denormalize = [True] * image.shape[0] image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], @@ -223,7 +213,7 @@ def __sample_inpainting( vae_scale_factor = self.pipeline.vae_scale_factor # prepare conditioning image - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") if sample_inpainting: t = transforms.Compose([ @@ -277,11 +267,8 @@ def __sample_inpainting( device=self.train_device ) - self.model.vae_to(self.temp_device) - torch_gc() - # prepare prompt - self.model.text_encoder_to(self.train_device) + self.model.materialize_only("text_encoder") prompt_embedding = self.model.encode_text( text=prompt, @@ -297,9 +284,6 @@ def __sample_inpainting( combined_prompt_embedding = torch.cat([negative_prompt_embedding, prompt_embedding]) \ .to(dtype=self.model.train_dtype.torch_dtype()) - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare timesteps noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) timesteps = noise_scheduler.timesteps @@ -328,7 +312,7 @@ def __sample_inpainting( extra_step_kwargs["generator"] = generator # denoising loop - self.model.unet_to(self.train_device) + self.model.materialize_only("unet") for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): latent_model_input = noise_scheduler.scale_model_input(latent_image, timestep) latent_model_input = torch.concat( @@ -365,11 +349,8 @@ def __sample_inpainting( on_update_progress(i + 1, len(timesteps)) - self.model.unet_to(self.temp_device) - torch_gc() - #decode - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") latent_image = latent_image.to(dtype=vae.dtype) image = vae.decode(latent_image / vae.config.scaling_factor, return_dict=False)[0] @@ -377,9 +358,6 @@ def __sample_inpainting( do_denormalize = [True] * image.shape[0] image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], diff --git a/modules/modelSampler/StableDiffusionVaeSampler.py b/modules/modelSampler/StableDiffusionVaeSampler.py index 254c73fdd..787b0cfdd 100644 --- a/modules/modelSampler/StableDiffusionVaeSampler.py +++ b/modules/modelSampler/StableDiffusionVaeSampler.py @@ -63,14 +63,12 @@ def sample( image_tensor = t_in(image).to(device=self.train_device, dtype=self.model.vae.dtype) image_tensor = image_tensor * 2 - 1 - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") with torch.no_grad(): latent_image_tensor = self.model.vae.encode(image_tensor.unsqueeze(0)).latent_dist.mean image_tensor = self.model.vae.decode(latent_image_tensor).sample.squeeze() - self.model.vae_to(self.temp_device) - image_tensor = (image_tensor + 1) * 0.5 image_tensor = image_tensor.clamp(0, 1) diff --git a/modules/modelSampler/StableDiffusionXLSampler.py b/modules/modelSampler/StableDiffusionXLSampler.py index 1f066268c..93d9f23d3 100644 --- a/modules/modelSampler/StableDiffusionXLSampler.py +++ b/modules/modelSampler/StableDiffusionXLSampler.py @@ -12,7 +12,6 @@ from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat from modules.util.image_util import load_image -from modules.util.torch_util import torch_gc import torch from torch import nn @@ -69,7 +68,7 @@ def __sample_base( vae_scale_factor = self.pipeline.vae_scale_factor # prepare prompt - self.model.text_encoder_to(self.train_device) + self.model.materialize_only_text_encoders() prompt_embedding, pooled_text_encoder_2_output = self.model.combine_text_encoder_output(*self.model.encode_text( text=prompt, @@ -88,9 +87,6 @@ def __sample_base( combined_prompt_embedding = torch.cat([negative_prompt_embedding, prompt_embedding]) \ .to(dtype=self.model.train_dtype.torch_dtype()) - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare timesteps noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) timesteps = noise_scheduler.timesteps @@ -144,7 +140,7 @@ def __sample_base( extra_step_kwargs["generator"] = generator # denoising loop - self.model.unet_to(self.train_device) + self.model.materialize_only("unet") for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): latent_model_input = torch.cat([latent_image] * 2) latent_model_input = noise_scheduler.scale_model_input(latent_model_input, timestep) @@ -177,11 +173,8 @@ def __sample_base( on_update_progress(i + 1, len(timesteps)) - self.model.unet_to(self.temp_device) - torch_gc() - # decode - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") latent_image = latent_image.to(dtype=self.model.vae_train_dtype.torch_dtype()) with self.model.vae_autocast_context: @@ -190,9 +183,6 @@ def __sample_base( do_denormalize = [True] * image.shape[0] image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], @@ -247,7 +237,7 @@ def __sample_inpainting( vae_scale_factor = self.pipeline.vae_scale_factor # prepare conditioning image - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") with self.model.vae_autocast_context: if sample_inpainting: @@ -303,11 +293,8 @@ def __sample_inpainting( device=self.train_device ) - self.model.vae_to(self.temp_device) - torch_gc() - # prepare prompt - self.model.text_encoder_to(self.train_device) + self.model.materialize_only_text_encoders() prompt_embedding, pooled_text_encoder_2_output = self.model.combine_text_encoder_output( *self.model.encode_text( @@ -328,9 +315,6 @@ def __sample_inpainting( combined_prompt_embedding = torch.cat([negative_prompt_embedding, prompt_embedding]) \ .to(dtype=self.model.train_dtype.torch_dtype()) - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare timesteps noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) timesteps = noise_scheduler.timesteps @@ -392,7 +376,7 @@ def __sample_inpainting( extra_step_kwargs["generator"] = generator # denoising loop - self.model.unet_to(self.train_device) + self.model.materialize_only("unet") for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): latent_model_input = noise_scheduler.scale_model_input(latent_image, timestep) latent_model_input = torch.concat( @@ -428,11 +412,8 @@ def __sample_inpainting( on_update_progress(i + 1, len(timesteps)) - self.model.unet_to(self.temp_device) - torch_gc() - # decode - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") latent_image = latent_image.to(dtype=self.model.vae_train_dtype.torch_dtype()) with self.model.vae_autocast_context: @@ -441,9 +422,6 @@ def __sample_inpainting( do_denormalize = [True] * image.shape[0] image = image_processor.postprocess(image, output_type='pil', do_denormalize=do_denormalize) - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], diff --git a/modules/modelSampler/WuerstchenSampler.py b/modules/modelSampler/WuerstchenSampler.py index d679757c2..a1e9d8ba0 100644 --- a/modules/modelSampler/WuerstchenSampler.py +++ b/modules/modelSampler/WuerstchenSampler.py @@ -11,7 +11,6 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat -from modules.util.torch_util import torch_gc import torch @@ -50,7 +49,7 @@ def __sample_prior( on_update_progress, ): # prepare prompt - self.model.prior_text_encoder_to(self.train_device) + self.model.materialize_only("text_encoder") prompt_embedding, pooled_prompt_embedding = self.model.encode_text( text=prompt, @@ -70,9 +69,6 @@ def __sample_prior( combined_pooled_prompt_embedding = torch.cat([pooled_negative_prompt_embedding, pooled_prompt_embedding]) \ .to(dtype=self.model.prior_train_dtype.torch_dtype()) - self.model.prior_text_encoder_to(self.temp_device) - torch_gc() - # prepare timesteps prior_noise_scheduler.set_timesteps(diffusion_steps, device=self.train_device) timesteps = prior_noise_scheduler.timesteps @@ -95,7 +91,7 @@ def __sample_prior( clip_img = torch.zeros(size=(2, 1, 768), dtype=self.model.prior_train_dtype.torch_dtype(), device=combined_prompt_embedding.device) - self.model.prior_prior_to(self.train_device) + self.model.materialize_only("prior") for i, timestep in enumerate(tqdm(timesteps[:-1], desc="sampling")): timestep = torch.stack([timestep]).to(dtype=self.model.prior_train_dtype.torch_dtype()) @@ -134,9 +130,6 @@ def __sample_prior( on_update_progress(i + 1, len(timesteps)) - self.model.prior_prior_to(self.temp_device) - torch_gc() - if self.model_type.is_wuerstchen_v2(): latent_image = latent_image * 42.0 - 1.0 @@ -161,9 +154,9 @@ def __sample_decoder( ): # prepare prompt if self.model_type.is_wuerstchen_v2(): - self.model.decoder_text_encoder_to(self.train_device) + self.model.materialize_only("decoder_text_encoder") elif self.model_type.is_stable_cascade(): - self.model.prior_text_encoder_to(self.train_device) + self.model.materialize_only("text_encoder") tokenizer_output = decoder_tokenizer( prompt, padding='max_length', @@ -188,12 +181,6 @@ def __sample_decoder( if self.model_type.is_stable_cascade(): prompt_embedding = text_encoder_output.text_embeds.unsqueeze(1) - if self.model_type.is_wuerstchen_v2(): - self.model.decoder_text_encoder_to(self.temp_device) - elif self.model_type.is_stable_cascade(): - self.model.prior_text_encoder_to(self.temp_device) - torch_gc() - # prepare timesteps decoder_noise_scheduler.set_timesteps(10, device=self.train_device) timesteps = decoder_noise_scheduler.timesteps @@ -214,7 +201,7 @@ def __sample_decoder( if "generator" in set(inspect.signature(decoder_noise_scheduler.step).parameters.keys()): extra_step_kwargs["generator"] = generator - self.model.decoder_decoder_to(self.train_device) + self.model.materialize_only("decoder") for i, timestep in enumerate(tqdm(timesteps[:-1], desc="sampling")): timestep = torch.stack([timestep]).to(dtype=self.model.prior_train_dtype.torch_dtype()) @@ -248,9 +235,6 @@ def __sample_decoder( on_update_progress(i + 1, len(timesteps)) - self.model.decoder_decoder_to(self.temp_device) - torch_gc() - return latent_image @torch.no_grad() @@ -322,16 +306,13 @@ def __sample_base( ) # decode vqgan - self.model.decoder_vqgan_to(self.train_device) + self.model.materialize_only("decoder_vqgan") latents = decoder_vqgan.config.scale_factor * latent_image image_tensor = decoder_vqgan.decode(latents).sample.clamp(0, 1) image_array = image_tensor.permute(0, 2, 3, 1).cpu().squeeze().float().numpy() image_array = (image_array * 255).round().astype("uint8") - self.model.decoder_vqgan_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=Image.fromarray(image_array), diff --git a/modules/modelSampler/ZImageSampler.py b/modules/modelSampler/ZImageSampler.py index 0cfcb46d6..0e001df2a 100644 --- a/modules/modelSampler/ZImageSampler.py +++ b/modules/modelSampler/ZImageSampler.py @@ -12,7 +12,6 @@ from modules.util.enum.ModelType import ModelType from modules.util.enum.NoiseScheduler import NoiseScheduler from modules.util.enum.VideoFormat import VideoFormat -from modules.util.torch_util import torch_gc import torch @@ -65,7 +64,7 @@ def __sample_base( #patch_size = 2 # prepare prompt - self.model.text_encoder_to(self.train_device) + self.model.materialize_only("text_encoder") batch_size = 2 if cfg_scale > 1.0 else 1 prompt_embedding = self.model.encode_text( @@ -74,9 +73,6 @@ def __sample_base( train_device=self.train_device, ) - self.model.text_encoder_to(self.temp_device) - torch_gc() - # prepare latent image latent_image = torch.randn( size=(1, num_latent_channels, height // vae_scale_factor, width // vae_scale_factor), @@ -94,7 +90,7 @@ def __sample_base( if "generator" in set(inspect.signature(noise_scheduler.step).parameters.keys()): extra_step_kwargs["generator"] = generator - self.model.transformer_to(self.train_device) + self.model.materialize_only("transformer") for i, timestep in enumerate(tqdm(timesteps, desc="sampling")): latent_model_input = latent_image.unsqueeze(2).to(dtype=self.model.train_dtype.torch_dtype()) latent_model_input = torch.cat([latent_model_input] * batch_size) @@ -118,18 +114,13 @@ def __sample_base( on_update_progress(i + 1, len(timesteps)) - self.model.transformer_to(self.temp_device) - torch_gc() - self.model.vae_to(self.train_device) + self.model.materialize_only("vae") latents = self.model.unscale_latents(latent_image) image = vae.decode(latents, return_dict=False)[0] image = image_processor.postprocess(image, output_type='pil') - self.model.vae_to(self.temp_device) - torch_gc() - return ModelSamplerOutput( file_type=FileType.IMAGE, data=image[0], diff --git a/modules/modelSetup/AnimaFineTuneSetup.py b/modules/modelSetup/AnimaFineTuneSetup.py index 2825d7cf7..3c62b006f 100644 --- a/modules/modelSetup/AnimaFineTuneSetup.py +++ b/modules/modelSetup/AnimaFineTuneSetup.py @@ -10,25 +10,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.ANIMA, TrainingMethod.FINE_TUNE) class AnimaFineTuneSetup( BaseAnimaSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: AnimaModel, @@ -70,10 +56,14 @@ def setup_train_device( config: TrainConfig, ): vae_on_train_device = not config.latent_caching - - model.text_encoder_to(self.temp_device if config.latent_caching else self.train_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + text_encoder_on_train_device = not config.latent_caching + + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) model.text_encoder.eval() model.text_conditioner.eval() diff --git a/modules/modelSetup/AnimaLoRASetup.py b/modules/modelSetup/AnimaLoRASetup.py index ef3086a10..b5edd80b7 100644 --- a/modules/modelSetup/AnimaLoRASetup.py +++ b/modules/modelSetup/AnimaLoRASetup.py @@ -10,25 +10,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.ANIMA, TrainingMethod.LORA) class AnimaLoRASetup( BaseAnimaSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: AnimaModel, @@ -82,10 +68,14 @@ def setup_train_device( config: TrainConfig, ): vae_on_train_device = not config.latent_caching - - model.text_encoder_to(self.temp_device if config.latent_caching else self.train_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + text_encoder_on_train_device = not config.latent_caching + + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) model.text_encoder.eval() model.text_conditioner.eval() diff --git a/modules/modelSetup/BaseAnimaSetup.py b/modules/modelSetup/BaseAnimaSetup.py index 0ed9bbb19..763462696 100644 --- a/modules/modelSetup/BaseAnimaSetup.py +++ b/modules/modelSetup/BaseAnimaSetup.py @@ -14,8 +14,6 @@ enable_checkpointing_for_qwen_transformer, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -23,7 +21,6 @@ from torch import Tensor -#TODO share more code with other models class BaseAnimaSetup( BaseModelSetup, ModelSetupDiffusionLossMixin, @@ -46,25 +43,10 @@ def setup_optimizations( model: AnimaModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_qwen_transformer(model.transformer, config, config.transformer) - model.text_encoder_offload_conductor = enable_checkpointing_for_qwen3_encoder_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) - - self._set_attention_backend(model.transformer, config.attention_mechanism, mask=False) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_qwen_transformer, attention_mask=False) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_qwen3_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) def predict( self, @@ -166,9 +148,7 @@ def calculate_loss( ).mean() def prepare_text_caching(self, model: AnimaModel, config: TrainConfig): - model.to(self.temp_device) - - model.text_encoder_to(self.train_device) + model.materialize_only("text_encoder") model.eval() torch_gc() diff --git a/modules/modelSetup/BaseChromaSetup.py b/modules/modelSetup/BaseChromaSetup.py index 01dcfab0c..adf449e49 100644 --- a/modules/modelSetup/BaseChromaSetup.py +++ b/modules/modelSetup/BaseChromaSetup.py @@ -16,16 +16,12 @@ enable_checkpointing_for_t5_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch from torch import Tensor -#TODO share more code with Flux and other models class BaseChromaSetup( BaseModelSetup, ModelSetupDiffusionLossMixin, @@ -48,25 +44,10 @@ def setup_optimizations( model: ChromaModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_chroma_transformer(model.transformer, config, config.transformer) - model.text_encoder_offload_conductor = enable_checkpointing_for_t5_encoder_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) - - self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_chroma_transformer, attention_mask=True) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_t5_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) def _setup_embeddings( self, @@ -265,10 +246,7 @@ def calculate_loss( def prepare_text_caching(self, model: ChromaModel, config: TrainConfig): - model.to(self.temp_device) - if not config.train_text_encoder_or_embedding(): - model.text_encoder_to(self.train_device) + model.materialize_only("text_encoder") model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseErnieSetup.py b/modules/modelSetup/BaseErnieSetup.py index c412073e8..379747ba1 100644 --- a/modules/modelSetup/BaseErnieSetup.py +++ b/modules/modelSetup/BaseErnieSetup.py @@ -14,9 +14,6 @@ enable_checkpointing_for_mistral_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -44,25 +41,10 @@ def setup_optimizations( model: ErnieModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_ernie_transformer(model.transformer, config, config.transformer) - model.text_encoder_offload_conductor = enable_checkpointing_for_mistral_encoder_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) - - self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_ernie_transformer, attention_mask=True) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_mistral_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) def predict( self, @@ -160,7 +142,5 @@ def calculate_loss( ).mean() def prepare_text_caching(self, model: ErnieModel, config: TrainConfig): - model.to(self.temp_device) - model.text_encoder_to(self.train_device) + model.materialize_only("text_encoder") model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseFlux2Setup.py b/modules/modelSetup/BaseFlux2Setup.py index 6b2c1bb3e..a7509e77c 100644 --- a/modules/modelSetup/BaseFlux2Setup.py +++ b/modules/modelSetup/BaseFlux2Setup.py @@ -16,9 +16,6 @@ enable_checkpointing_for_qwen3_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -44,28 +41,12 @@ def setup_optimizations( model: Flux2Model, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_flux2_transformer(model.transformer, config, config.transformer) - if model.is_dev(): - model.text_encoder_offload_conductor = enable_checkpointing_for_mistral_encoder_layers(model.text_encoder, config, config.text_encoder) - else: - model.text_encoder_offload_conductor = enable_checkpointing_for_qwen3_encoder_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) - - self._set_attention_backend(model.transformer, config.attention_mechanism, mask=False) + text_encoder_checkpointing_fn = enable_checkpointing_for_mistral_encoder_layers if model.is_dev() \ + else enable_checkpointing_for_qwen3_encoder_layers + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_flux2_transformer, attention_mask=False) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, text_encoder_checkpointing_fn, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) def predict( self, @@ -182,7 +163,5 @@ def calculate_loss( ).mean() def prepare_text_caching(self, model: FluxModel, config: TrainConfig): - model.to(self.temp_device) - model.text_encoder_to(self.train_device) + model.materialize_only("text_encoder") model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseFluxSetup.py b/modules/modelSetup/BaseFluxSetup.py index 55999fb20..2bb5619dd 100644 --- a/modules/modelSetup/BaseFluxSetup.py +++ b/modules/modelSetup/BaseFluxSetup.py @@ -17,8 +17,6 @@ enable_checkpointing_for_t5_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress @@ -48,29 +46,11 @@ def setup_optimizations( model: FluxModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_flux_transformer(model.transformer, config, config.transformer) - if model.text_encoder_1 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_1, config, config.text_encoder) - if model.text_encoder_2 is not None: - model.text_encoder_2_offload_conductor = enable_checkpointing_for_t5_encoder_layers(model.text_encoder_2, config, config.text_encoder_2) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_2_autocast_context, model.text_encoder_2_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder_1, self.train_device, model.train_dtype, config) - quantize_layers(model.text_encoder_2, self.train_device, model.text_encoder_2_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) - - self._set_attention_backend(model.transformer, config.attention_mechanism, mask=False) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_flux_transformer, attention_mask=False) + self._setup_model_part(model, config, "text_encoder_1", config.text_encoder, enable_checkpointing_for_clip_encoder_layers) + self._setup_model_part(model, config, "text_encoder_2", config.text_encoder_2, enable_checkpointing_for_t5_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) def _setup_embeddings( self, @@ -326,13 +306,12 @@ def calculate_loss( ).mean() def prepare_text_caching(self, model: FluxModel, config: TrainConfig): - model.to(self.temp_device) - + parts = [] if not config.train_text_encoder_or_embedding(): - model.text_encoder_to(self.train_device) - + parts.append("text_encoder") if not config.train_text_encoder_2_or_embedding(): - model.text_encoder_2_to(self.train_device) + parts.append("text_encoder_2") + model.materialize_only(*parts) model.eval() torch_gc() diff --git a/modules/modelSetup/BaseHiDreamSetup.py b/modules/modelSetup/BaseHiDreamSetup.py index 1a06f961b..71cbe8b24 100644 --- a/modules/modelSetup/BaseHiDreamSetup.py +++ b/modules/modelSetup/BaseHiDreamSetup.py @@ -17,9 +17,6 @@ enable_checkpointing_for_t5_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -48,43 +45,13 @@ def setup_optimizations( model: HiDreamModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_hi_dream_transformer(model.transformer, config, config.transformer) - if model.text_encoder_1 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_1, config, config.text_encoder) - if model.text_encoder_2 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_2, config, config.text_encoder_2) - if model.text_encoder_3 is not None: - model.text_encoder_3_offload_conductor = enable_checkpointing_for_t5_encoder_layers(model.text_encoder_3, config, config.text_encoder_3) - if model.text_encoder_4 is not None: - model.text_encoder_4_offload_conductor = enable_checkpointing_for_llama_encoder_layers(model.text_encoder_4, config, config.text_encoder_4) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_3_autocast_context, model.text_encoder_3_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - model.transformer_autocast_context, model.transformer_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder_1, self.train_device, model.train_dtype, config) - quantize_layers(model.text_encoder_2, self.train_device, model.train_dtype, config) - quantize_layers(model.text_encoder_3, self.train_device, model.text_encoder_3_train_dtype, config) - quantize_layers(model.text_encoder_4, self.train_device, model.train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.transformer_train_dtype, config) - - self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_hi_dream_transformer, disable_fp16_autocast=True, attention_mask=True) + self._setup_model_part(model, config, "text_encoder_1", config.text_encoder, enable_checkpointing_for_clip_encoder_layers) + self._setup_model_part(model, config, "text_encoder_2", config.text_encoder_2, enable_checkpointing_for_clip_encoder_layers) + self._setup_model_part(model, config, "text_encoder_3", config.text_encoder_3, enable_checkpointing_for_t5_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "text_encoder_4", config.text_encoder_4, enable_checkpointing_for_llama_encoder_layers) + self._setup_model_part(model, config, "vae", config.vae) def _setup_embeddings( self, @@ -411,19 +378,15 @@ def calculate_loss( ).mean() def prepare_text_caching(self, model: HiDreamModel, config: TrainConfig): - model.to(self.temp_device) - + parts = [] if not config.train_text_encoder_or_embedding(): - model.text_encoder_to(self.train_device) - + parts.append("text_encoder") if not config.train_text_encoder_2_or_embedding(): - model.text_encoder_2_to(self.train_device) - + parts.append("text_encoder_2") if not config.train_text_encoder_3_or_embedding(): - model.text_encoder_3_to(self.train_device) - + parts.append("text_encoder_3") if not config.train_text_encoder_4_or_embedding(): - model.text_encoder_4_to(self.train_device) + parts.append("text_encoder_4") + model.materialize_only(*parts) model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseHunyuanVideoSetup.py b/modules/modelSetup/BaseHunyuanVideoSetup.py index 8e99dce84..a2749fef7 100644 --- a/modules/modelSetup/BaseHunyuanVideoSetup.py +++ b/modules/modelSetup/BaseHunyuanVideoSetup.py @@ -17,9 +17,6 @@ enable_checkpointing_for_llama_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -48,30 +45,13 @@ def setup_optimizations( model: HunyuanVideoModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_hunyuan_video_transformer(model.transformer, config, config.transformer) - if model.text_encoder_1 is not None: - model.text_encoder_1_offload_conductor = enable_checkpointing_for_llama_encoder_layers(model.text_encoder_1, config, config.text_encoder) - if model.text_encoder_2 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_2, config, config.text_encoder_2) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.transformer_autocast_context, model.transformer_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder_1, self.train_device, model.train_dtype, config) - quantize_layers(model.text_encoder_2, self.train_device, model.train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.transformer_train_dtype, config) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_hunyuan_video_transformer, disable_fp16_autocast=True, attention_mask=True) + self._setup_model_part(model, config, "text_encoder_1", config.text_encoder, enable_checkpointing_for_llama_encoder_layers) + self._setup_model_part(model, config, "text_encoder_2", config.text_encoder_2, enable_checkpointing_for_clip_encoder_layers) + self._setup_model_part(model, config, "vae", config.vae) model.vae.enable_tiling() - self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) def _setup_embeddings( self, @@ -294,13 +274,11 @@ def calculate_loss( ).mean() def prepare_text_caching(self, model: HunyuanVideoModel, config: TrainConfig): - model.to(self.temp_device) - + parts = [] if not config.train_text_encoder_or_embedding(): - model.text_encoder_to(self.train_device) - + parts.append("text_encoder") if not config.train_text_encoder_2_or_embedding(): - model.text_encoder_2_to(self.train_device) + parts.append("text_encoder_2") + model.materialize_only(*parts) model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseIdeogramSetup.py b/modules/modelSetup/BaseIdeogramSetup.py index f067e8c2e..966a4eb2f 100644 --- a/modules/modelSetup/BaseIdeogramSetup.py +++ b/modules/modelSetup/BaseIdeogramSetup.py @@ -14,9 +14,6 @@ enable_checkpointing_for_qwen3vl_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -44,37 +41,13 @@ def setup_optimizations( model: IdeogramModel, config: TrainConfig, ): - # Only the conditional transformer is trained, so gradient checkpointing applies there. - model.transformer_offload_conductor = \ - enable_checkpointing_for_ideogram_transformer(model.transformer, config, config.transformer) - - # The unconditional transformer is frozen, but it still benefits from layer offloading - # since both transformers need to fit in VRAM during sampling. It is optional, so may be unloaded. - if model.unconditional_transformer is not None: - model.unconditional_transformer_offload_conductor = \ - enable_checkpointing_for_ideogram_transformer(model.unconditional_transformer, config, config.unconditional_transformer) - - model.text_encoder_offload_conductor = enable_checkpointing_for_qwen3vl_encoder_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) - quantize_layers(model.unconditional_transformer, self.train_device, model.train_dtype, config) - - self._set_attention_backend(model.transformer, config.attention_mechanism, mask=False) - if model.unconditional_transformer is not None: - self._set_attention_backend(model.unconditional_transformer, config.attention_mechanism, mask=False) + super().setup_optimizations(model, config) + # The unconditional transformer is frozen but still layer-offloaded so both transformers fit in VRAM + # during sampling; it is optional, so _setup_model_part skips it when unloaded. + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_ideogram_transformer, attention_mask=False) + self._setup_model_part(model, config, "unconditional_transformer", config.unconditional_transformer, enable_checkpointing_for_ideogram_transformer, attention_mask=False) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_qwen3vl_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) def predict( self, @@ -204,7 +177,5 @@ def calculate_loss( ).mean() def prepare_text_caching(self, model: IdeogramModel, config: TrainConfig): - model.to(self.temp_device) - model.text_encoder_to(self.train_device) + model.materialize_only("text_encoder") model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseKrea2Setup.py b/modules/modelSetup/BaseKrea2Setup.py index 52e9b5959..bacf50741 100644 --- a/modules/modelSetup/BaseKrea2Setup.py +++ b/modules/modelSetup/BaseKrea2Setup.py @@ -14,9 +14,6 @@ enable_checkpointing_for_qwen3vl_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -46,25 +43,10 @@ def setup_optimizations( model: Krea2Model, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_krea2_transformer(model.transformer, config, config.transformer) - model.text_encoder_offload_conductor = enable_checkpointing_for_qwen3vl_encoder_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) - - self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_krea2_transformer, attention_mask=True) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_qwen3vl_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) def predict( self, @@ -178,10 +160,7 @@ def calculate_loss( ).mean() def prepare_text_caching(self, model: Krea2Model, config: TrainConfig): - model.to(self.temp_device) - if not config.train_text_encoder_or_embedding(): - model.text_encoder_to(self.train_device) + model.materialize_only("text_encoder") model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseModelSetup.py b/modules/modelSetup/BaseModelSetup.py index 8a6818908..fea9c94b9 100644 --- a/modules/modelSetup/BaseModelSetup.py +++ b/modules/modelSetup/BaseModelSetup.py @@ -3,10 +3,12 @@ from modules.model.BaseModel import BaseModel from modules.util.config.TrainConfig import TrainConfig, TrainEmbeddingConfig, TrainModelPartConfig +from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context from modules.util.enum.AttentionMechanism import AttentionMechanism from modules.util.enum.TrainingMethod import TrainingMethod from modules.util.ModuleFilter import ModuleFilter from modules.util.NamedParameterGroup import NamedParameterGroup, NamedParameterGroupCollection +from modules.util.quantization_util import quantize_layers from modules.util.TimedActionMixin import TimedActionMixin from modules.util.TrainProgress import TrainProgress @@ -47,7 +49,10 @@ def setup_optimizations( model: BaseModel, config: TrainConfig, ): - pass + # Model-wide dtype/autocast, shared by every leaf. Leaves call super() first so model.train_dtype is + # set before their first _setup_model_part, which reads it for the non-fp16 quantize path. + model.train_dtype = config.train_dtype + model.autocast_context = create_autocast_context(self.train_device, config.train_dtype, config.enable_autocast_cache) @abstractmethod def setup_model( @@ -237,6 +242,43 @@ def _setup_model_part_requires_grad( for param in self.frozen_parameters[unique_name]: param.requires_grad_(False) + def _setup_model_part( + self, + model, + config: TrainConfig, + attr: str, + config_part: TrainModelPartConfig, + checkpointing_fn=None, + *, + disable_fp16_autocast: bool = False, + attention_mask: bool | None = None, + ): + # Per-part optimization wiring, called once per model part from each leaf. The optional + # disable_fp16_autocast context and its dtype are stored per-part and can differ per part + # (e.g. HiDream disables fp16 for both text_encoder_3 and the transformer). checkpointing_fn returns + # None for non-offloadable parts (SD/SDXL UNet), so no conductor is stored for those. + module = getattr(model, attr) + if module is None: + return + + if checkpointing_fn is not None: + conductor = checkpointing_fn(module, config, config_part) + if conductor is not None: + setattr(model, f"{attr}_offload_conductor", conductor) + + if disable_fp16_autocast: + autocast_context, train_dtype = disable_fp16_autocast_context( + self.train_device, config.train_dtype, config.fallback_train_dtype, config.enable_autocast_cache) + setattr(model, f"{attr}_autocast_context", autocast_context) + setattr(model, f"{attr}_train_dtype", train_dtype) + else: + train_dtype = model.train_dtype + + quantize_layers(module, self.train_device, train_dtype, config) + + if attention_mask is not None: + self._set_attention_backend(module, config.attention_mechanism, mask=attention_mask) + @staticmethod def _set_attention_backend(component, attn: AttentionMechanism, mask: bool): match attn: diff --git a/modules/modelSetup/BasePixArtAlphaSetup.py b/modules/modelSetup/BasePixArtAlphaSetup.py index 57bf3a40b..c169cafe3 100644 --- a/modules/modelSetup/BasePixArtAlphaSetup.py +++ b/modules/modelSetup/BasePixArtAlphaSetup.py @@ -16,9 +16,6 @@ enable_checkpointing_for_t5_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -50,23 +47,10 @@ def setup_optimizations( model: PixArtAlphaModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_basic_transformer_blocks(model.transformer, config, config.transformer, offload_enabled=True) - model.text_encoder_offload_conductor = enable_checkpointing_for_t5_encoder_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) - self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_basic_transformer_blocks, attention_mask=True) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_t5_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) def _setup_embeddings( self, @@ -326,10 +310,7 @@ def calculate_loss( ).mean() def prepare_text_caching(self, model: PixArtAlphaModel, config: TrainConfig): - model.to(self.temp_device) - if not config.train_text_encoder_or_embedding(): - model.text_encoder_to(self.train_device) + model.materialize_only("text_encoder") model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseQwenSetup.py b/modules/modelSetup/BaseQwenSetup.py index a618dc28f..01852ae42 100644 --- a/modules/modelSetup/BaseQwenSetup.py +++ b/modules/modelSetup/BaseQwenSetup.py @@ -14,16 +14,12 @@ enable_checkpointing_for_qwen_transformer, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch from torch import Tensor -#TODO share more code with other models class BaseQwenSetup( BaseModelSetup, ModelSetupDiffusionLossMixin, @@ -45,24 +41,10 @@ def setup_optimizations( model: QwenModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_qwen_transformer(model.transformer, config, config.transformer) - model.text_encoder_offload_conductor = enable_checkpointing_for_qwen25vl_encoder_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) - self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_qwen_transformer, attention_mask=True) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_qwen25vl_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) def predict( self, @@ -177,10 +159,7 @@ def calculate_loss( ).mean() def prepare_text_caching(self, model: QwenModel, config: TrainConfig): - model.to(self.temp_device) - if not config.train_text_encoder_or_embedding(): - model.text_encoder_to(self.train_device) + model.materialize_only("text_encoder") model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseSanaSetup.py b/modules/modelSetup/BaseSanaSetup.py index a96afd770..b3d2137b0 100644 --- a/modules/modelSetup/BaseSanaSetup.py +++ b/modules/modelSetup/BaseSanaSetup.py @@ -16,9 +16,7 @@ enable_checkpointing_for_sana_transformer, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc +from modules.util.dtype_util import disable_fp16_autocast_context from modules.util.TrainProgress import TrainProgress import torch @@ -51,19 +49,12 @@ def setup_optimizations( config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_sana_transformer(model.transformer, config, config.transformer) - model.text_encoder_offload_conductor = enable_checkpointing_for_gemma_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_autocast_context, model.text_encoder_train_dtype = disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) + super().setup_optimizations(model, config) + # Sana's vae runs under its own fp16-disabled autocast in predict(), set here rather than via + # _setup_model_part's autocast handling. Note a preexisting inconsistency (predates this refactor, + # kept as-is since Sana is largely outdated): unlike SDXL, model.vae_train_dtype is computed but never + # read anywhere, and the vae below is quantized with model.train_dtype, not this fp16-disabled dtype. model.vae_autocast_context, model.vae_train_dtype = disable_fp16_autocast_context( self.train_device, config.train_dtype, @@ -71,10 +62,9 @@ def setup_optimizations( config.enable_autocast_cache, ) - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) - self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_sana_transformer, attention_mask=True) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_gemma_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) def _setup_embeddings( self, @@ -246,10 +236,7 @@ def calculate_loss( ).mean() def prepare_text_caching(self, model: SanaModel, config: TrainConfig): - model.to(self.temp_device) - if not config.train_text_encoder_or_embedding(): - model.text_encoder_to(self.train_device) + model.materialize_only("text_encoder") model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseStableDiffusion3Setup.py b/modules/modelSetup/BaseStableDiffusion3Setup.py index 678d727ff..9fd69ccf6 100644 --- a/modules/modelSetup/BaseStableDiffusion3Setup.py +++ b/modules/modelSetup/BaseStableDiffusion3Setup.py @@ -17,9 +17,6 @@ enable_checkpointing_for_t5_encoder_layers, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -47,31 +44,12 @@ def setup_optimizations( model: StableDiffusion3Model, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_stable_diffusion_3_transformer(model.transformer, config, config.transformer) - if model.text_encoder_1 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_1, config, config.text_encoder) - if model.text_encoder_2 is not None: - enable_checkpointing_for_clip_encoder_layers(model.text_encoder_2, config, config.text_encoder_2) - if model.text_encoder_3 is not None: - model.text_encoder_3_offload_conductor = enable_checkpointing_for_t5_encoder_layers(model.text_encoder_3, config, config.text_encoder_3) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - model.text_encoder_3_autocast_context, model.text_encoder_3_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder_1, self.train_device, model.train_dtype, config) - quantize_layers(model.text_encoder_2, self.train_device, model.train_dtype, config) - quantize_layers(model.text_encoder_3, self.train_device, model.text_encoder_3_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) - self._set_attention_backend(model.transformer, config.attention_mechanism, mask=False) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_stable_diffusion_3_transformer, attention_mask=False) + self._setup_model_part(model, config, "text_encoder_1", config.text_encoder, enable_checkpointing_for_clip_encoder_layers) + self._setup_model_part(model, config, "text_encoder_2", config.text_encoder_2, enable_checkpointing_for_clip_encoder_layers) + self._setup_model_part(model, config, "text_encoder_3", config.text_encoder_3, enable_checkpointing_for_t5_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) def _setup_embeddings( self, @@ -346,16 +324,13 @@ def calculate_loss( ).mean() def prepare_text_caching(self, model: StableDiffusion3Model, config: TrainConfig): - model.to(self.temp_device) - + parts = [] if not config.train_text_encoder_or_embedding(): - model.text_encoder_to(self.train_device) - + parts.append("text_encoder") if not config.train_text_encoder_2_or_embedding(): - model.text_encoder_2_to(self.train_device) - + parts.append("text_encoder_2") if not config.train_text_encoder_3_or_embedding(): - model.text_encoder_3_to(self.train_device) + parts.append("text_encoder_3") + model.materialize_only(*parts) model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseStableDiffusionSetup.py b/modules/modelSetup/BaseStableDiffusionSetup.py index 5be9fc97e..0e10939ef 100644 --- a/modules/modelSetup/BaseStableDiffusionSetup.py +++ b/modules/modelSetup/BaseStableDiffusionSetup.py @@ -17,9 +17,7 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.conv_util import apply_circular_padding_to_conv2d -from modules.util.dtype_util import create_autocast_context from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -50,9 +48,11 @@ def setup_optimizations( model: StableDiffusionModel, config: TrainConfig, ): + # Not routed through _setup_model_part: the UNet's checkpointing needs supports_offloading=False, which + # _setup_model_part's checkpointing_fn slot doesn't pass, so the parts are wired by hand here. if config.unet.checkpointing_enabled(): model.unet.enable_gradient_checkpointing() - enable_checkpointing_for_basic_transformer_blocks(model.unet, config, config.unet, offload_enabled=False) + enable_checkpointing_for_basic_transformer_blocks(model.unet, config, config.unet, supports_offloading=False) enable_checkpointing_for_clip_encoder_layers(model.text_encoder, config, config.text_encoder) if config.force_circular_padding: @@ -61,8 +61,7 @@ def setup_optimizations( if model.unet_lora is not None: apply_circular_padding_to_conv2d(model.unet_lora) - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) + super().setup_optimizations(model, config) quantize_layers(model.text_encoder, self.train_device, model.train_dtype, config) quantize_layers(model.vae, self.train_device, model.train_dtype, config) @@ -334,10 +333,9 @@ def calculate_loss( ).mean() def prepare_text_caching(self, model: StableDiffusionModel, config: TrainConfig): - model.to(self.temp_device) - if not config.train_text_encoder_or_embedding(): - model.text_encoder_to(self.train_device) + model.materialize_only("text_encoder") + if model.depth_estimator is not None: + model.depth_estimator.to(self.temp_device) model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseStableDiffusionXLSetup.py b/modules/modelSetup/BaseStableDiffusionXLSetup.py index e5c7bf0c3..dd7a4fe6f 100644 --- a/modules/modelSetup/BaseStableDiffusionXLSetup.py +++ b/modules/modelSetup/BaseStableDiffusionXLSetup.py @@ -17,9 +17,8 @@ ) from modules.util.config.TrainConfig import TrainConfig from modules.util.conv_util import apply_circular_padding_to_conv2d -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context +from modules.util.dtype_util import disable_fp16_autocast_context from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -47,9 +46,11 @@ def setup_optimizations( model: StableDiffusionXLModel, config: TrainConfig, ): + # Not routed through _setup_model_part: the UNet's checkpointing needs supports_offloading=False, which + # _setup_model_part's checkpointing_fn slot doesn't pass, so the parts are wired by hand here. if config.unet.checkpointing_enabled(): model.unet.enable_gradient_checkpointing() - enable_checkpointing_for_basic_transformer_blocks(model.unet, config, config.unet, offload_enabled=False) + enable_checkpointing_for_basic_transformer_blocks(model.unet, config, config.unet, supports_offloading=False) enable_checkpointing_for_clip_encoder_layers(model.text_encoder_1, config, config.text_encoder) enable_checkpointing_for_clip_encoder_layers(model.text_encoder_2, config, config.text_encoder_2) @@ -59,8 +60,7 @@ def setup_optimizations( if model.unet_lora is not None: apply_circular_padding_to_conv2d(model.unet_lora) - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) + super().setup_optimizations(model, config) model.vae_autocast_context, model.vae_train_dtype = disable_fp16_autocast_context( self.train_device, @@ -379,13 +379,11 @@ def calculate_loss( ).mean() def prepare_text_caching(self, model: StableDiffusionXLModel, config: TrainConfig): - model.to(self.temp_device) - + parts = [] if not config.train_text_encoder_or_embedding(): - model.text_encoder_to(self.train_device) - + parts.append("text_encoder") if not config.train_text_encoder_2_or_embedding(): - model.text_encoder_2_to(self.train_device) + parts.append("text_encoder_2") + model.materialize_only(*parts) model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseWuerstchenSetup.py b/modules/modelSetup/BaseWuerstchenSetup.py index 078e14fff..c6b807127 100644 --- a/modules/modelSetup/BaseWuerstchenSetup.py +++ b/modules/modelSetup/BaseWuerstchenSetup.py @@ -15,12 +15,10 @@ from modules.util.config.TrainConfig import TrainConfig from modules.util.conv_util import apply_circular_padding_to_conv2d from modules.util.dtype_util import ( - create_autocast_context, disable_bf16_on_fp16_autocast_context, disable_fp16_autocast_context, ) from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -53,6 +51,10 @@ def setup_optimizations( model: WuerstchenModel, config: TrainConfig, ): + # Not routed through _setup_model_part: Wuerstchen's parts (prior_prior, decoder_*, effnet_encoder, + # prior_text_encoder) don't match the transformer/text_encoder/vae shape _setup_model_part assumes and + # take bespoke per-part contexts (stable-cascade prior fp16-disable, effnet bf16-on-fp16), so this + # setup is fully hand-rolled. if config.prior.checkpointing_enabled(): model.prior_prior.enable_gradient_checkpointing() enable_checkpointing_for_clip_encoder_layers(model.prior_text_encoder, config, config.text_encoder) @@ -64,8 +66,7 @@ def setup_optimizations( if model.prior_prior_lora is not None: apply_circular_padding_to_conv2d(model.prior_prior_lora) - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) + super().setup_optimizations(model, config) if model.model_type.is_stable_cascade(): model.prior_autocast_context, model.prior_train_dtype = disable_fp16_autocast_context( @@ -345,10 +346,7 @@ def calculate_loss( ).mean() def prepare_text_caching(self, model: WuerstchenModel, config: TrainConfig): - model.to(self.temp_device) - if not config.train_text_encoder_or_embedding(): - model.prior_text_encoder_to(self.train_device) + model.materialize_only("text_encoder") model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseZImageSetup.py b/modules/modelSetup/BaseZImageSetup.py index a180f57d2..df0f0912c 100644 --- a/modules/modelSetup/BaseZImageSetup.py +++ b/modules/modelSetup/BaseZImageSetup.py @@ -15,9 +15,6 @@ enable_checkpointing_for_z_image_transformer, ) from modules.util.config.TrainConfig import TrainConfig -from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context -from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -46,25 +43,10 @@ def setup_optimizations( model: ZImageModel, config: TrainConfig, ): - model.transformer_offload_conductor = enable_checkpointing_for_z_image_transformer(model.transformer, config, config.transformer) - model.text_encoder_offload_conductor = enable_checkpointing_for_qwen3_encoder_layers(model.text_encoder, config, config.text_encoder) - - model.autocast_context, model.train_dtype = create_autocast_context( - self.train_device, config.train_dtype, config.enable_autocast_cache) - - #TODO necessary if we don't train it? - model.text_encoder_autocast_context, model.text_encoder_train_dtype = \ - disable_fp16_autocast_context( - self.train_device, - config.train_dtype, - config.fallback_train_dtype, - config.enable_autocast_cache, - ) - - quantize_layers(model.text_encoder, self.train_device, model.text_encoder_train_dtype, config) - quantize_layers(model.vae, self.train_device, model.train_dtype, config) - quantize_layers(model.transformer, self.train_device, model.train_dtype, config) - self._set_attention_backend(model.transformer, config.attention_mechanism, mask=True) + super().setup_optimizations(model, config) + self._setup_model_part(model, config, "transformer", config.transformer, enable_checkpointing_for_z_image_transformer, attention_mask=True) + self._setup_model_part(model, config, "text_encoder", config.text_encoder, enable_checkpointing_for_qwen3_encoder_layers, disable_fp16_autocast=True) + self._setup_model_part(model, config, "vae", config.vae) def predict( self, @@ -160,8 +142,6 @@ def calculate_loss( ).mean() def prepare_text_caching(self, model: ZImageModel, config: TrainConfig): - model.to(self.temp_device) - model.text_encoder_to(self.train_device) + model.materialize_only("text_encoder") model.eval() - torch_gc() diff --git a/modules/modelSetup/ChromaEmbeddingSetup.py b/modules/modelSetup/ChromaEmbeddingSetup.py index 88aef8db4..ce49b9db2 100644 --- a/modules/modelSetup/ChromaEmbeddingSetup.py +++ b/modules/modelSetup/ChromaEmbeddingSetup.py @@ -9,25 +9,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.CHROMA_1, TrainingMethod.EMBEDDING) class ChromaEmbeddingSetup( BaseChromaSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: ChromaModel, @@ -74,9 +60,12 @@ def setup_train_device( ): vae_on_train_device = not config.latent_caching - model.text_encoder_to(self.train_device if config.text_encoder.train_embedding else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if config.text_encoder.train_embedding: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) model.text_encoder.eval() model.vae.eval() diff --git a/modules/modelSetup/ChromaFineTuneSetup.py b/modules/modelSetup/ChromaFineTuneSetup.py index 95bd8bac6..64c2aa67f 100644 --- a/modules/modelSetup/ChromaFineTuneSetup.py +++ b/modules/modelSetup/ChromaFineTuneSetup.py @@ -10,25 +10,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.CHROMA_1, TrainingMethod.FINE_TUNE) class ChromaFineTuneSetup( BaseChromaSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: ChromaModel, @@ -87,9 +73,12 @@ def setup_train_device( config.train_text_encoder_or_embedding() \ or not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if config.text_encoder.train: model.text_encoder.train() diff --git a/modules/modelSetup/ChromaLoRASetup.py b/modules/modelSetup/ChromaLoRASetup.py index 43e9c6572..63e3ef964 100644 --- a/modules/modelSetup/ChromaLoRASetup.py +++ b/modules/modelSetup/ChromaLoRASetup.py @@ -11,25 +11,11 @@ from modules.util.torch_util import state_dict_has_prefix from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.CHROMA_1, TrainingMethod.LORA) class ChromaLoRASetup( BaseChromaSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: ChromaModel, @@ -114,9 +100,12 @@ def setup_train_device( config.train_text_encoder_or_embedding() \ or not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if config.text_encoder.train: model.text_encoder.train() diff --git a/modules/modelSetup/ErnieFineTuneSetup.py b/modules/modelSetup/ErnieFineTuneSetup.py index feadd4fca..dc7f2f770 100644 --- a/modules/modelSetup/ErnieFineTuneSetup.py +++ b/modules/modelSetup/ErnieFineTuneSetup.py @@ -10,25 +10,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.ERNIE, TrainingMethod.FINE_TUNE) class ErnieFineTuneSetup( BaseErnieSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: ErnieModel, @@ -67,9 +53,12 @@ def setup_train_device( vae_on_train_device = not config.latent_caching text_encoder_on_train_device = not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) model.text_encoder.eval() model.vae.eval() diff --git a/modules/modelSetup/ErnieLoRASetup.py b/modules/modelSetup/ErnieLoRASetup.py index 1a70ef269..70b449dd0 100644 --- a/modules/modelSetup/ErnieLoRASetup.py +++ b/modules/modelSetup/ErnieLoRASetup.py @@ -10,25 +10,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.ERNIE, TrainingMethod.LORA) class ErnieLoRASetup( BaseErnieSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: ErnieModel, @@ -77,9 +63,12 @@ def setup_train_device( vae_on_train_device = not config.latent_caching text_encoder_on_train_device = not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) model.text_encoder.eval() model.vae.eval() diff --git a/modules/modelSetup/Flux2FineTuneSetup.py b/modules/modelSetup/Flux2FineTuneSetup.py index 7ca128a09..be310d512 100644 --- a/modules/modelSetup/Flux2FineTuneSetup.py +++ b/modules/modelSetup/Flux2FineTuneSetup.py @@ -10,25 +10,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.FLUX_2, TrainingMethod.FINE_TUNE) class Flux2FineTuneSetup( BaseFlux2Setup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: Flux2Model, @@ -67,9 +53,12 @@ def setup_train_device( vae_on_train_device = not config.latent_caching text_encoder_on_train_device = not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) model.text_encoder.eval() model.vae.eval() diff --git a/modules/modelSetup/Flux2LoRASetup.py b/modules/modelSetup/Flux2LoRASetup.py index fe1750528..0358e62c1 100644 --- a/modules/modelSetup/Flux2LoRASetup.py +++ b/modules/modelSetup/Flux2LoRASetup.py @@ -10,25 +10,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.FLUX_2, TrainingMethod.LORA) class Flux2LoRASetup( BaseFlux2Setup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: Flux2Model, @@ -80,9 +66,12 @@ def setup_train_device( vae_on_train_device = not config.latent_caching text_encoder_on_train_device = not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) model.text_encoder.eval() model.vae.eval() diff --git a/modules/modelSetup/FluxEmbeddingSetup.py b/modules/modelSetup/FluxEmbeddingSetup.py index f31c5e922..4751d5e37 100644 --- a/modules/modelSetup/FluxEmbeddingSetup.py +++ b/modules/modelSetup/FluxEmbeddingSetup.py @@ -9,26 +9,12 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.FLUX_DEV_1, TrainingMethod.EMBEDDING) @factory.register(BaseModelSetup, ModelType.FLUX_FILL_DEV_1, TrainingMethod.EMBEDDING) class FluxEmbeddingSetup( BaseFluxSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: FluxModel, @@ -87,10 +73,14 @@ def setup_train_device( ): vae_on_train_device = not config.latent_caching - model.text_encoder_1_to(self.train_device if config.text_encoder.train_embedding else self.temp_device) - model.text_encoder_2_to(self.train_device if config.text_encoder_2.train_embedding else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if config.text_encoder.train_embedding: + parts.append("text_encoder") + if config.text_encoder_2.train_embedding: + parts.append("text_encoder_2") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if model.text_encoder_1 is not None: model.text_encoder_1.eval() diff --git a/modules/modelSetup/FluxFineTuneSetup.py b/modules/modelSetup/FluxFineTuneSetup.py index 5d2aa61cd..4dc713d97 100644 --- a/modules/modelSetup/FluxFineTuneSetup.py +++ b/modules/modelSetup/FluxFineTuneSetup.py @@ -10,26 +10,12 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.FLUX_DEV_1, TrainingMethod.FINE_TUNE) @factory.register(BaseModelSetup, ModelType.FLUX_FILL_DEV_1, TrainingMethod.FINE_TUNE) class FluxFineTuneSetup( BaseFluxSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: FluxModel, @@ -103,10 +89,14 @@ def setup_train_device( config.train_text_encoder_2_or_embedding() \ or not config.latent_caching - model.text_encoder_1_to(self.train_device if text_encoder_1_on_train_device else self.temp_device) - model.text_encoder_2_to(self.train_device if text_encoder_2_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_1_on_train_device: + parts.append("text_encoder") + if text_encoder_2_on_train_device: + parts.append("text_encoder_2") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if model.text_encoder_1: if config.text_encoder.train: diff --git a/modules/modelSetup/FluxLoRASetup.py b/modules/modelSetup/FluxLoRASetup.py index b883dc7b3..557710459 100644 --- a/modules/modelSetup/FluxLoRASetup.py +++ b/modules/modelSetup/FluxLoRASetup.py @@ -11,26 +11,12 @@ from modules.util.torch_util import state_dict_has_prefix from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.FLUX_DEV_1, TrainingMethod.LORA) @factory.register(BaseModelSetup, ModelType.FLUX_FILL_DEV_1, TrainingMethod.LORA) class FluxLoRASetup( BaseFluxSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: FluxModel, @@ -146,10 +132,14 @@ def setup_train_device( config.train_text_encoder_2_or_embedding() \ or not config.latent_caching - model.text_encoder_1_to(self.train_device if text_encoder_1_on_train_device else self.temp_device) - model.text_encoder_2_to(self.train_device if text_encoder_2_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_1_on_train_device: + parts.append("text_encoder") + if text_encoder_2_on_train_device: + parts.append("text_encoder_2") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if model.text_encoder_1: if config.text_encoder.train: diff --git a/modules/modelSetup/HiDreamEmbeddingSetup.py b/modules/modelSetup/HiDreamEmbeddingSetup.py index 659f3834e..b2e80df44 100644 --- a/modules/modelSetup/HiDreamEmbeddingSetup.py +++ b/modules/modelSetup/HiDreamEmbeddingSetup.py @@ -9,25 +9,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.HI_DREAM_FULL, TrainingMethod.EMBEDDING) class HiDreamEmbeddingSetup( BaseHiDreamSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: HiDreamModel, @@ -106,12 +92,18 @@ def setup_train_device( ): vae_on_train_device = not config.latent_caching - model.text_encoder_1_to(self.train_device if config.text_encoder.train_embedding else self.temp_device) - model.text_encoder_2_to(self.train_device if config.text_encoder_2.train_embedding else self.temp_device) - model.text_encoder_3_to(self.train_device if config.text_encoder_3.train_embedding else self.temp_device) - model.text_encoder_4_to(self.train_device if config.text_encoder_4.train_embedding else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if config.text_encoder.train_embedding: + parts.append("text_encoder") + if config.text_encoder_2.train_embedding: + parts.append("text_encoder_2") + if config.text_encoder_3.train_embedding: + parts.append("text_encoder_3") + if config.text_encoder_4.train_embedding: + parts.append("text_encoder_4") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if model.text_encoder_1 is not None: model.text_encoder_1.eval() diff --git a/modules/modelSetup/HiDreamFineTuneSetup.py b/modules/modelSetup/HiDreamFineTuneSetup.py index 2361e69ff..c1bad4337 100644 --- a/modules/modelSetup/HiDreamFineTuneSetup.py +++ b/modules/modelSetup/HiDreamFineTuneSetup.py @@ -12,25 +12,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.HI_DREAM_FULL, TrainingMethod.FINE_TUNE) class HiDreamFineTuneSetup( BaseHiDreamSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: HiDreamModel, @@ -137,12 +123,18 @@ def setup_train_device( config.train_text_encoder_4_or_embedding() \ or not config.latent_caching - model.text_encoder_1_to(self.train_device if text_encoder_1_on_train_device else self.temp_device) - model.text_encoder_2_to(self.train_device if text_encoder_2_on_train_device else self.temp_device) - model.text_encoder_3_to(self.train_device if text_encoder_3_on_train_device else self.temp_device) - model.text_encoder_4_to(self.train_device if text_encoder_4_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_1_on_train_device: + parts.append("text_encoder") + if text_encoder_2_on_train_device: + parts.append("text_encoder_2") + if text_encoder_3_on_train_device: + parts.append("text_encoder_3") + if text_encoder_4_on_train_device: + parts.append("text_encoder_4") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if model.text_encoder_1: if config.text_encoder.train: diff --git a/modules/modelSetup/HiDreamLoRASetup.py b/modules/modelSetup/HiDreamLoRASetup.py index a63a30ab7..6bb90324d 100644 --- a/modules/modelSetup/HiDreamLoRASetup.py +++ b/modules/modelSetup/HiDreamLoRASetup.py @@ -13,25 +13,11 @@ from modules.util.torch_util import state_dict_has_prefix from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.HI_DREAM_FULL, TrainingMethod.LORA) class HiDreamLoRASetup( BaseHiDreamSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: HiDreamModel, @@ -209,12 +195,18 @@ def setup_train_device( config.train_text_encoder_4_or_embedding() \ or not config.latent_caching - model.text_encoder_1_to(self.train_device if text_encoder_1_on_train_device else self.temp_device) - model.text_encoder_2_to(self.train_device if text_encoder_2_on_train_device else self.temp_device) - model.text_encoder_3_to(self.train_device if text_encoder_3_on_train_device else self.temp_device) - model.text_encoder_4_to(self.train_device if text_encoder_4_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_1_on_train_device: + parts.append("text_encoder") + if text_encoder_2_on_train_device: + parts.append("text_encoder_2") + if text_encoder_3_on_train_device: + parts.append("text_encoder_3") + if text_encoder_4_on_train_device: + parts.append("text_encoder_4") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if model.text_encoder_1: if config.text_encoder.train: diff --git a/modules/modelSetup/HunyuanVideoEmbeddingSetup.py b/modules/modelSetup/HunyuanVideoEmbeddingSetup.py index 91022c9f6..4b1168e4b 100644 --- a/modules/modelSetup/HunyuanVideoEmbeddingSetup.py +++ b/modules/modelSetup/HunyuanVideoEmbeddingSetup.py @@ -9,25 +9,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.HUNYUAN_VIDEO, TrainingMethod.EMBEDDING) class HunyuanVideoEmbeddingSetup( BaseHunyuanVideoSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: HunyuanVideoModel, @@ -86,10 +72,14 @@ def setup_train_device( ): vae_on_train_device = not config.latent_caching - model.text_encoder_1_to(self.train_device if config.text_encoder.train_embedding else self.temp_device) - model.text_encoder_2_to(self.train_device if config.text_encoder_2.train_embedding else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if config.text_encoder.train_embedding: + parts.append("text_encoder") + if config.text_encoder_2.train_embedding: + parts.append("text_encoder_2") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if model.text_encoder_1 is not None: model.text_encoder_1.eval() diff --git a/modules/modelSetup/HunyuanVideoFineTuneSetup.py b/modules/modelSetup/HunyuanVideoFineTuneSetup.py index 9bb26e5d0..66786ef19 100644 --- a/modules/modelSetup/HunyuanVideoFineTuneSetup.py +++ b/modules/modelSetup/HunyuanVideoFineTuneSetup.py @@ -17,18 +17,6 @@ class HunyuanVideoFineTuneSetup( BaseHunyuanVideoSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: HunyuanVideoModel, @@ -103,10 +91,14 @@ def setup_train_device( config.train_text_encoder_2_or_embedding() \ or not config.latent_caching - model.text_encoder_1_to(self.train_device if text_encoder_1_on_train_device else self.temp_device) - model.text_encoder_2_to(self.train_device if text_encoder_2_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_1_on_train_device: + parts.append("text_encoder") + if text_encoder_2_on_train_device: + parts.append("text_encoder_2") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if model.text_encoder_1: if config.text_encoder.train: diff --git a/modules/modelSetup/HunyuanVideoLoRASetup.py b/modules/modelSetup/HunyuanVideoLoRASetup.py index 45a5f619e..987df90ee 100644 --- a/modules/modelSetup/HunyuanVideoLoRASetup.py +++ b/modules/modelSetup/HunyuanVideoLoRASetup.py @@ -13,25 +13,11 @@ from modules.util.torch_util import state_dict_has_prefix from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.HUNYUAN_VIDEO, TrainingMethod.LORA) class HunyuanVideoLoRASetup( BaseHunyuanVideoSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: HunyuanVideoModel, @@ -150,10 +136,14 @@ def setup_train_device( config.train_text_encoder_2_or_embedding() \ or not config.latent_caching - model.text_encoder_1_to(self.train_device if text_encoder_1_on_train_device else self.temp_device) - model.text_encoder_2_to(self.train_device if text_encoder_2_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_1_on_train_device: + parts.append("text_encoder") + if text_encoder_2_on_train_device: + parts.append("text_encoder_2") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if model.text_encoder_1: if config.text_encoder.train: diff --git a/modules/modelSetup/IdeogramFineTuneSetup.py b/modules/modelSetup/IdeogramFineTuneSetup.py index a259af03e..79dd5a05c 100644 --- a/modules/modelSetup/IdeogramFineTuneSetup.py +++ b/modules/modelSetup/IdeogramFineTuneSetup.py @@ -10,25 +10,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.IDEOGRAM_4, TrainingMethod.FINE_TUNE) class IdeogramFineTuneSetup( BaseIdeogramSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: IdeogramModel, @@ -71,11 +57,14 @@ def setup_train_device( vae_on_train_device = not config.latent_caching text_encoder_on_train_device = not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) - # the unconditional transformer is only needed for sampling; keep it off the train device during training - model.unconditional_transformer_to(self.temp_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + # the unconditional transformer is only needed for sampling; materialize_only() evicts it as it's + # not in parts, keeping it off the train device during training + model.materialize_only(*parts) model.text_encoder.eval() model.vae.eval() diff --git a/modules/modelSetup/IdeogramLoRASetup.py b/modules/modelSetup/IdeogramLoRASetup.py index e3f83d8bf..fbdcec35e 100644 --- a/modules/modelSetup/IdeogramLoRASetup.py +++ b/modules/modelSetup/IdeogramLoRASetup.py @@ -10,25 +10,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.IDEOGRAM_4, TrainingMethod.LORA) class IdeogramLoRASetup( BaseIdeogramSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: IdeogramModel, @@ -80,11 +66,14 @@ def setup_train_device( vae_on_train_device = not config.latent_caching text_encoder_on_train_device = not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) - # the unconditional transformer is only needed for sampling; keep it off the train device during training - model.unconditional_transformer_to(self.temp_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + # the unconditional transformer is only needed for sampling; materialize_only() evicts it as it's + # not in parts, keeping it off the train device during training + model.materialize_only(*parts) model.text_encoder.eval() model.vae.eval() diff --git a/modules/modelSetup/Krea2FineTuneSetup.py b/modules/modelSetup/Krea2FineTuneSetup.py index c2ff053f1..07cc43944 100644 --- a/modules/modelSetup/Krea2FineTuneSetup.py +++ b/modules/modelSetup/Krea2FineTuneSetup.py @@ -10,25 +10,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.KREA_2, TrainingMethod.FINE_TUNE) class Krea2FineTuneSetup( BaseKrea2Setup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: Krea2Model, @@ -73,9 +59,12 @@ def setup_train_device( vae_on_train_device = not config.latent_caching text_encoder_on_train_device = not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) model.text_encoder.eval() model.vae.eval() diff --git a/modules/modelSetup/Krea2LoRASetup.py b/modules/modelSetup/Krea2LoRASetup.py index 7310317bc..507ec8f59 100644 --- a/modules/modelSetup/Krea2LoRASetup.py +++ b/modules/modelSetup/Krea2LoRASetup.py @@ -10,25 +10,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.KREA_2, TrainingMethod.LORA) class Krea2LoRASetup( BaseKrea2Setup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: Krea2Model, @@ -83,9 +69,12 @@ def setup_train_device( vae_on_train_device = not config.latent_caching text_encoder_on_train_device = not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) model.text_encoder.eval() model.vae.eval() diff --git a/modules/modelSetup/PixArtAlphaEmbeddingSetup.py b/modules/modelSetup/PixArtAlphaEmbeddingSetup.py index 232bf2e01..1c8413f3f 100644 --- a/modules/modelSetup/PixArtAlphaEmbeddingSetup.py +++ b/modules/modelSetup/PixArtAlphaEmbeddingSetup.py @@ -9,26 +9,12 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.PIXART_ALPHA, TrainingMethod.EMBEDDING) @factory.register(BaseModelSetup, ModelType.PIXART_SIGMA, TrainingMethod.EMBEDDING) class PixArtAlphaEmbeddingSetup( BasePixArtAlphaSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: PixArtAlphaModel, @@ -74,9 +60,10 @@ def setup_train_device( ): vae_on_train_device = self.debug_mode - model.text_encoder_to(self.train_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer", "text_encoder"] + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) model.text_encoder.eval() model.vae.eval() diff --git a/modules/modelSetup/PixArtAlphaFineTuneSetup.py b/modules/modelSetup/PixArtAlphaFineTuneSetup.py index e34db3ac5..13bec1203 100644 --- a/modules/modelSetup/PixArtAlphaFineTuneSetup.py +++ b/modules/modelSetup/PixArtAlphaFineTuneSetup.py @@ -10,26 +10,12 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.PIXART_ALPHA, TrainingMethod.FINE_TUNE) @factory.register(BaseModelSetup, ModelType.PIXART_SIGMA, TrainingMethod.FINE_TUNE) class PixArtAlphaFineTuneSetup( BasePixArtAlphaSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: PixArtAlphaModel, @@ -94,9 +80,12 @@ def setup_train_device( or config.train_any_embedding() \ or not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if config.text_encoder.train: model.text_encoder.train() diff --git a/modules/modelSetup/PixArtAlphaLoRASetup.py b/modules/modelSetup/PixArtAlphaLoRASetup.py index 0456b4165..8c7de7766 100644 --- a/modules/modelSetup/PixArtAlphaLoRASetup.py +++ b/modules/modelSetup/PixArtAlphaLoRASetup.py @@ -11,26 +11,12 @@ from modules.util.torch_util import state_dict_has_prefix from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.PIXART_ALPHA, TrainingMethod.LORA) @factory.register(BaseModelSetup, ModelType.PIXART_SIGMA, TrainingMethod.LORA) class PixArtAlphaLoRASetup( BasePixArtAlphaSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: PixArtAlphaModel, @@ -115,9 +101,12 @@ def setup_train_device( or config.train_any_embedding() \ or not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if config.text_encoder.train: model.text_encoder.train() diff --git a/modules/modelSetup/QwenFineTuneSetup.py b/modules/modelSetup/QwenFineTuneSetup.py index 7b9f5dc60..5c4296218 100644 --- a/modules/modelSetup/QwenFineTuneSetup.py +++ b/modules/modelSetup/QwenFineTuneSetup.py @@ -10,25 +10,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.QWEN, TrainingMethod.FINE_TUNE) class QwenFineTuneSetup( BaseQwenSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: QwenModel, @@ -74,9 +60,12 @@ def setup_train_device( config.train_text_encoder_or_embedding() \ or not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if config.text_encoder.train: model.text_encoder.train() diff --git a/modules/modelSetup/QwenLoRASetup.py b/modules/modelSetup/QwenLoRASetup.py index bf114b0ea..15e637b73 100644 --- a/modules/modelSetup/QwenLoRASetup.py +++ b/modules/modelSetup/QwenLoRASetup.py @@ -11,25 +11,11 @@ from modules.util.torch_util import state_dict_has_prefix from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.QWEN, TrainingMethod.LORA) class QwenLoRASetup( BaseQwenSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: QwenModel, @@ -101,9 +87,12 @@ def setup_train_device( config.train_text_encoder_or_embedding() \ or not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if config.text_encoder.train: model.text_encoder.train() diff --git a/modules/modelSetup/SanaEmbeddingSetup.py b/modules/modelSetup/SanaEmbeddingSetup.py index fad567b75..f3812a0f4 100644 --- a/modules/modelSetup/SanaEmbeddingSetup.py +++ b/modules/modelSetup/SanaEmbeddingSetup.py @@ -9,25 +9,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.SANA, TrainingMethod.EMBEDDING) class SanaEmbeddingSetup( BaseSanaSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: SanaModel, @@ -73,9 +59,10 @@ def setup_train_device( ): vae_on_train_device = self.debug_mode - model.text_encoder_to(self.train_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer", "text_encoder"] + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) model.text_encoder.eval() model.vae.eval() diff --git a/modules/modelSetup/SanaFineTuneSetup.py b/modules/modelSetup/SanaFineTuneSetup.py index e782e856a..bf3f69740 100644 --- a/modules/modelSetup/SanaFineTuneSetup.py +++ b/modules/modelSetup/SanaFineTuneSetup.py @@ -10,25 +10,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.SANA, TrainingMethod.FINE_TUNE) class SanaFineTuneSetup( BaseSanaSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: SanaModel, @@ -87,9 +73,12 @@ def setup_train_device( or config.train_any_embedding() \ or not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if config.text_encoder.train: model.text_encoder.train() diff --git a/modules/modelSetup/SanaLoRASetup.py b/modules/modelSetup/SanaLoRASetup.py index 70f378a9f..73ced8854 100644 --- a/modules/modelSetup/SanaLoRASetup.py +++ b/modules/modelSetup/SanaLoRASetup.py @@ -11,25 +11,11 @@ from modules.util.torch_util import state_dict_has_prefix from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.SANA, TrainingMethod.LORA) class SanaLoRASetup( BaseSanaSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: SanaModel, @@ -113,9 +99,12 @@ def setup_train_device( or config.train_any_embedding() \ or not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if config.text_encoder.train: model.text_encoder.train() diff --git a/modules/modelSetup/StableDiffusion3EmbeddingSetup.py b/modules/modelSetup/StableDiffusion3EmbeddingSetup.py index af857f996..437f17b38 100644 --- a/modules/modelSetup/StableDiffusion3EmbeddingSetup.py +++ b/modules/modelSetup/StableDiffusion3EmbeddingSetup.py @@ -9,26 +9,12 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_3, TrainingMethod.EMBEDDING) @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_35, TrainingMethod.EMBEDDING) class StableDiffusion3EmbeddingSetup( BaseStableDiffusion3Setup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: StableDiffusion3Model, @@ -97,11 +83,16 @@ def setup_train_device( ): vae_on_train_device = not config.latent_caching - model.text_encoder_1_to(self.train_device if config.text_encoder.train_embedding else self.temp_device) - model.text_encoder_2_to(self.train_device if config.text_encoder_2.train_embedding else self.temp_device) - model.text_encoder_3_to(self.train_device if config.text_encoder_3.train_embedding else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if config.text_encoder.train_embedding: + parts.append("text_encoder") + if config.text_encoder_2.train_embedding: + parts.append("text_encoder_2") + if config.text_encoder_3.train_embedding: + parts.append("text_encoder_3") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if model.text_encoder_1 is not None: model.text_encoder_1.eval() diff --git a/modules/modelSetup/StableDiffusion3FineTuneSetup.py b/modules/modelSetup/StableDiffusion3FineTuneSetup.py index 0c93025ed..e3e1d74f3 100644 --- a/modules/modelSetup/StableDiffusion3FineTuneSetup.py +++ b/modules/modelSetup/StableDiffusion3FineTuneSetup.py @@ -10,26 +10,12 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_3, TrainingMethod.FINE_TUNE) @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_35, TrainingMethod.FINE_TUNE) class StableDiffusion3FineTuneSetup( BaseStableDiffusion3Setup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: StableDiffusion3Model, @@ -117,11 +103,16 @@ def setup_train_device( config.train_text_encoder_3_or_embedding() \ or not config.latent_caching - model.text_encoder_1_to(self.train_device if text_encoder_1_on_train_device else self.temp_device) - model.text_encoder_2_to(self.train_device if text_encoder_2_on_train_device else self.temp_device) - model.text_encoder_3_to(self.train_device if text_encoder_3_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_1_on_train_device: + parts.append("text_encoder") + if text_encoder_2_on_train_device: + parts.append("text_encoder_2") + if text_encoder_3_on_train_device: + parts.append("text_encoder_3") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if model.text_encoder_1: if config.text_encoder.train: diff --git a/modules/modelSetup/StableDiffusion3LoRASetup.py b/modules/modelSetup/StableDiffusion3LoRASetup.py index db851027e..4dd2168bc 100644 --- a/modules/modelSetup/StableDiffusion3LoRASetup.py +++ b/modules/modelSetup/StableDiffusion3LoRASetup.py @@ -11,26 +11,12 @@ from modules.util.torch_util import state_dict_has_prefix from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_3, TrainingMethod.LORA) @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_35, TrainingMethod.LORA) class StableDiffusion3LoRASetup( BaseStableDiffusion3Setup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: StableDiffusion3Model, @@ -176,11 +162,16 @@ def setup_train_device( config.train_text_encoder_3_or_embedding() \ or not config.latent_caching - model.text_encoder_1_to(self.train_device if text_encoder_1_on_train_device else self.temp_device) - model.text_encoder_2_to(self.train_device if text_encoder_2_on_train_device else self.temp_device) - model.text_encoder_3_to(self.train_device if text_encoder_3_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_1_on_train_device: + parts.append("text_encoder") + if text_encoder_2_on_train_device: + parts.append("text_encoder_2") + if text_encoder_3_on_train_device: + parts.append("text_encoder_3") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if model.text_encoder_1: if config.text_encoder.train: diff --git a/modules/modelSetup/StableDiffusionEmbeddingSetup.py b/modules/modelSetup/StableDiffusionEmbeddingSetup.py index b2f339ef2..2e53f5c98 100644 --- a/modules/modelSetup/StableDiffusionEmbeddingSetup.py +++ b/modules/modelSetup/StableDiffusionEmbeddingSetup.py @@ -9,8 +9,6 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_15, TrainingMethod.EMBEDDING) @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_15_INPAINTING, TrainingMethod.EMBEDDING) @@ -23,18 +21,6 @@ class StableDiffusionEmbeddingSetup( BaseStableDiffusionSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: StableDiffusionModel, @@ -84,10 +70,12 @@ def setup_train_device( ): vae_on_train_device = self.debug_mode or not config.latent_caching - model.text_encoder_to(self.train_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.unet_to(self.train_device) - model.depth_estimator_to(self.temp_device) + parts = ["unet", "text_encoder"] + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) + if model.depth_estimator is not None: + model.depth_estimator.to(self.temp_device) model.text_encoder.eval() model.vae.eval() diff --git a/modules/modelSetup/StableDiffusionFineTuneSetup.py b/modules/modelSetup/StableDiffusionFineTuneSetup.py index e32008411..57e0dc632 100644 --- a/modules/modelSetup/StableDiffusionFineTuneSetup.py +++ b/modules/modelSetup/StableDiffusionFineTuneSetup.py @@ -10,8 +10,6 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_15, TrainingMethod.FINE_TUNE) @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_15_INPAINTING, TrainingMethod.FINE_TUNE) @@ -24,18 +22,6 @@ class StableDiffusionFineTuneSetup( BaseStableDiffusionSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: StableDiffusionModel, @@ -102,10 +88,14 @@ def setup_train_device( or config.train_any_embedding() \ or not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.unet_to(self.train_device) - model.depth_estimator_to(self.temp_device) + parts = ["unet"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) + if model.depth_estimator is not None: + model.depth_estimator.to(self.temp_device) if config.text_encoder.train: model.text_encoder.train() diff --git a/modules/modelSetup/StableDiffusionFineTuneVaeSetup.py b/modules/modelSetup/StableDiffusionFineTuneVaeSetup.py index 168020b53..8f947a3d6 100644 --- a/modules/modelSetup/StableDiffusionFineTuneVaeSetup.py +++ b/modules/modelSetup/StableDiffusionFineTuneVaeSetup.py @@ -23,18 +23,6 @@ class StableDiffusionFineTuneVaeSetup( BaseStableDiffusionSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: StableDiffusionModel, @@ -67,9 +55,7 @@ def setup_train_device( model: StableDiffusionModel, config: TrainConfig, ): - model.text_encoder.to(self.temp_device) - model.vae.to(self.train_device) - model.unet.to(self.temp_device) + model.materialize_only("vae") if model.depth_estimator is not None: model.depth_estimator.to(self.temp_device) diff --git a/modules/modelSetup/StableDiffusionLoRASetup.py b/modules/modelSetup/StableDiffusionLoRASetup.py index b6c4e8746..2ecc660e9 100644 --- a/modules/modelSetup/StableDiffusionLoRASetup.py +++ b/modules/modelSetup/StableDiffusionLoRASetup.py @@ -11,8 +11,6 @@ from modules.util.torch_util import state_dict_has_prefix from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_15, TrainingMethod.LORA) @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_15_INPAINTING, TrainingMethod.LORA) @@ -25,18 +23,6 @@ class StableDiffusionLoRASetup( BaseStableDiffusionSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: StableDiffusionModel, @@ -125,10 +111,14 @@ def setup_train_device( or config.train_any_embedding() \ or not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.unet_to(self.train_device) - model.depth_estimator_to(self.temp_device) + parts = ["unet"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) + if model.depth_estimator is not None: + model.depth_estimator.to(self.temp_device) if config.text_encoder.train: model.text_encoder.train() diff --git a/modules/modelSetup/StableDiffusionXLEmbeddingSetup.py b/modules/modelSetup/StableDiffusionXLEmbeddingSetup.py index a2fe88ac7..acf1cc5f1 100644 --- a/modules/modelSetup/StableDiffusionXLEmbeddingSetup.py +++ b/modules/modelSetup/StableDiffusionXLEmbeddingSetup.py @@ -9,26 +9,12 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_XL_10_BASE, TrainingMethod.EMBEDDING) @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_XL_10_BASE_INPAINTING, TrainingMethod.EMBEDDING) class StableDiffusionXLEmbeddingSetup( BaseStableDiffusionXLSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: StableDiffusionXLModel, @@ -87,10 +73,14 @@ def setup_train_device( ): vae_on_train_device = not config.latent_caching - model.text_encoder_1_to(self.train_device if config.text_encoder.train_embedding else self.temp_device) - model.text_encoder_2_to(self.train_device if config.text_encoder_2.train_embedding else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.unet_to(self.train_device) + parts = ["unet"] + if config.text_encoder.train_embedding: + parts.append("text_encoder") + if config.text_encoder_2.train_embedding: + parts.append("text_encoder_2") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) model.text_encoder_1.eval() model.text_encoder_2.eval() diff --git a/modules/modelSetup/StableDiffusionXLFineTuneSetup.py b/modules/modelSetup/StableDiffusionXLFineTuneSetup.py index 7ed64132c..0cfe94d22 100644 --- a/modules/modelSetup/StableDiffusionXLFineTuneSetup.py +++ b/modules/modelSetup/StableDiffusionXLFineTuneSetup.py @@ -10,26 +10,12 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_XL_10_BASE, TrainingMethod.FINE_TUNE) @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_XL_10_BASE_INPAINTING, TrainingMethod.FINE_TUNE) class StableDiffusionXLFineTuneSetup( BaseStableDiffusionXLSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: StableDiffusionXLModel, @@ -111,10 +97,14 @@ def setup_train_device( or config.train_any_embedding() \ or not config.latent_caching - model.text_encoder_1_to(self.train_device if text_encoder_1_on_train_device else self.temp_device) - model.text_encoder_2_to(self.train_device if text_encoder_2_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.unet_to(self.train_device) + parts = ["unet"] + if text_encoder_1_on_train_device: + parts.append("text_encoder") + if text_encoder_2_on_train_device: + parts.append("text_encoder_2") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if config.text_encoder.train: model.text_encoder_1.train() diff --git a/modules/modelSetup/StableDiffusionXLLoRASetup.py b/modules/modelSetup/StableDiffusionXLLoRASetup.py index 78aeb7f63..afa1b0cb2 100644 --- a/modules/modelSetup/StableDiffusionXLLoRASetup.py +++ b/modules/modelSetup/StableDiffusionXLLoRASetup.py @@ -11,26 +11,12 @@ from modules.util.torch_util import state_dict_has_prefix from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_XL_10_BASE, TrainingMethod.LORA) @factory.register(BaseModelSetup, ModelType.STABLE_DIFFUSION_XL_10_BASE_INPAINTING, TrainingMethod.LORA) class StableDiffusionXLLoRASetup( BaseStableDiffusionXLSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: StableDiffusionXLModel, @@ -142,10 +128,14 @@ def setup_train_device( config.train_text_encoder_2_or_embedding() \ or not config.latent_caching - model.text_encoder_1_to(self.train_device if text_encoder_1_on_train_device else self.temp_device) - model.text_encoder_2_to(self.train_device if text_encoder_2_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.unet_to(self.train_device) + parts = ["unet"] + if text_encoder_1_on_train_device: + parts.append("text_encoder") + if text_encoder_2_on_train_device: + parts.append("text_encoder_2") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) if config.text_encoder.train: model.text_encoder_1.train() diff --git a/modules/modelSetup/WuerstchenEmbeddingSetup.py b/modules/modelSetup/WuerstchenEmbeddingSetup.py index 7d0d65506..7e6c8ee30 100644 --- a/modules/modelSetup/WuerstchenEmbeddingSetup.py +++ b/modules/modelSetup/WuerstchenEmbeddingSetup.py @@ -9,26 +9,12 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.WUERSTCHEN_2, TrainingMethod.EMBEDDING) @factory.register(BaseModelSetup, ModelType.STABLE_CASCADE_1, TrainingMethod.EMBEDDING) class WuerstchenEmbeddingSetup( BaseWuerstchenSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: WuerstchenModel, @@ -78,14 +64,12 @@ def setup_train_device( ): effnet_on_train_device = not config.latent_caching - if model.model_type.is_wuerstchen_v2(): - model.decoder_text_encoder_to(self.temp_device) - model.decoder_decoder_to(self.temp_device) - model.decoder_vqgan_to(self.temp_device) - model.effnet_encoder_to(self.train_device if effnet_on_train_device else self.temp_device) - - model.prior_text_encoder_to(self.train_device) - model.prior_prior_to(self.train_device) + # decoder/decoder_text_encoder/decoder_vqgan are never needed during prior training; materialize_only() + # evicts them (decoder_text_encoder only exists in model_parts() for Wuerstchen v2, not Stable Cascade) + parts = ["prior", "text_encoder"] + if effnet_on_train_device: + parts.append("effnet_encoder") + model.materialize_only(*parts) if model.model_type.is_wuerstchen_v2(): model.decoder_text_encoder.eval() diff --git a/modules/modelSetup/WuerstchenFineTuneSetup.py b/modules/modelSetup/WuerstchenFineTuneSetup.py index 4f6262d63..830709a68 100644 --- a/modules/modelSetup/WuerstchenFineTuneSetup.py +++ b/modules/modelSetup/WuerstchenFineTuneSetup.py @@ -10,26 +10,12 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.WUERSTCHEN_2, TrainingMethod.FINE_TUNE) @factory.register(BaseModelSetup, ModelType.STABLE_CASCADE_1, TrainingMethod.FINE_TUNE) class WuerstchenFineTuneSetup( BaseWuerstchenSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: WuerstchenModel, @@ -86,20 +72,19 @@ def setup_train_device( config: TrainConfig, ): effnet_on_train_device = not config.latent_caching - - if model.model_type.is_wuerstchen_v2(): - model.decoder_text_encoder_to(self.temp_device) - model.decoder_decoder_to(self.temp_device) - model.decoder_vqgan_to(self.temp_device) - model.effnet_encoder_to(self.train_device if effnet_on_train_device else self.temp_device) - text_encoder_on_train_device = \ config.text_encoder.train \ or config.train_any_embedding() \ or not config.latent_caching - model.prior_text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.prior_prior_to(self.train_device) + # decoder/decoder_text_encoder/decoder_vqgan are never needed during prior training; materialize_only() + # evicts them (decoder_text_encoder only exists in model_parts() for Wuerstchen v2, not Stable Cascade) + parts = ["prior"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if effnet_on_train_device: + parts.append("effnet_encoder") + model.materialize_only(*parts) if model.model_type.is_wuerstchen_v2(): model.decoder_text_encoder.eval() diff --git a/modules/modelSetup/WuerstchenLoRASetup.py b/modules/modelSetup/WuerstchenLoRASetup.py index 1bdc15f3a..c9b780e9d 100644 --- a/modules/modelSetup/WuerstchenLoRASetup.py +++ b/modules/modelSetup/WuerstchenLoRASetup.py @@ -11,26 +11,12 @@ from modules.util.torch_util import state_dict_has_prefix from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.WUERSTCHEN_2, TrainingMethod.LORA) @factory.register(BaseModelSetup, ModelType.STABLE_CASCADE_1, TrainingMethod.LORA) class WuerstchenLoRASetup( BaseWuerstchenSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: WuerstchenModel, @@ -113,20 +99,19 @@ def setup_train_device( config: TrainConfig, ): effnet_on_train_device = not config.latent_caching - - if model.model_type.is_wuerstchen_v2(): - model.decoder_text_encoder_to(self.temp_device) - model.decoder_decoder_to(self.temp_device) - model.decoder_vqgan_to(self.temp_device) - model.effnet_encoder_to(self.train_device if effnet_on_train_device else self.temp_device) - text_encoder_on_train_device = \ config.text_encoder.train \ or config.train_any_embedding() \ or not config.latent_caching - model.prior_text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.prior_prior_to(self.train_device) + # decoder/decoder_text_encoder/decoder_vqgan are never needed during prior training; materialize_only() + # evicts them (decoder_text_encoder only exists in model_parts() for Wuerstchen v2, not Stable Cascade) + parts = ["prior"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if effnet_on_train_device: + parts.append("effnet_encoder") + model.materialize_only(*parts) if model.model_type.is_wuerstchen_v2(): model.decoder_text_encoder.eval() diff --git a/modules/modelSetup/ZImageFineTuneSetup.py b/modules/modelSetup/ZImageFineTuneSetup.py index 6f2642de7..8fd711ae4 100644 --- a/modules/modelSetup/ZImageFineTuneSetup.py +++ b/modules/modelSetup/ZImageFineTuneSetup.py @@ -10,25 +10,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.Z_IMAGE, TrainingMethod.FINE_TUNE) class ZImageFineTuneSetup( BaseZImageSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: ZImageModel, @@ -67,9 +53,12 @@ def setup_train_device( vae_on_train_device = not config.latent_caching text_encoder_on_train_device = not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) model.text_encoder.eval() model.vae.eval() diff --git a/modules/modelSetup/ZImageLoRASetup.py b/modules/modelSetup/ZImageLoRASetup.py index 85cd789d6..7b360a995 100644 --- a/modules/modelSetup/ZImageLoRASetup.py +++ b/modules/modelSetup/ZImageLoRASetup.py @@ -10,25 +10,11 @@ from modules.util.optimizer_util import init_model_parameters from modules.util.TrainProgress import TrainProgress -import torch - @factory.register(BaseModelSetup, ModelType.Z_IMAGE, TrainingMethod.LORA) class ZImageLoRASetup( BaseZImageSetup, ): - def __init__( - self, - train_device: torch.device, - temp_device: torch.device, - debug_mode: bool, - ): - super().__init__( - train_device=train_device, - temp_device=temp_device, - debug_mode=debug_mode, - ) - def create_parameters( self, model: ZImageModel, @@ -80,9 +66,12 @@ def setup_train_device( vae_on_train_device = not config.latent_caching text_encoder_on_train_device = not config.latent_caching - model.text_encoder_to(self.train_device if text_encoder_on_train_device else self.temp_device) - model.vae_to(self.train_device if vae_on_train_device else self.temp_device) - model.transformer_to(self.train_device) + parts = ["transformer"] + if text_encoder_on_train_device: + parts.append("text_encoder") + if vae_on_train_device: + parts.append("vae") + model.materialize_only(*parts) model.text_encoder.eval() model.vae.eval() diff --git a/modules/trainer/GenericTrainer.py b/modules/trainer/GenericTrainer.py index dd17ad76e..5b516cf39 100644 --- a/modules/trainer/GenericTrainer.py +++ b/modules/trainer/GenericTrainer.py @@ -139,9 +139,7 @@ def start(self): self.model_setup.setup_optimizations(self.model, self.config) self.model_setup.setup_train_device(self.model, self.config) self.model_setup.setup_model(self.model, self.config) - self.model.to(self.temp_device) self.model.eval() - torch_gc() self.callbacks.on_update_status("creating the data loader/caching") @@ -253,7 +251,6 @@ def on_sample_custom(sampler_output: ModelSamplerOutput): on_sample = on_sample_custom if is_custom_sample else on_sample_default on_update_progress = self.callbacks.on_update_sample_custom_progress if is_custom_sample else self.callbacks.on_update_sample_default_progress - self.model.to(self.temp_device) self.model.eval() sample_config = copy.copy(sample_config) @@ -717,7 +714,7 @@ def sample_commands_fun(): backup = self.commands.get_and_reset_backup_command() save = self.commands.get_and_reset_save_command() if multi.is_master() and (backup or save): - self.model.to(self.temp_device) + self.model.evict() if backup: self.__backup(train_progress, True, step_tqdm.write) if save: @@ -842,7 +839,7 @@ def sample_commands_fun(): def end(self): if self.one_step_trained: - self.model.to(self.temp_device) + self.model.evict() if self.config.backup_before_save and multi.is_master(): self.__backup(self.model.train_progress) @@ -875,7 +872,7 @@ def end(self): ) if self.model is not None: - self.model.to(self.temp_device) + self.model.evict() if multi.is_master(): self.tensorboard.close() diff --git a/modules/ui/SampleWindowController.py b/modules/ui/SampleWindowController.py index d1d24d643..eae78c015 100644 --- a/modules/ui/SampleWindowController.py +++ b/modules/ui/SampleWindowController.py @@ -96,7 +96,6 @@ def load_model(self) -> BaseModel: model_setup.setup_optimizations(model, self.initial_train_config) model_setup.setup_train_device(model, self.initial_train_config) model_setup.setup_model(model, self.initial_train_config) - model.to(torch.device(self.initial_train_config.temp_device)) return model @@ -145,3 +144,7 @@ def do_sample(self, on_sample, on_update_progress): on_sample=on_sample, on_update_progress=on_update_progress, ) + + # the sampler materializes parts on demand and no longer self-evicts; + # release VRAM now that this standalone sample window is idle again + self.model.evict() diff --git a/modules/util/LayerOffloadConductor.py b/modules/util/LayerOffloadConductor.py index a69094c75..684bb6972 100644 --- a/modules/util/LayerOffloadConductor.py +++ b/modules/util/LayerOffloadConductor.py @@ -618,63 +618,69 @@ def __init__( def offload_activated(self) -> bool: return self.__offload_activations or self.__offload_layers - def to(self, device: torch.device): + def evict(self): torch_gc() self.__wait_all_layer_transfers() self.__wait_all_activation_transfers() - if device_equals(device, self.__temp_device): - log("to temp device") - - # deallocate the cache before to take advantage of the gc - self.__train_device_layer_allocator.deallocate_cache() - self.__temp_device_layer_allocator.deallocate_cache() - self.__temp_device_activations_allocator.deallocate_cache() - - self.__module_to_device_except_layers(self.__temp_device) - for layer_index, layer in enumerate(self.__layers): - self.__layers[layer_index].to(self.__temp_device) - for module in layer.modules(): - offload_quantized(module, self.__temp_device, allocator=clone_tensor_allocator) - self.__layer_device_map[layer_index] = None - - self.__is_active = False - - elif device_equals(device, self.__train_device): - log("to train device") - - self.__offload_strategy = LayerOffloadStrategy(self.__layers, self.__layer_offload_fraction) - - self.__train_device_layer_allocator.allocate_cache( - self.__layers, self.__offload_strategy.max_loaded_bytes) - self.__temp_device_layer_allocator.allocate_cache( - self.__layers, self.__offload_strategy.max_offloaded_bytes) - self.__module_to_device_except_layers(self.__train_device) - - # move all layers to the train device, then move offloadable tensors back to the temp device - for layer_index, layer in enumerate(self.__layers): - if self.__layer_device_map[layer_index] is None: - log(f"layer {layer_index} to train device") - layer.to(self.__train_device) - - if layer_index in self.__offload_strategy.initial_loaded_layers: - allocator = self.__train_device_layer_allocator.get_allocator( - layer_index, allocate_forward=True) - for module in layer.modules(): - offload_quantized(module, self.__train_device, allocator=allocator.allocate_like) - self.__layer_device_map[layer_index] = self.__train_device - else: - allocator = self.__temp_device_layer_allocator.get_allocator(layer_index, allocate_forward=True) - for module in layer.modules(): - offload_quantized(module, self.__temp_device, allocator=allocator.allocate_like) - self.__layer_device_map[layer_index] = self.__temp_device - - if self.__async_transfer: - event = SyncEvent(self.__train_stream.record_event(), f"train on {self.__train_device}") - self.__layer_train_event_map[layer_index] = event - - self.__is_active = True + log("to temp device") + + # deallocate the cache before to take advantage of the gc + self.__train_device_layer_allocator.deallocate_cache() + self.__temp_device_layer_allocator.deallocate_cache() + self.__temp_device_activations_allocator.deallocate_cache() + + self.__module_to_device_except_layers(self.__temp_device) + for layer_index, layer in enumerate(self.__layers): + self.__layers[layer_index].to(self.__temp_device) + for module in layer.modules(): + offload_quantized(module, self.__temp_device, allocator=clone_tensor_allocator) + self.__layer_device_map[layer_index] = None + + self.__is_active = False + + torch_gc() + + def materialize(self): + torch_gc() + + self.__wait_all_layer_transfers() + self.__wait_all_activation_transfers() + + log("to train device") + + self.__offload_strategy = LayerOffloadStrategy(self.__layers, self.__layer_offload_fraction) + + self.__train_device_layer_allocator.allocate_cache( + self.__layers, self.__offload_strategy.max_loaded_bytes) + self.__temp_device_layer_allocator.allocate_cache( + self.__layers, self.__offload_strategy.max_offloaded_bytes) + self.__module_to_device_except_layers(self.__train_device) + + # move all layers to the train device, then move offloadable tensors back to the temp device + for layer_index, layer in enumerate(self.__layers): + if self.__layer_device_map[layer_index] is None: + log(f"layer {layer_index} to train device") + layer.to(self.__train_device) + + if layer_index in self.__offload_strategy.initial_loaded_layers: + allocator = self.__train_device_layer_allocator.get_allocator( + layer_index, allocate_forward=True) + for module in layer.modules(): + offload_quantized(module, self.__train_device, allocator=allocator.allocate_like) + self.__layer_device_map[layer_index] = self.__train_device + else: + allocator = self.__temp_device_layer_allocator.get_allocator(layer_index, allocate_forward=True) + for module in layer.modules(): + offload_quantized(module, self.__temp_device, allocator=allocator.allocate_like) + self.__layer_device_map[layer_index] = self.__temp_device + + if self.__async_transfer: + event = SyncEvent(self.__train_stream.record_event(), f"train on {self.__train_device}") + self.__layer_train_event_map[layer_index] = event + + self.__is_active = True torch_gc() diff --git a/modules/util/checkpointing_util.py b/modules/util/checkpointing_util.py index 1f669e9a5..2df40065e 100644 --- a/modules/util/checkpointing_util.py +++ b/modules/util/checkpointing_util.py @@ -243,13 +243,14 @@ def enable_checkpointing( part: TrainModelPartConfig, compile: bool, lists, # if there are multiple entries in this list, they must be in the exact order they are executed - otherwise offloading fails - offload_enabled: bool = True, + supports_offloading: bool = True, ) -> LayerOffloadConductor | None: if not part.checkpointing_or_offloading_enabled() and not compile: return None - # a conductor exists iff this part actually offloads (and the component supports conductor offloading) - offload = offload_enabled and part.offloading_enabled() + # a conductor exists iff this part actually offloads: the user enabled it (part.offloading_enabled()) and the + # architecture can be driven by the conductor (supports_offloading). + offload = supports_offloading and part.offloading_enabled() conductor = LayerOffloadConductor(model, config, part) if offload else None checkpointing = part.checkpointing_enabled() @@ -298,12 +299,12 @@ def enable_checkpointing_for_basic_transformer_blocks( model: nn.Module, config: TrainConfig, part: TrainModelPartConfig, - offload_enabled: bool, + supports_offloading: bool = True, ) -> LayerOffloadConductor | None: return enable_checkpointing(model, config, part, config.compile, [ (BasicTransformerBlock , []), ], - offload_enabled = offload_enabled, + supports_offloading = supports_offloading, ) def enable_checkpointing_for_clip_encoder_layers( @@ -313,7 +314,7 @@ def enable_checkpointing_for_clip_encoder_layers( ): return enable_checkpointing(model, config, part, False, [ (CLIPEncoderLayer, []), # No activation offloading for text encoders, because the output might be taken from the middle of the network - ], offload_enabled=False) # CLIP is non-offloadable; keep it plain-checkpointed so a migrated offload_fraction can't build a self-activating conductor + ], supports_offloading=False) # CLIP is non-offloadable; keep it plain-checkpointed so a migrated offload_fraction can't build a self-activating conductor def enable_checkpointing_for_t5_encoder_layers( model: nn.Module, diff --git a/modules/util/dtype_util.py b/modules/util/dtype_util.py index b6d6b2643..7dde18d45 100644 --- a/modules/util/dtype_util.py +++ b/modules/util/dtype_util.py @@ -20,7 +20,7 @@ def create_autocast_context( device: torch.device, train_dtype: DataType | None, enable_autocast_cache: bool, -) -> tuple[torch.autocast | nullcontext, DataType]: +) -> torch.autocast | nullcontext: torch_train_dtype = train_dtype.torch_dtype() if torch_train_dtype in (torch.float16, torch.bfloat16): @@ -34,13 +34,13 @@ def create_autocast_context( # MPS additionally needs macOS >= 14. print(f"Warning: Mixed precision training is untested on device type '{device.type}'.") return torch.autocast(device_type=device.type, dtype=torch_train_dtype, - cache_enabled=enable_autocast_cache), train_dtype + cache_enabled=enable_autocast_cache) elif device.type == "cuda": # float32/tfloat32 on CUDA (and ROCm, which also reports device type "cuda"): # CUDA accepts float32 as an autocast dtype and upcasts lower-precision weights # on the fly (this is undocumented but works). return torch.autocast(device_type=device.type, dtype=torch_train_dtype, - cache_enabled=enable_autocast_cache), train_dtype + cache_enabled=enable_autocast_cache) else: # float32/tfloat32 on a non-CUDA backend (cpu, mps, xpu, ...): those backends # reject fp32 autocast, so disable autocast and let the model run at its weight @@ -49,7 +49,7 @@ def create_autocast_context( print("Warning: float32 training does not upcast lower-precision weights on this device " "(only CUDA can autocast to float32); the model runs at its weight dtype. " "Set the weight data types to float32 for full precision.") - return torch.autocast(device_type=device.type, enabled=False), train_dtype + return torch.autocast(device_type=device.type, enabled=False) def disable_fp16_autocast_context(