From 807b243e56754bd89e7d499405c8c40bc9fa574f Mon Sep 17 00:00:00 2001 From: filipstrand Date: Wed, 22 Jan 2025 08:18:39 +0100 Subject: [PATCH] Support variable numbers of transformer blocks --- src/mflux/controlnet/flux_controlnet.py | 6 ++++-- src/mflux/flux/flux.py | 6 ++++-- src/mflux/models/transformer/transformer.py | 4 ++-- src/mflux/weights/weight_handler.py | 3 +++ 4 files changed, 13 insertions(+), 6 deletions(-) diff --git a/src/mflux/controlnet/flux_controlnet.py b/src/mflux/controlnet/flux_controlnet.py index 0ff1a47..b961514 100644 --- a/src/mflux/controlnet/flux_controlnet.py +++ b/src/mflux/controlnet/flux_controlnet.py @@ -56,14 +56,16 @@ class Flux1Controlnet: self.t5_tokenizer = TokenizerT5(tokenizers.t5, max_length=self.model_config.max_sequence_length) self.clip_tokenizer = TokenizerCLIP(tokenizers.clip) + # Load the weights + weights = WeightHandler.load_regular_weights(repo_id=model_config.model_name, local_path=local_path) + # Initialize the models self.vae = VAE() - self.transformer = Transformer(model_config) + self.transformer = Transformer(model_config, num_transformer_blocks=weights.num_transformer_blocks()) self.t5_text_encoder = T5Encoder() self.clip_text_encoder = CLIPEncoder() # Set the weights and quantize the model - weights = WeightHandler.load_regular_weights(repo_id=model_config.model_name, local_path=local_path) self.bits = WeightUtil.set_weights_and_quantize( quantize_arg=quantize, weights=weights, diff --git a/src/mflux/flux/flux.py b/src/mflux/flux/flux.py index f01b541..fc6b8d6 100644 --- a/src/mflux/flux/flux.py +++ b/src/mflux/flux/flux.py @@ -46,14 +46,16 @@ class Flux1(nn.Module): self.t5_tokenizer = TokenizerT5(tokenizers.t5, max_length=self.model_config.max_sequence_length) self.clip_tokenizer = TokenizerCLIP(tokenizers.clip) + # Load the weights + weights = WeightHandler.load_regular_weights(repo_id=model_config.model_name, local_path=local_path) + # Initialize the models self.vae = VAE() - self.transformer = Transformer(model_config) + self.transformer = Transformer(model_config, num_transformer_blocks=weights.num_transformer_blocks()) self.t5_text_encoder = T5Encoder() self.clip_text_encoder = CLIPEncoder() # Set the weights and quantize the model - weights = WeightHandler.load_regular_weights(repo_id=model_config.model_name, local_path=local_path) self.bits = WeightUtil.set_weights_and_quantize( quantize_arg=quantize, weights=weights, diff --git a/src/mflux/models/transformer/transformer.py b/src/mflux/models/transformer/transformer.py index dbabac8..d0e5f23 100644 --- a/src/mflux/models/transformer/transformer.py +++ b/src/mflux/models/transformer/transformer.py @@ -19,13 +19,13 @@ from mflux.models.transformer.time_text_embed import TimeTextEmbed class Transformer(nn.Module): - def __init__(self, model_config: ModelConfig): + def __init__(self, model_config: ModelConfig, num_transformer_blocks: int): super().__init__() self.pos_embed = EmbedND() self.x_embedder = nn.Linear(64, 3072) self.time_text_embed = TimeTextEmbed(model_config=model_config) self.context_embedder = nn.Linear(4096, 3072) - self.transformer_blocks = [JointTransformerBlock(i) for i in range(19)] + self.transformer_blocks = [JointTransformerBlock(i) for i in range(num_transformer_blocks)] self.single_transformer_blocks = [SingleTransformerBlock(i) for i in range(38)] self.norm_out = AdaLayerNormContinuous(3072, 3072) self.proj_out = nn.Linear(3072, 64) diff --git a/src/mflux/weights/weight_handler.py b/src/mflux/weights/weight_handler.py index 7fb0470..74db4ce 100644 --- a/src/mflux/weights/weight_handler.py +++ b/src/mflux/weights/weight_handler.py @@ -57,6 +57,9 @@ class WeightHandler: ), ) + def num_transformer_blocks(self) -> int: + return len(self.transformer["transformer_blocks"]) + @staticmethod def _load_clip_encoder(root_path: Path) -> (dict, int): weights, quantization_level, _ = WeightHandler._get_weights("text_encoder", root_path)