diff --git a/modules/modelLoader/AnimaModelLoader.py b/modules/modelLoader/AnimaModelLoader.py index f4d3662bc..77f963d3c 100644 --- a/modules/modelLoader/AnimaModelLoader.py +++ b/modules/modelLoader/AnimaModelLoader.py @@ -18,7 +18,6 @@ AutoencoderKLQwenImage, CosmosTransformer3DModel, FlowMatchEulerDiscreteScheduler, - GGUFQuantizationConfig, ) from transformers import Qwen2Tokenizer, Qwen3Model, T5TokenizerFast @@ -71,7 +70,7 @@ def __load_diffusers( subfolder="scheduler", ) - text_encoder = self._load_transformers_sub_module( + text_encoder = self._load_text_encoder( Qwen3Model, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, @@ -86,43 +85,22 @@ def __load_diffusers( torch_dtype=torch.bfloat16, ) - if vae_model_name: #TODO simplify - vae = self._load_diffusers_sub_module( - AutoencoderKLQwenImage, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKLQwenImage, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + vae = self._load_vae( + AutoencoderKLQwenImage, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - if transformer_model_name: - transformer = CosmosTransformer3DModel.from_single_file( - transformer_model_name, - config=base_model_name, - subfolder="transformer", - #avoid loading the transformer in float32: - torch_dtype=torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, - ) - else: - transformer = self._load_diffusers_sub_module( - CosmosTransformer3DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + transformer = self._load_transformer( + CosmosTransformer3DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + config=base_model_name, + ) model.model_type = model_type model.tokenizer = tokenizer diff --git a/modules/modelLoader/ErnieModelLoader.py b/modules/modelLoader/ErnieModelLoader.py index af268c25c..a9f0c9259 100644 --- a/modules/modelLoader/ErnieModelLoader.py +++ b/modules/modelLoader/ErnieModelLoader.py @@ -11,13 +11,10 @@ from modules.util.ModelNames import ModelNames from modules.util.ModelWeightDtypes import ModelWeightDtypes -import torch - from diffusers import ( AutoencoderKLFlux2, ErnieImageTransformer2DModel, FlowMatchEulerDiscreteScheduler, - GGUFQuantizationConfig, ) from transformers import AutoTokenizer, Mistral3Model @@ -56,33 +53,21 @@ def __load_diffusers( vae_model_name: str, quantization: QuantizationConfig, ): - if transformer_model_name: - transformer = ErnieImageTransformer2DModel.from_single_file( - transformer_model_name, - config=base_model_name, - subfolder="transformer", - torch_dtype=torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, - ) - else: - transformer = self._load_diffusers_sub_module( - ErnieImageTransformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + transformer = self._load_transformer( + ErnieImageTransformer2DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + config=base_model_name, + ) tokenizer = AutoTokenizer.from_pretrained( base_model_name, subfolder="tokenizer", ) - text_encoder = self._load_transformers_sub_module( + text_encoder = self._load_text_encoder( Mistral3Model, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, @@ -95,21 +80,13 @@ def __load_diffusers( subfolder="scheduler", ) - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKLFlux2, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKLFlux2, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + vae = self._load_vae( + AutoencoderKLFlux2, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) model.model_type = model_type model.tokenizer = tokenizer diff --git a/modules/modelLoader/Flux2ModelLoader.py b/modules/modelLoader/Flux2ModelLoader.py index 33f3fe518..3b502b8de 100644 --- a/modules/modelLoader/Flux2ModelLoader.py +++ b/modules/modelLoader/Flux2ModelLoader.py @@ -11,13 +11,10 @@ from modules.util.ModelNames import ModelNames from modules.util.ModelWeightDtypes import ModelWeightDtypes -import torch - from diffusers import ( AutoencoderKLFlux2, FlowMatchEulerDiscreteScheduler, Flux2Transformer2DModel, - GGUFQuantizationConfig, ) from transformers import ( Mistral3ForConditionalGeneration, @@ -60,27 +57,14 @@ def __load_diffusers( vae_model_name: str, quantization: QuantizationConfig, ): - if transformer_model_name: - transformer = Flux2Transformer2DModel.from_single_file( - transformer_model_name, - config=base_model_name, - subfolder="transformer", - #avoid loading the transformer in float32: - torch_dtype=torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, - ) - else: - transformer = self._load_diffusers_sub_module( - Flux2Transformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + transformer = self._load_transformer( + Flux2Transformer2DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + config=base_model_name, + ) if transformer.config.num_attention_heads == 48: #Flux2.Dev tokenizer = PixtralProcessor.from_pretrained( @@ -88,7 +72,7 @@ def __load_diffusers( subfolder="tokenizer", ).tokenizer - text_encoder = self._load_transformers_sub_module( + text_encoder = self._load_text_encoder( Mistral3ForConditionalGeneration, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, @@ -100,7 +84,7 @@ def __load_diffusers( base_model_name, subfolder="tokenizer", ) - text_encoder = self._load_transformers_sub_module( + text_encoder = self._load_text_encoder( Qwen3ForCausalLM, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, @@ -113,21 +97,13 @@ def __load_diffusers( subfolder="scheduler", ) - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKLFlux2, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKLFlux2, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + vae = self._load_vae( + AutoencoderKLFlux2, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) model.model_type = model_type model.tokenizer = tokenizer diff --git a/modules/modelLoader/IdeogramModelLoader.py b/modules/modelLoader/IdeogramModelLoader.py index 5a91d51c2..3858ce732 100644 --- a/modules/modelLoader/IdeogramModelLoader.py +++ b/modules/modelLoader/IdeogramModelLoader.py @@ -31,11 +31,12 @@ def __load_internal( model_type: ModelType, weight_dtypes: ModelWeightDtypes, base_model_name: str, + vae_model_name: str, include_unconditional_transformer: bool, quantization: QuantizationConfig, ): if os.path.isfile(os.path.join(base_model_name, "meta.json")): - self.__load_diffusers(model, model_type, weight_dtypes, base_model_name, include_unconditional_transformer, quantization) + self.__load_diffusers(model, model_type, weight_dtypes, base_model_name, vae_model_name, include_unconditional_transformer, quantization) else: raise Exception("not an internal model") @@ -45,6 +46,7 @@ def __load_diffusers( model_type: ModelType, weight_dtypes: ModelWeightDtypes, base_model_name: str, + vae_model_name: str, include_unconditional_transformer: bool, quantization: QuantizationConfig, ): @@ -71,7 +73,7 @@ def __load_diffusers( else: unconditional_transformer = None - text_encoder = self._load_transformers_sub_module( + text_encoder = self._load_text_encoder( Qwen3VLModel, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, @@ -89,12 +91,12 @@ def __load_diffusers( subfolder="scheduler", ) - vae = self._load_diffusers_sub_module( + vae = self._load_vae( AutoencoderKLFlux2, weight_dtypes.vae, weight_dtypes.train_dtype, base_model_name, - "vae", + vae_model_name, ) model.model_type = model_type @@ -129,7 +131,7 @@ def load( try: self.__load_internal( - model, model_type, weight_dtypes, model_names.base_model, + model, model_type, weight_dtypes, model_names.base_model, model_names.vae_model, model_names.include_unconditional_transformer, quantization, ) return @@ -138,7 +140,7 @@ def load( try: self.__load_diffusers( - model, model_type, weight_dtypes, model_names.base_model, + model, model_type, weight_dtypes, model_names.base_model, model_names.vae_model, model_names.include_unconditional_transformer, quantization, ) return diff --git a/modules/modelLoader/ZImageModelLoader.py b/modules/modelLoader/ZImageModelLoader.py index 308232823..f4b0b8db7 100644 --- a/modules/modelLoader/ZImageModelLoader.py +++ b/modules/modelLoader/ZImageModelLoader.py @@ -11,12 +11,9 @@ from modules.util.ModelNames import ModelNames from modules.util.ModelWeightDtypes import ModelWeightDtypes -import torch - from diffusers import ( AutoencoderKL, FlowMatchEulerDiscreteScheduler, - GGUFQuantizationConfig, ZImageTransformer2DModel, ) from transformers import ( @@ -68,7 +65,7 @@ def __load_diffusers( subfolder="scheduler", ) - text_encoder = self._load_transformers_sub_module( + text_encoder = self._load_text_encoder( Qwen3ForCausalLM, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, @@ -76,41 +73,21 @@ def __load_diffusers( "text_encoder", ) - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + vae = self._load_vae( + AutoencoderKL, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - if transformer_model_name: - transformer = ZImageTransformer2DModel.from_single_file( - transformer_model_name, - #avoid loading the transformer in float32: - torch_dtype = torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, - ) - else: - transformer = self._load_diffusers_sub_module( - ZImageTransformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + transformer = self._load_transformer( + ZImageTransformer2DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + ) model.model_type = model_type model.tokenizer = tokenizer diff --git a/modules/modelLoader/chroma/ChromaModelLoader.py b/modules/modelLoader/chroma/ChromaModelLoader.py index 7dcbef794..4a93cc674 100644 --- a/modules/modelLoader/chroma/ChromaModelLoader.py +++ b/modules/modelLoader/chroma/ChromaModelLoader.py @@ -9,13 +9,10 @@ from modules.util.ModelNames import ModelNames from modules.util.ModelWeightDtypes import ModelWeightDtypes -import torch - from diffusers import ( AutoencoderKL, ChromaTransformer2DModel, FlowMatchEulerDiscreteScheduler, - GGUFQuantizationConfig, ) from transformers import T5EncoderModel, T5Tokenizer @@ -63,7 +60,7 @@ def __load_diffusers( subfolder="scheduler", ) - text_encoder = self._load_transformers_sub_module( + text_encoder = self._load_text_encoder( T5EncoderModel, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, @@ -71,41 +68,21 @@ def __load_diffusers( "text_encoder", ) - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + vae = self._load_vae( + AutoencoderKL, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - if transformer_model_name: - transformer = ChromaTransformer2DModel.from_single_file( - transformer_model_name, - #avoid loading the transformer in float32: - torch_dtype = torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, - ) - else: - transformer = self._load_diffusers_sub_module( - ChromaTransformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + transformer = self._load_transformer( + ChromaTransformer2DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + ) model.model_type = model_type model.tokenizer = tokenizer diff --git a/modules/modelLoader/flux/FluxModelLoader.py b/modules/modelLoader/flux/FluxModelLoader.py index d4f21ea2a..02547a950 100644 --- a/modules/modelLoader/flux/FluxModelLoader.py +++ b/modules/modelLoader/flux/FluxModelLoader.py @@ -16,7 +16,6 @@ FlowMatchEulerDiscreteScheduler, FluxPipeline, FluxTransformer2DModel, - GGUFQuantizationConfig, ) from transformers import CLIPTextModel, CLIPTokenizer, T5EncoderModel, T5Tokenizer @@ -81,7 +80,7 @@ def __load_diffusers( ) if include_text_encoder_1: - text_encoder_1 = self._load_transformers_sub_module( + text_encoder_1 = self._load_text_encoder( CLIPTextModel, weight_dtypes.text_encoder, weight_dtypes.train_dtype, @@ -92,7 +91,7 @@ def __load_diffusers( text_encoder_1 = None if include_text_encoder_2: - text_encoder_2 = self._load_transformers_sub_module( + text_encoder_2 = self._load_text_encoder( T5EncoderModel, weight_dtypes.text_encoder_2, weight_dtypes.fallback_train_dtype, @@ -102,41 +101,21 @@ def __load_diffusers( else: text_encoder_2 = None - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + vae = self._load_vae( + AutoencoderKL, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - if transformer_model_name: - transformer = FluxTransformer2DModel.from_single_file( - transformer_model_name, - #avoid loading the transformer in float32: - torch_dtype = torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, - ) - else: - transformer = self._load_diffusers_sub_module( - FluxTransformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + transformer = self._load_transformer( + FluxTransformer2DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + ) model.model_type = model_type model.tokenizer_1 = tokenizer_1 diff --git a/modules/modelLoader/hiDream/HiDreamModelLoader.py b/modules/modelLoader/hiDream/HiDreamModelLoader.py index b3e20f23c..4c0f002bb 100644 --- a/modules/modelLoader/hiDream/HiDreamModelLoader.py +++ b/modules/modelLoader/hiDream/HiDreamModelLoader.py @@ -92,7 +92,7 @@ def __load_diffusers( ) if include_text_encoder_1: - text_encoder_1 = self._load_transformers_sub_module( + text_encoder_1 = self._load_text_encoder( CLIPTextModelWithProjection, weight_dtypes.text_encoder, weight_dtypes.train_dtype, @@ -103,7 +103,7 @@ def __load_diffusers( text_encoder_1 = None if include_text_encoder_2: - text_encoder_2 = self._load_transformers_sub_module( + text_encoder_2 = self._load_text_encoder( CLIPTextModelWithProjection, weight_dtypes.text_encoder_2, weight_dtypes.train_dtype, @@ -114,7 +114,7 @@ def __load_diffusers( text_encoder_2 = None if include_text_encoder_3: - text_encoder_3 = self._load_transformers_sub_module( + text_encoder_3 = self._load_text_encoder( T5EncoderModel, weight_dtypes.text_encoder_3, weight_dtypes.fallback_train_dtype, @@ -126,6 +126,8 @@ def __load_diffusers( if include_text_encoder_4: if text_encoder_4_model_name: + # override repo holds text_encoder_4 at its root, not in a base-model subfolder, so it bypasses + # _load_text_encoder (which always loads from a base-repo subfolder) and loads directly text_encoder_4 = self._load_transformers_sub_module( LlamaForCausalLM, weight_dtypes.text_encoder_4, @@ -133,7 +135,7 @@ def __load_diffusers( text_encoder_4_model_name, ) else: - text_encoder_4 = self._load_transformers_sub_module( + text_encoder_4 = self._load_text_encoder( LlamaForCausalLM, weight_dtypes.text_encoder_4, weight_dtypes.train_dtype, @@ -144,21 +146,13 @@ def __load_diffusers( else: text_encoder_4 = None - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + vae = self._load_vae( + AutoencoderKL, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) transformer = self._load_diffusers_sub_module( HiDreamImageTransformer2DModel, diff --git a/modules/modelLoader/hunyuanVideo/HunyuanVideoModelLoader.py b/modules/modelLoader/hunyuanVideo/HunyuanVideoModelLoader.py index 85c91699b..e1f306d58 100644 --- a/modules/modelLoader/hunyuanVideo/HunyuanVideoModelLoader.py +++ b/modules/modelLoader/hunyuanVideo/HunyuanVideoModelLoader.py @@ -81,7 +81,7 @@ def __load_diffusers( ) if include_text_encoder_1: - text_encoder_1 = self._load_transformers_sub_module( + text_encoder_1 = self._load_text_encoder( LlamaModel, weight_dtypes.text_encoder, weight_dtypes.train_dtype, @@ -92,7 +92,7 @@ def __load_diffusers( text_encoder_1 = None if include_text_encoder_2: - text_encoder_2 = self._load_transformers_sub_module( + text_encoder_2 = self._load_text_encoder( CLIPTextModel, weight_dtypes.text_encoder_2, weight_dtypes.fallback_train_dtype, @@ -102,43 +102,22 @@ def __load_diffusers( else: text_encoder_2 = None - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKLHunyuanVideo, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKLHunyuanVideo, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + vae = self._load_vae( + AutoencoderKLHunyuanVideo, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - if transformer_model_name: - transformer = HunyuanVideoTransformer3DModel.from_single_file( - transformer_model_name, - config=base_model_name, - subfolder="transformer", - #avoid loading the transformer in float32: - torch_dtype = torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization - ) - else: - transformer = self._load_diffusers_sub_module( - HunyuanVideoTransformer3DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + transformer = self._load_transformer( + HunyuanVideoTransformer3DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + config=base_model_name, + ) model.model_type = model_type model.tokenizer_1 = tokenizer_1 diff --git a/modules/modelLoader/krea2/Krea2ModelLoader.py b/modules/modelLoader/krea2/Krea2ModelLoader.py index c3987e97c..df51d2e5e 100644 --- a/modules/modelLoader/krea2/Krea2ModelLoader.py +++ b/modules/modelLoader/krea2/Krea2ModelLoader.py @@ -8,12 +8,9 @@ from modules.util.ModelNames import ModelNames from modules.util.ModelWeightDtypes import ModelWeightDtypes -import torch - from diffusers import ( AutoencoderKLQwenImage, FlowMatchEulerDiscreteScheduler, - GGUFQuantizationConfig, Krea2Transformer2DModel, ) from transformers import Qwen2Tokenizer, Qwen3VLModel @@ -62,7 +59,7 @@ def __load_diffusers( subfolder="scheduler", ) - text_encoder = self._load_transformers_sub_module( + text_encoder = self._load_text_encoder( Qwen3VLModel, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, @@ -70,43 +67,22 @@ def __load_diffusers( "text_encoder", ) - if vae_model_name: #TODO simplify - vae = self._load_diffusers_sub_module( - AutoencoderKLQwenImage, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKLQwenImage, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + vae = self._load_vae( + AutoencoderKLQwenImage, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - if transformer_model_name: - transformer = Krea2Transformer2DModel.from_single_file( - transformer_model_name, - config=base_model_name, - subfolder="transformer", - #avoid loading the transformer in float32: - torch_dtype = torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, - ) - else: - transformer = self._load_diffusers_sub_module( - Krea2Transformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + transformer = self._load_transformer( + Krea2Transformer2DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + config=base_model_name, + ) model.model_type = model_type model.tokenizer = tokenizer diff --git a/modules/modelLoader/mixin/HFModelLoaderMixin.py b/modules/modelLoader/mixin/HFModelLoaderMixin.py index f2f196257..4cdc457bd 100644 --- a/modules/modelLoader/mixin/HFModelLoaderMixin.py +++ b/modules/modelLoader/mixin/HFModelLoaderMixin.py @@ -6,6 +6,7 @@ from modules.util.config.TrainConfig import QuantizationConfig from modules.util.enum.DataType import DataType +from modules.util.ModelWeightDtypes import ModelWeightDtypes from modules.util.quantization_util import ( is_quantized_parameter, replace_linear_with_quantized_layers, @@ -14,6 +15,7 @@ import torch from torch import nn +from diffusers import GGUFQuantizationConfig from transformers.conversion_mapping import get_checkpoint_conversion_mapping from transformers.core_model_loading import rename_source_key @@ -328,3 +330,87 @@ def _convert_diffusers_sub_module_to_dtype( None, quantization, ) + + def _load_transformer( + self, + module_type, + weight_dtypes: ModelWeightDtypes, + base_model_name: str, + transformer_model_name: str, + quantization: QuantizationConfig, + config: str | None = None, + ): + # a single-file (optionally GGUF-quantized) checkpoint is loaded directly, using + # a separate repo to source the model config if the checkpoint doesn't carry one; + # otherwise the transformer is loaded from its subfolder in the base model repo + if transformer_model_name: + single_file_kwargs = {} + if config is not None: + single_file_kwargs["config"] = config + single_file_kwargs["subfolder"] = "transformer" + + transformer = module_type.from_single_file( + transformer_model_name, + **single_file_kwargs, + #avoid loading the transformer in float32: + torch_dtype=torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), + quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, + ) + return self._convert_diffusers_sub_module_to_dtype( + transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, + ) + else: + return self._load_diffusers_sub_module( + module_type, + weight_dtypes.transformer, + weight_dtypes.train_dtype, + base_model_name, + "transformer", + quantization, + ) + + def _load_text_encoder( + self, + module_type, + dtype: DataType, + train_dtype: DataType, + base_model_name: str, + subfolder: str, + ): + # text encoders have no single-file override and always load from their subfolder in the base model + # repo; kept as a per-model entry point alongside _load_transformer / _load_vae. dtype/train_dtype are + # explicit rather than a weight_dtypes bundle since a model can hold several encoders (text_encoder, + # text_encoder_2, ...) with differing dtypes + return self._load_transformers_sub_module( + module_type, + dtype, + train_dtype, + base_model_name, + subfolder, + ) + + def _load_vae( + self, + module_type, + dtype: DataType, + train_dtype: DataType, + base_model_name: str, + vae_model_name: str, + ): + # a separate vae repo overrides the base model's vae subfolder when given. train_dtype is explicit + # since some models (e.g. SDXL) upgrade the vae to fallback_train_dtype to avoid fp16 overflow + if vae_model_name: + return self._load_diffusers_sub_module( + module_type, + dtype, + train_dtype, + vae_model_name, + ) + else: + return self._load_diffusers_sub_module( + module_type, + dtype, + train_dtype, + base_model_name, + "vae", + ) diff --git a/modules/modelLoader/pixartAlpha/PixArtAlphaModelLoader.py b/modules/modelLoader/pixartAlpha/PixArtAlphaModelLoader.py index 467c29c8a..69e280c9d 100644 --- a/modules/modelLoader/pixartAlpha/PixArtAlphaModelLoader.py +++ b/modules/modelLoader/pixartAlpha/PixArtAlphaModelLoader.py @@ -52,7 +52,7 @@ def __load_diffusers( subfolder="scheduler", ) - text_encoder = self._load_transformers_sub_module( + text_encoder = self._load_text_encoder( T5EncoderModel, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, @@ -60,21 +60,13 @@ def __load_diffusers( "text_encoder", ) - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + vae = self._load_vae( + AutoencoderKL, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) transformer = self._load_diffusers_sub_module( Transformer2DModel, diff --git a/modules/modelLoader/qwen/QwenModelLoader.py b/modules/modelLoader/qwen/QwenModelLoader.py index 953f15bfb..77a12bed7 100644 --- a/modules/modelLoader/qwen/QwenModelLoader.py +++ b/modules/modelLoader/qwen/QwenModelLoader.py @@ -8,12 +8,9 @@ from modules.util.ModelNames import ModelNames from modules.util.ModelWeightDtypes import ModelWeightDtypes -import torch - from diffusers import ( AutoencoderKLQwenImage, FlowMatchEulerDiscreteScheduler, - GGUFQuantizationConfig, QwenImageTransformer2DModel, ) from transformers import Qwen2_5_VLForConditionalGeneration, Qwen2Tokenizer @@ -62,7 +59,7 @@ def __load_diffusers( subfolder="scheduler", ) - text_encoder = self._load_transformers_sub_module( + text_encoder = self._load_text_encoder( Qwen2_5_VLForConditionalGeneration, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, @@ -70,43 +67,22 @@ def __load_diffusers( "text_encoder", ) - if vae_model_name: #TODO simplify - vae = self._load_diffusers_sub_module( - AutoencoderKLQwenImage, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKLQwenImage, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + vae = self._load_vae( + AutoencoderKLQwenImage, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) - if transformer_model_name: - transformer = QwenImageTransformer2DModel.from_single_file( - transformer_model_name, - config=base_model_name, - subfolder="transformer", - #avoid loading the transformer in float32: - torch_dtype = torch.bfloat16 if weight_dtypes.transformer.torch_dtype() is None else weight_dtypes.transformer.torch_dtype(), - quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16) if weight_dtypes.transformer.is_gguf() else None, - ) - transformer = self._convert_diffusers_sub_module_to_dtype( - transformer, weight_dtypes.transformer, weight_dtypes.train_dtype, quantization, - ) - else: - transformer = self._load_diffusers_sub_module( - QwenImageTransformer2DModel, - weight_dtypes.transformer, - weight_dtypes.train_dtype, - base_model_name, - "transformer", - quantization, - ) + transformer = self._load_transformer( + QwenImageTransformer2DModel, + weight_dtypes, + base_model_name, + transformer_model_name, + quantization, + config=base_model_name, + ) model.model_type = model_type model.tokenizer = tokenizer diff --git a/modules/modelLoader/sana/SanaModelLoader.py b/modules/modelLoader/sana/SanaModelLoader.py index a904e3996..ec4d31239 100644 --- a/modules/modelLoader/sana/SanaModelLoader.py +++ b/modules/modelLoader/sana/SanaModelLoader.py @@ -52,7 +52,7 @@ def __load_diffusers( subfolder="scheduler", ) - text_encoder = self._load_transformers_sub_module( + text_encoder = self._load_text_encoder( Gemma2Model, weight_dtypes.text_encoder, weight_dtypes.fallback_train_dtype, @@ -60,21 +60,13 @@ def __load_diffusers( "text_encoder", ) - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderDC, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderDC, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + vae = self._load_vae( + AutoencoderDC, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) transformer = self._load_diffusers_sub_module( SanaTransformer2DModel, diff --git a/modules/modelLoader/stableDiffusion/StableDiffusionModelLoader.py b/modules/modelLoader/stableDiffusion/StableDiffusionModelLoader.py index aa610f485..223d1da61 100644 --- a/modules/modelLoader/stableDiffusion/StableDiffusionModelLoader.py +++ b/modules/modelLoader/stableDiffusion/StableDiffusionModelLoader.py @@ -87,7 +87,7 @@ def __load_diffusers( original_noise_scheduler=noise_scheduler, ) - text_encoder = self._load_transformers_sub_module( + text_encoder = self._load_text_encoder( CLIPTextModel, weight_dtypes.text_encoder, weight_dtypes.train_dtype, @@ -95,21 +95,13 @@ def __load_diffusers( "text_encoder", ) - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + vae = self._load_vae( + AutoencoderKL, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) unet = self._load_diffusers_sub_module( UNet2DConditionModel, diff --git a/modules/modelLoader/stableDiffusion3/StableDiffusion3ModelLoader.py b/modules/modelLoader/stableDiffusion3/StableDiffusion3ModelLoader.py index 47d87da74..0c4f4348a 100644 --- a/modules/modelLoader/stableDiffusion3/StableDiffusion3ModelLoader.py +++ b/modules/modelLoader/stableDiffusion3/StableDiffusion3ModelLoader.py @@ -81,7 +81,7 @@ def __load_diffusers( ) if include_text_encoder_1: - text_encoder_1 = self._load_transformers_sub_module( + text_encoder_1 = self._load_text_encoder( CLIPTextModelWithProjection, weight_dtypes.text_encoder, weight_dtypes.train_dtype, @@ -92,7 +92,7 @@ def __load_diffusers( text_encoder_1 = None if include_text_encoder_2: - text_encoder_2 = self._load_transformers_sub_module( + text_encoder_2 = self._load_text_encoder( CLIPTextModelWithProjection, weight_dtypes.text_encoder_2, weight_dtypes.train_dtype, @@ -103,7 +103,7 @@ def __load_diffusers( text_encoder_2 = None if include_text_encoder_3: - text_encoder_3 = self._load_transformers_sub_module( + text_encoder_3 = self._load_text_encoder( T5EncoderModel, weight_dtypes.text_encoder_3, weight_dtypes.fallback_train_dtype, @@ -113,21 +113,13 @@ def __load_diffusers( else: text_encoder_3 = None - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.train_dtype, - base_model_name, - "vae", - ) + vae = self._load_vae( + AutoencoderKL, + weight_dtypes.vae, + weight_dtypes.train_dtype, + base_model_name, + vae_model_name, + ) transformer = self._load_diffusers_sub_module( SD3Transformer2DModel, diff --git a/modules/modelLoader/stableDiffusionXL/StableDiffusionXLModelLoader.py b/modules/modelLoader/stableDiffusionXL/StableDiffusionXLModelLoader.py index afbab6581..ba340006a 100644 --- a/modules/modelLoader/stableDiffusionXL/StableDiffusionXLModelLoader.py +++ b/modules/modelLoader/stableDiffusionXL/StableDiffusionXLModelLoader.py @@ -83,7 +83,7 @@ def __load_diffusers( original_noise_scheduler=noise_scheduler, ) - text_encoder_1 = self._load_transformers_sub_module( + text_encoder_1 = self._load_text_encoder( CLIPTextModel, weight_dtypes.text_encoder, weight_dtypes.train_dtype, @@ -91,7 +91,7 @@ def __load_diffusers( "text_encoder", ) - text_encoder_2 = self._load_transformers_sub_module( + text_encoder_2 = self._load_text_encoder( CLIPTextModelWithProjection, weight_dtypes.text_encoder_2, weight_dtypes.train_dtype, @@ -99,21 +99,13 @@ def __load_diffusers( "text_encoder_2", ) - if vae_model_name: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.fallback_train_dtype, - vae_model_name, - ) - else: - vae = self._load_diffusers_sub_module( - AutoencoderKL, - weight_dtypes.vae, - weight_dtypes.fallback_train_dtype, - base_model_name, - "vae", - ) + vae = self._load_vae( + AutoencoderKL, + weight_dtypes.vae, + weight_dtypes.fallback_train_dtype, + base_model_name, + vae_model_name, + ) unet = self._load_diffusers_sub_module( UNet2DConditionModel, diff --git a/modules/modelLoader/wuerstchen/WuerstchenModelLoader.py b/modules/modelLoader/wuerstchen/WuerstchenModelLoader.py index 188107a2c..d8e86a19b 100644 --- a/modules/modelLoader/wuerstchen/WuerstchenModelLoader.py +++ b/modules/modelLoader/wuerstchen/WuerstchenModelLoader.py @@ -75,7 +75,7 @@ def __load_diffusers( ) if model_type.is_wuerstchen_v2(): - decoder_text_encoder = self._load_transformers_sub_module( + decoder_text_encoder = self._load_text_encoder( CLIPTextModel, weight_dtypes.decoder_text_encoder, weight_dtypes.train_dtype, @@ -164,7 +164,7 @@ def __load_diffusers( ) if model_type.is_wuerstchen_v2(): - prior_text_encoder = self._load_transformers_sub_module( + prior_text_encoder = self._load_text_encoder( CLIPTextModel, weight_dtypes.text_encoder, weight_dtypes.train_dtype, @@ -172,7 +172,7 @@ def __load_diffusers( "text_encoder", ) elif model_type.is_stable_cascade(): - prior_text_encoder = self._load_transformers_sub_module( + prior_text_encoder = self._load_text_encoder( CLIPTextModelWithProjection, weight_dtypes.text_encoder, weight_dtypes.train_dtype,