Small additions and refactoring
This commit is contained in:
parent
a2900c3b13
commit
25e375d45d
@ -41,7 +41,7 @@ class RuntimeConfig:
|
|||||||
if isinstance(self.config, ConfigControlnet):
|
if isinstance(self.config, ConfigControlnet):
|
||||||
return self.config.controlnet_strength
|
return self.config.controlnet_strength
|
||||||
else:
|
else:
|
||||||
return NotImplementedError("Controlnet conditioning scale is only available for ConfigControlnet")
|
raise NotImplementedError("Controlnet conditioning scale is only available for ConfigControlnet")
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _create_sigmas(config, model) -> mx.array:
|
def _create_sigmas(config, model) -> mx.array:
|
||||||
|
|||||||
@ -1,3 +1,4 @@
|
|||||||
|
import logging
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Tuple
|
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.controlnet.utils_controlnet import preprocess_canny
|
||||||
from mflux.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder
|
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.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.transformer.transformer import Transformer
|
||||||
from mflux.models.vae.vae import VAE
|
from mflux.models.vae.vae import VAE
|
||||||
from mflux.post_processing.image import GeneratedImage
|
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.tokenizer.tokenizer_handler import TokenizerHandler
|
||||||
from mflux.weights.weight_handler import WeightHandler
|
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__)
|
log = logging.getLogger(__name__)
|
||||||
|
|
||||||
CONTROLNET_ID = "InstantX/FLUX.1-dev-Controlnet-Canny"
|
CONTROLNET_ID = "InstantX/FLUX.1-dev-Controlnet-Canny"
|
||||||
|
|
||||||
|
|
||||||
class Flux1Controlnet:
|
class Flux1Controlnet:
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@ -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('--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('--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('--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('--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('--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')
|
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:
|
if args.path and args.model is None:
|
||||||
parser.error("--model must be specified when using --path")
|
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
|
# Load the model
|
||||||
flux = Flux1(
|
flux = Flux1(
|
||||||
model_config=ModelConfig.from_alias(args.model),
|
model_config=ModelConfig.from_alias(args.model),
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user