diff --git a/src/mflux/config/runtime_config.py b/src/mflux/config/runtime_config.py index 20205bc..efb33cf 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 NotImplementedError("Controlnet conditioning scale is only available for ConfigControlnet") + raise 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 2ed4c70..41174ba 100644 --- a/src/mflux/controlnet/flux_controlnet.py +++ b/src/mflux/controlnet/flux_controlnet.py @@ -1,3 +1,4 @@ +import logging from pathlib import Path from typing import Tuple @@ -13,6 +14,10 @@ from mflux.config.runtime_config import RuntimeConfig from mflux.controlnet.utils_controlnet import preprocess_canny from mflux.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder from mflux.models.text_encoder.t5_encoder.t5_encoder import T5Encoder +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 from mflux.models.transformer.transformer import Transformer from mflux.models.vae.vae import VAE from mflux.post_processing.image import GeneratedImage @@ -22,19 +27,11 @@ from mflux.tokenizer.t5_tokenizer import TokenizerT5 from mflux.tokenizer.tokenizer_handler import TokenizerHandler 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.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 logging - log = logging.getLogger(__name__) CONTROLNET_ID = "InstantX/FLUX.1-dev-Controlnet-Canny" + class Flux1Controlnet: def __init__( self, diff --git a/src/mflux/generate.py b/src/mflux/generate.py index 1b8e59f..9bf281d 100644 --- a/src/mflux/generate.py +++ b/src/mflux/generate.py @@ -18,7 +18,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('--quantize', "-q", type=int, choices=[4, 8], default=None, help='Quantize the model (4 or 8, Default is None)') parser.add_argument('--path', type=str, default=None, help='Local path for loading a model from disk') @@ -31,6 +31,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 = Flux1( model_config=ModelConfig.from_alias(args.model),