diff --git a/pyproject.toml b/pyproject.toml index 60b6838..17bba62 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -25,6 +25,7 @@ dependencies = [ "huggingface-hub>=0.24.5", "safetensors>=0.4.4", "piexif>=1.1.3", + "opencv-python>=4.10.0", ] [project.urls] diff --git a/requirements.txt b/requirements.txt index b9a4c8f..f6d689e 100644 --- a/requirements.txt +++ b/requirements.txt @@ -7,4 +7,5 @@ torch>=2.3.1 tqdm>=4.66.5 huggingface-hub>=0.24.5 safetensors>=0.4.4 -piexif>=1.1.3 \ No newline at end of file +piexif>=1.1.3 +opencv-python>=4.10.0 \ No newline at end of file diff --git a/src/mflux/config/runtime_config.py b/src/mflux/config/runtime_config.py index 96bd89b..67a357a 100644 --- a/src/mflux/config/runtime_config.py +++ b/src/mflux/config/runtime_config.py @@ -41,7 +41,7 @@ class RuntimeConfig: if isinstance(self.config, ConfigControlnet): return self.config.controlnet_strength else: - return ValueError("Controlnet conditioning scale is only available for ConfigControlnet") + return NotImplementedError("Controlnet conditioning scale is only available for ConfigControlnet") @staticmethod def _create_sigmas(config, model) -> mx.array: diff --git a/src/mflux/controlnet/flux_controlnet.py b/src/mflux/controlnet/flux_controlnet.py index 5a5a71b..a1aafac 100644 --- a/src/mflux/controlnet/flux_controlnet.py +++ b/src/mflux/controlnet/flux_controlnet.py @@ -7,7 +7,7 @@ from mlx import nn from mlx.utils import tree_flatten from tqdm import tqdm -from mflux.config.config import Config, ConfigControlnet +from mflux.config.config import ConfigControlnet from mflux.config.model_config import ModelConfig from mflux.config.runtime_config import RuntimeConfig from mflux.controlnet.utils_controlnet import preprocess_canny @@ -24,19 +24,16 @@ from mflux.weights.weight_handler import WeightHandler from mflux.config.model_config import ModelConfig from mflux.config.runtime_config import RuntimeConfig -from mflux.models.transformer.ada_layer_norm_continous import AdaLayerNormContinuous from mflux.models.transformer.embed_nd import EmbedND from mflux.models.transformer.joint_transformer_block import JointTransformerBlock from mflux.models.transformer.single_transformer_block import SingleTransformerBlock from mflux.models.transformer.time_text_embed import TimeTextEmbed -import numpy as np -import cv2 import logging log = logging.getLogger(__name__) -CONTOLNET_ID = "InstantX/FLUX.1-dev-Controlnet-Canny" +CONTROLNET_ID = "InstantX/FLUX.1-dev-Controlnet-Canny" class Flux1Controlnet: def __init__( @@ -88,7 +85,7 @@ class Flux1Controlnet: if weights.quantization_level is not None: self._set_model_weights(weights) - weights_controlnet, ctrlnet_quantization_level, controlnet_config = WeightHandler.load_controlnet_transformer(controlnet_id=CONTOLNET_ID) + 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"], @@ -135,7 +132,7 @@ class Flux1Controlnet: pooled_prompt_embeds = self.clip_text_encoder.forward(clip_tokens) for t in time_steps: - controlnet_block_samples = self.transformer_controlnet.forward( + ctrlnet_block_samples, ctrlnet_single_block_samples = self.transformer_controlnet.forward( t=t, prompt_embeds=prompt_embeds, pooled_prompt_embeds=pooled_prompt_embeds, @@ -150,8 +147,8 @@ class Flux1Controlnet: pooled_prompt_embeds=pooled_prompt_embeds, hidden_states=latents, config=config, - controlnet_block_samples=controlnet_block_samples, - # controlnet_single_block_samples=controlnet_single_block_samples, + controlnet_block_samples=ctrlnet_block_samples, + controlnet_single_block_samples=ctrlnet_single_block_samples, ) # 4.t Take one denoise step @@ -318,4 +315,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 \ No newline at end of file + return controlnet_block_samples, controlnet_single_block_samples \ No newline at end of file diff --git a/src/mflux/generate_controlnet.py b/src/mflux/generate_controlnet.py index a8b3c57..4c5bb1a 100644 --- a/src/mflux/generate_controlnet.py +++ b/src/mflux/generate_controlnet.py @@ -20,7 +20,7 @@ def main(): parser.add_argument('--seed', type=int, default=None, help='Entropy Seed (Default is time-based random-seed)') parser.add_argument('--height', type=int, default=1024, help='Image height (Default is 1024)') parser.add_argument('--width', type=int, default=1024, help='Image width (Default is 1024)') - parser.add_argument('--steps', type=int, default=4, help='Inference Steps') + parser.add_argument('--steps', type=int, default=None, help='Inference Steps') parser.add_argument('--guidance', type=float, default=3.5, help='Guidance Scale (Default is 3.5)') parser.add_argument('--controlnet-strength', type=float, default=0.7, help='Controls how strongly the control image influences the output image. A value of 0.0 means no influence. (Default is 0.7)') parser.add_argument('--quantize', "-q", type=int, choices=[4, 8], default=None, help='Quantize the model (4 or 8, Default is None)') @@ -34,6 +34,9 @@ def main(): if args.path and args.model is None: parser.error("--model must be specified when using --path") + if args.steps is None: + args.steps = 4 if args.model == "schnell" else 14 + # Load the model flux = Flux1Controlnet( model_config=ModelConfig.from_alias(args.model), diff --git a/src/mflux/models/transformer/transformer.py b/src/mflux/models/transformer/transformer.py index 10af56e..e57d514 100644 --- a/src/mflux/models/transformer/transformer.py +++ b/src/mflux/models/transformer/transformer.py @@ -54,7 +54,7 @@ class Transformer(nn.Module): text_embeddings=text_embeddings, rotary_embeddings=image_rotary_emb ) - if controlnet_block_samples is not None: + if controlnet_block_samples is not None and len(controlnet_block_samples) > 0: interval_control = len(self.transformer_blocks) / len(controlnet_block_samples) interval_control = int(math.ceil(interval_control)) hidden_states = hidden_states + controlnet_block_samples[idx // interval_control] @@ -67,7 +67,7 @@ class Transformer(nn.Module): text_embeddings=text_embeddings, rotary_embeddings=image_rotary_emb ) - if controlnet_single_block_samples is not None: + if controlnet_single_block_samples is not None and len(controlnet_single_block_samples) > 0: interval_control = len(self.single_transformer_blocks) / len(controlnet_single_block_samples) interval_control = int(math.ceil(interval_control)) hidden_states[:, encoder_hidden_states.shape[1] :, ...] = ( diff --git a/src/mflux/weights/weight_handler.py b/src/mflux/weights/weight_handler.py index bc3a7fe..dbd25f1 100644 --- a/src/mflux/weights/weight_handler.py +++ b/src/mflux/weights/weight_handler.py @@ -90,7 +90,7 @@ class WeightHandler: return weights, quantization_level @staticmethod - def load_controlnet_transformer(controlnet_id: Path | None = None) -> (dict, int): + def load_controlnet_transformer(controlnet_id: str) -> (dict, int): controlnet_path = Path(snapshot_download(repo_id=controlnet_id,allow_patterns=["*.safetensors","config.json"])) file = next(controlnet_path.glob("diffusion_pytorch_model.safetensors")) quantization_level = mx.load(str(file), return_metadata=True)[1].get("quantization_level")