From 78bbf3d8fb6affc3adc1f2a8649ba6008e32ac29 Mon Sep 17 00:00:00 2001 From: filipstrand Date: Tue, 17 Sep 2024 06:41:05 +0200 Subject: [PATCH] Fix formatting and some warnings --- src/mflux/controlnet/flux_controlnet.py | 14 +++++++------- src/mflux/models/transformer/transformer.py | 13 ++++++------- 2 files changed, 13 insertions(+), 14 deletions(-) diff --git a/src/mflux/controlnet/flux_controlnet.py b/src/mflux/controlnet/flux_controlnet.py index 41174ba..36e92bf 100644 --- a/src/mflux/controlnet/flux_controlnet.py +++ b/src/mflux/controlnet/flux_controlnet.py @@ -85,8 +85,8 @@ class Flux1Controlnet: weights_controlnet, ctrlnet_quantization_level, controlnet_config = WeightHandler.load_controlnet_transformer(controlnet_id=CONTROLNET_ID) self.transformer_controlnet = TransformerControlnet( model_config=model_config, - num_blocks= controlnet_config["num_layers"], - num_single_blocks= controlnet_config["num_single_layers"], + num_blocks=controlnet_config["num_layers"], + num_single_blocks=controlnet_config["num_single_layers"], ) if ctrlnet_quantization_level is None: @@ -230,6 +230,7 @@ class Flux1Controlnet: ControlNetOutput = Tuple[list[mx.array], list[mx.array]] + class TransformerControlnet(nn.Module): def __init__( @@ -237,7 +238,7 @@ class TransformerControlnet(nn.Module): model_config: ModelConfig, num_blocks: int, num_single_blocks: int, - ): + ): super().__init__() self.pos_embed = EmbedND() self.x_embedder = nn.Linear(64, 3072) @@ -270,8 +271,8 @@ class TransformerControlnet(nn.Module): guidance = mx.broadcast_to(config.guidance * config.num_train_steps, (1,)).astype(config.precision) text_embeddings = self.time_text_embed.forward(time_step, pooled_prompt_embeds, guidance) encoder_hidden_states = self.context_embedder(prompt_embeds) - txt_ids = Transformer._prepare_text_ids(seq_len=prompt_embeds.shape[1]) - img_ids = Transformer._prepare_latent_image_ids(config.height, config.width) + txt_ids = Transformer.prepare_text_ids(seq_len=prompt_embeds.shape[1]) + img_ids = Transformer.prepare_latent_image_ids(config.height, config.width) ids = mx.concatenate((txt_ids, img_ids), axis=1) image_rotary_emb = self.pos_embed.forward(ids) @@ -293,7 +294,6 @@ class TransformerControlnet(nn.Module): block_sample = controlnet_block(block_sample) controlnet_block_samples = controlnet_block_samples + (block_sample,) - single_block_samples = () for block in self.single_transformer_blocks: ctrlnet_hidden_states = block.forward( @@ -312,4 +312,4 @@ class TransformerControlnet(nn.Module): controlnet_block_samples = [sample * conditioning_scale for sample in controlnet_block_samples] controlnet_single_block_samples = [sample * conditioning_scale for sample in controlnet_single_block_samples] - return controlnet_block_samples, controlnet_single_block_samples \ No newline at end of file + return controlnet_block_samples, controlnet_single_block_samples diff --git a/src/mflux/models/transformer/transformer.py b/src/mflux/models/transformer/transformer.py index e57d514..6bbcc53 100644 --- a/src/mflux/models/transformer/transformer.py +++ b/src/mflux/models/transformer/transformer.py @@ -12,7 +12,6 @@ from mflux.models.transformer.single_transformer_block import SingleTransformerB from mflux.models.transformer.time_text_embed import TimeTextEmbed - class Transformer(nn.Module): def __init__(self, model_config: ModelConfig): @@ -33,8 +32,8 @@ class Transformer(nn.Module): pooled_prompt_embeds: mx.array, hidden_states: mx.array, config: RuntimeConfig, - controlnet_block_samples: Tuple[mx.array] | None = None, - controlnet_single_block_samples: Tuple[mx.array] | None = None, + controlnet_block_samples: list[mx.array] | None = None, + controlnet_single_block_samples: list[mx.array] | None = None, ) -> mx.array: time_step = config.sigmas[t] * config.num_train_steps time_step = mx.broadcast_to(time_step, (1,)).astype(config.precision) @@ -42,8 +41,8 @@ class Transformer(nn.Module): guidance = mx.broadcast_to(config.guidance * config.num_train_steps, (1,)).astype(config.precision) text_embeddings = self.time_text_embed.forward(time_step, pooled_prompt_embeds, guidance) encoder_hidden_states = self.context_embedder(prompt_embeds) - txt_ids = Transformer._prepare_text_ids(seq_len=prompt_embeds.shape[1]) - img_ids = Transformer._prepare_latent_image_ids(config.height, config.width) + txt_ids = Transformer.prepare_text_ids(seq_len=prompt_embeds.shape[1]) + img_ids = Transformer.prepare_latent_image_ids(config.height, config.width) ids = mx.concatenate((txt_ids, img_ids), axis=1) image_rotary_emb = self.pos_embed.forward(ids) @@ -82,7 +81,7 @@ class Transformer(nn.Module): return noise @staticmethod - def _prepare_latent_image_ids(height: int, width: int) -> mx.array: + def prepare_latent_image_ids(height: int, width: int) -> mx.array: latent_width = width // 16 latent_height = height // 16 latent_image_ids = mx.zeros((latent_height, latent_width, 3)) @@ -93,5 +92,5 @@ class Transformer(nn.Module): return latent_image_ids @staticmethod - def _prepare_text_ids(seq_len: mx.array) -> mx.array: + def prepare_text_ids(seq_len: mx.array) -> mx.array: return mx.zeros((1, seq_len, 3))