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..660d9ecd0 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("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..3da261272 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("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..11815e8df 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("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..1184a3d17 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("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..02a35899f 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("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..5c0e5ece9 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("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..f6b3c8f26 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("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..40ee02e27 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("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..1e7ff0e3f 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("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..944dd70a9 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("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..2ebad9c6f 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("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..fd929b672 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("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..ee6a246be 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("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..7e1da8b10 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("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..e481afba8 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("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..4ed9e1e54 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("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..21728dc75 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 torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -96,17 +97,80 @@ 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() + # print_fragmentation(f"evict({', '.join(parts) or 'all'})") + + 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: + conductor.to(device) + 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..974d3482a 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,8 +126,7 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, diff --git a/modules/modelSampler/ChromaSampler.py b/modules/modelSampler/ChromaSampler.py index bc89e5cab..15441b164 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,8 +145,7 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, diff --git a/modules/modelSampler/ErnieSampler.py b/modules/modelSampler/ErnieSampler.py index 2122e5343..78aee2405 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,8 +120,7 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, diff --git a/modules/modelSampler/Flux2Sampler.py b/modules/modelSampler/Flux2Sampler.py index 0a4cca9b9..6653840fd 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,8 +145,7 @@ def __sample_base( image = image_processor.postprocess(image, output_type='pil') - self.model.vae_to(self.temp_device) - torch_gc() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, diff --git a/modules/modelSampler/FluxSampler.py b/modules/modelSampler/FluxSampler.py index fbe532053..a7fb6a34d 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,8 +158,7 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, @@ -222,7 +215,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 +289,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 +300,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 +327,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 +362,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 +369,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,8 +377,7 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, diff --git a/modules/modelSampler/HiDreamSampler.py b/modules/modelSampler/HiDreamSampler.py index ee7569ccf..eefc3a358 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,8 +147,7 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, diff --git a/modules/modelSampler/HunyuanVideoSampler.py b/modules/modelSampler/HunyuanVideoSampler.py index c12056c92..54f508f1d 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,19 +135,15 @@ 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() + self.model.evict() is_image = image.shape[2] == 1 if is_image: diff --git a/modules/modelSampler/IdeogramSampler.py b/modules/modelSampler/IdeogramSampler.py index 0a5dd420b..cfc5862b0 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,8 +198,7 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, diff --git a/modules/modelSampler/Krea2Sampler.py b/modules/modelSampler/Krea2Sampler.py index 20aecd5b1..e5f2a68c4 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,8 +135,7 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, diff --git a/modules/modelSampler/PixArtAlphaSampler.py b/modules/modelSampler/PixArtAlphaSampler.py index 9adddf3f2..4f21e5250 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,8 +148,7 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, diff --git a/modules/modelSampler/QwenSampler.py b/modules/modelSampler/QwenSampler.py index 4ca604102..4ac4c3c17 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,8 +145,7 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, diff --git a/modules/modelSampler/SanaSampler.py b/modules/modelSampler/SanaSampler.py index 6251ab87e..17d0645b6 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,8 +133,7 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, diff --git a/modules/modelSampler/StableDiffusion3Sampler.py b/modules/modelSampler/StableDiffusion3Sampler.py index fe159e914..eaea37f06 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,8 +145,7 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, diff --git a/modules/modelSampler/StableDiffusionSampler.py b/modules/modelSampler/StableDiffusionSampler.py index 92bfcd759..1881a15bf 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,8 +160,7 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, @@ -223,7 +215,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 +269,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 +286,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 +314,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 +351,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,8 +360,7 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, diff --git a/modules/modelSampler/StableDiffusionVaeSampler.py b/modules/modelSampler/StableDiffusionVaeSampler.py index 254c73fdd..325a1d667 100644 --- a/modules/modelSampler/StableDiffusionVaeSampler.py +++ b/modules/modelSampler/StableDiffusionVaeSampler.py @@ -63,13 +63,13 @@ 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("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) + self.model.evict("vae") 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..e351dac78 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,8 +183,7 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, @@ -247,7 +239,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 +295,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 +317,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 +378,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 +414,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,8 +424,7 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, diff --git a/modules/modelSampler/WuerstchenSampler.py b/modules/modelSampler/WuerstchenSampler.py index d679757c2..ac18ae6bc 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,15 +306,14 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, diff --git a/modules/modelSampler/ZImageSampler.py b/modules/modelSampler/ZImageSampler.py index 0cfcb46d6..74cf22bb4 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,17 +114,14 @@ 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() + self.model.evict() return ModelSamplerOutput( file_type=FileType.IMAGE, 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..8c0b9e0aa 100644 --- a/modules/modelSetup/BaseAnimaSetup.py +++ b/modules/modelSetup/BaseAnimaSetup.py @@ -166,9 +166,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..11ab09c82 100644 --- a/modules/modelSetup/BaseChromaSetup.py +++ b/modules/modelSetup/BaseChromaSetup.py @@ -18,7 +18,6 @@ from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -265,10 +264,9 @@ 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") + else: + model.evict() model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseErnieSetup.py b/modules/modelSetup/BaseErnieSetup.py index c412073e8..d5050c49f 100644 --- a/modules/modelSetup/BaseErnieSetup.py +++ b/modules/modelSetup/BaseErnieSetup.py @@ -16,7 +16,6 @@ from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -160,7 +159,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..63e8dfc70 100644 --- a/modules/modelSetup/BaseFlux2Setup.py +++ b/modules/modelSetup/BaseFlux2Setup.py @@ -18,7 +18,6 @@ from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -182,7 +181,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..398d4a1b4 100644 --- a/modules/modelSetup/BaseFluxSetup.py +++ b/modules/modelSetup/BaseFluxSetup.py @@ -326,13 +326,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..301574535 100644 --- a/modules/modelSetup/BaseHiDreamSetup.py +++ b/modules/modelSetup/BaseHiDreamSetup.py @@ -19,7 +19,6 @@ 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 @@ -411,19 +410,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..4f5713008 100644 --- a/modules/modelSetup/BaseHunyuanVideoSetup.py +++ b/modules/modelSetup/BaseHunyuanVideoSetup.py @@ -19,7 +19,6 @@ 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 @@ -294,13 +293,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..7fb1a0bff 100644 --- a/modules/modelSetup/BaseIdeogramSetup.py +++ b/modules/modelSetup/BaseIdeogramSetup.py @@ -16,7 +16,6 @@ from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -204,7 +203,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..e4249e897 100644 --- a/modules/modelSetup/BaseKrea2Setup.py +++ b/modules/modelSetup/BaseKrea2Setup.py @@ -16,7 +16,6 @@ from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -178,10 +177,9 @@ 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") + else: + model.evict() model.eval() - torch_gc() diff --git a/modules/modelSetup/BasePixArtAlphaSetup.py b/modules/modelSetup/BasePixArtAlphaSetup.py index 57bf3a40b..008e9a0c0 100644 --- a/modules/modelSetup/BasePixArtAlphaSetup.py +++ b/modules/modelSetup/BasePixArtAlphaSetup.py @@ -18,7 +18,6 @@ from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -326,10 +325,9 @@ 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") + else: + model.evict() model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseQwenSetup.py b/modules/modelSetup/BaseQwenSetup.py index a618dc28f..0615a3c1e 100644 --- a/modules/modelSetup/BaseQwenSetup.py +++ b/modules/modelSetup/BaseQwenSetup.py @@ -16,7 +16,6 @@ from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -177,10 +176,9 @@ 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") + else: + model.evict() model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseSanaSetup.py b/modules/modelSetup/BaseSanaSetup.py index a96afd770..8718142a7 100644 --- a/modules/modelSetup/BaseSanaSetup.py +++ b/modules/modelSetup/BaseSanaSetup.py @@ -18,7 +18,6 @@ from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -246,10 +245,9 @@ 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") + else: + model.evict() model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseStableDiffusion3Setup.py b/modules/modelSetup/BaseStableDiffusion3Setup.py index 678d727ff..dd7cf0841 100644 --- a/modules/modelSetup/BaseStableDiffusion3Setup.py +++ b/modules/modelSetup/BaseStableDiffusion3Setup.py @@ -19,7 +19,6 @@ 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 @@ -346,16 +345,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..7a3f910c8 100644 --- a/modules/modelSetup/BaseStableDiffusionSetup.py +++ b/modules/modelSetup/BaseStableDiffusionSetup.py @@ -19,7 +19,6 @@ 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 @@ -334,10 +333,11 @@ 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") + else: + model.evict() + 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..54a080579 100644 --- a/modules/modelSetup/BaseStableDiffusionXLSetup.py +++ b/modules/modelSetup/BaseStableDiffusionXLSetup.py @@ -19,7 +19,6 @@ 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.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -379,13 +378,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..d73052f83 100644 --- a/modules/modelSetup/BaseWuerstchenSetup.py +++ b/modules/modelSetup/BaseWuerstchenSetup.py @@ -20,7 +20,6 @@ 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 @@ -345,10 +344,9 @@ 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") + else: + model.evict() model.eval() - torch_gc() diff --git a/modules/modelSetup/BaseZImageSetup.py b/modules/modelSetup/BaseZImageSetup.py index a180f57d2..b77f9bb72 100644 --- a/modules/modelSetup/BaseZImageSetup.py +++ b/modules/modelSetup/BaseZImageSetup.py @@ -17,7 +17,6 @@ from modules.util.config.TrainConfig import TrainConfig from modules.util.dtype_util import create_autocast_context, disable_fp16_autocast_context from modules.util.quantization_util import quantize_layers -from modules.util.torch_util import torch_gc from modules.util.TrainProgress import TrainProgress import torch @@ -160,8 +159,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..e5138fbb4 100644 --- a/modules/trainer/GenericTrainer.py +++ b/modules/trainer/GenericTrainer.py @@ -139,9 +139,8 @@ 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.evict() self.model.eval() - torch_gc() self.callbacks.on_update_status("creating the data loader/caching") @@ -253,7 +252,7 @@ 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.evict() self.model.eval() sample_config = copy.copy(sample_config) @@ -717,7 +716,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 +841,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 +874,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..45a02212f 100644 --- a/modules/ui/SampleWindowController.py +++ b/modules/ui/SampleWindowController.py @@ -96,7 +96,7 @@ 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)) + model.evict() return model diff --git a/modules/ui/TrainUIController.py b/modules/ui/TrainUIController.py index 72623def4..78527d7cd 100644 --- a/modules/ui/TrainUIController.py +++ b/modules/ui/TrainUIController.py @@ -236,6 +236,9 @@ def __training_thread_function(self): self.view.sync_cloud_secrets() error_caught = True traceback.print_exc() + finally: + # print_fragmentation("training run, at stop") + pass trainer.end() @@ -247,6 +250,9 @@ def __training_thread_function(self): torch.clear_autocast_cache() torch_gc() + # after trainer.end() + torch_gc: quantifies the reserved-but-unallocated that persists through unload + # print_fragmentation("training run, after unload") + if error_caught: self.on_update_status("Error: check the console for details") else: diff --git a/modules/util/torch_util.py b/modules/util/torch_util.py index 408100bf9..8c5a28a80 100644 --- a/modules/util/torch_util.py +++ b/modules/util/torch_util.py @@ -233,6 +233,26 @@ def torch_gc(): torch.mps.empty_cache() +def print_fragmentation(label: str): + # Live (not peak) snapshot of allocator fragmentation for the current device. reserved-allocated is the total free + # memory the allocator holds but hasn't handed out; ext_frag is how shattered that free space is -- the fraction + # NOT in the single largest contiguous hole. An OOM is about contiguity, not total free bytes, so ext_frag near 1 + # (free memory in many small holes) is the fragmentation that blocks a large allocation even when free > request. + if not torch.cuda.is_available(): + return + segments = torch.cuda.memory_snapshot() + reserved = sum(s["total_size"] for s in segments) + allocated = sum(b["size"] for s in segments for b in s["blocks"] if b["state"] == "active_allocated") + free_total = reserved - allocated + largest_free = max( + (b["size"] for s in segments for b in s["blocks"] if b["state"] != "active_allocated"), + default=0) + ext_frag = 1 - largest_free / free_total if free_total > 0 else 0.0 + gib = 1024 ** 3 + print(f"[fragmentation] {label}: allocated={allocated / gib:.2f} GiB reserved={reserved / gib:.2f} GiB " + f"free={free_total / gib:.2f} GiB largest_free={largest_free / gib:.2f} GiB ext_frag={ext_frag:.1%}") + + def torch_sync(): if torch.cuda.is_available(): torch.cuda.synchronize()