Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
54 changes: 16 additions & 38 deletions modules/modelLoader/AnimaModelLoader.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,6 @@
AutoencoderKLQwenImage,
CosmosTransformer3DModel,
FlowMatchEulerDiscreteScheduler,
GGUFQuantizationConfig,
)
from transformers import Qwen2Tokenizer, Qwen3Model, T5TokenizerFast

Expand Down Expand Up @@ -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,
Expand All @@ -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
Expand Down
55 changes: 16 additions & 39 deletions modules/modelLoader/ErnieModelLoader.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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,
Expand All @@ -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
Expand Down
58 changes: 17 additions & 41 deletions modules/modelLoader/Flux2ModelLoader.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -60,35 +57,22 @@ 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(
base_model_name,
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,
Expand All @@ -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,
Expand All @@ -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
Expand Down
14 changes: 8 additions & 6 deletions modules/modelLoader/IdeogramModelLoader.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")

Expand All @@ -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,
):
Expand All @@ -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,
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
Loading