fixes from review
This commit is contained in:
parent
a04d139ff5
commit
dee452b3bb
@ -25,6 +25,7 @@ dependencies = [
|
|||||||
"huggingface-hub>=0.24.5",
|
"huggingface-hub>=0.24.5",
|
||||||
"safetensors>=0.4.4",
|
"safetensors>=0.4.4",
|
||||||
"piexif>=1.1.3",
|
"piexif>=1.1.3",
|
||||||
|
"opencv-python>=4.10.0",
|
||||||
]
|
]
|
||||||
|
|
||||||
[project.urls]
|
[project.urls]
|
||||||
|
|||||||
@ -8,3 +8,4 @@ tqdm>=4.66.5
|
|||||||
huggingface-hub>=0.24.5
|
huggingface-hub>=0.24.5
|
||||||
safetensors>=0.4.4
|
safetensors>=0.4.4
|
||||||
piexif>=1.1.3
|
piexif>=1.1.3
|
||||||
|
opencv-python>=4.10.0
|
||||||
@ -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 ValueError("Controlnet conditioning scale is only available for ConfigControlnet")
|
return 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:
|
||||||
|
|||||||
@ -7,7 +7,7 @@ from mlx import nn
|
|||||||
from mlx.utils import tree_flatten
|
from mlx.utils import tree_flatten
|
||||||
from tqdm import tqdm
|
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.model_config import ModelConfig
|
||||||
from mflux.config.runtime_config import RuntimeConfig
|
from mflux.config.runtime_config import RuntimeConfig
|
||||||
from mflux.controlnet.utils_controlnet import preprocess_canny
|
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.model_config import ModelConfig
|
||||||
from mflux.config.runtime_config import RuntimeConfig
|
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.embed_nd import EmbedND
|
||||||
from mflux.models.transformer.joint_transformer_block import JointTransformerBlock
|
from mflux.models.transformer.joint_transformer_block import JointTransformerBlock
|
||||||
from mflux.models.transformer.single_transformer_block import SingleTransformerBlock
|
from mflux.models.transformer.single_transformer_block import SingleTransformerBlock
|
||||||
from mflux.models.transformer.time_text_embed import TimeTextEmbed
|
from mflux.models.transformer.time_text_embed import TimeTextEmbed
|
||||||
import numpy as np
|
|
||||||
import cv2
|
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
|
||||||
log = logging.getLogger(__name__)
|
log = logging.getLogger(__name__)
|
||||||
|
|
||||||
CONTOLNET_ID = "InstantX/FLUX.1-dev-Controlnet-Canny"
|
CONTROLNET_ID = "InstantX/FLUX.1-dev-Controlnet-Canny"
|
||||||
|
|
||||||
class Flux1Controlnet:
|
class Flux1Controlnet:
|
||||||
def __init__(
|
def __init__(
|
||||||
@ -88,7 +85,7 @@ class Flux1Controlnet:
|
|||||||
if weights.quantization_level is not None:
|
if weights.quantization_level is not None:
|
||||||
self._set_model_weights(weights)
|
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(
|
self.transformer_controlnet = TransformerControlnet(
|
||||||
model_config=model_config,
|
model_config=model_config,
|
||||||
num_blocks= controlnet_config["num_layers"],
|
num_blocks= controlnet_config["num_layers"],
|
||||||
@ -135,7 +132,7 @@ class Flux1Controlnet:
|
|||||||
pooled_prompt_embeds = self.clip_text_encoder.forward(clip_tokens)
|
pooled_prompt_embeds = self.clip_text_encoder.forward(clip_tokens)
|
||||||
|
|
||||||
for t in time_steps:
|
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,
|
t=t,
|
||||||
prompt_embeds=prompt_embeds,
|
prompt_embeds=prompt_embeds,
|
||||||
pooled_prompt_embeds=pooled_prompt_embeds,
|
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||||
@ -150,8 +147,8 @@ class Flux1Controlnet:
|
|||||||
pooled_prompt_embeds=pooled_prompt_embeds,
|
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||||
hidden_states=latents,
|
hidden_states=latents,
|
||||||
config=config,
|
config=config,
|
||||||
controlnet_block_samples=controlnet_block_samples,
|
controlnet_block_samples=ctrlnet_block_samples,
|
||||||
# controlnet_single_block_samples=controlnet_single_block_samples,
|
controlnet_single_block_samples=ctrlnet_single_block_samples,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 4.t Take one denoise step
|
# 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_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]
|
controlnet_single_block_samples = [sample * conditioning_scale for sample in controlnet_single_block_samples]
|
||||||
|
|
||||||
return controlnet_block_samples
|
return controlnet_block_samples, controlnet_single_block_samples
|
||||||
@ -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('--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('--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('--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)')
|
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:
|
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 = Flux1Controlnet(
|
flux = Flux1Controlnet(
|
||||||
model_config=ModelConfig.from_alias(args.model),
|
model_config=ModelConfig.from_alias(args.model),
|
||||||
|
|||||||
@ -54,7 +54,7 @@ class Transformer(nn.Module):
|
|||||||
text_embeddings=text_embeddings,
|
text_embeddings=text_embeddings,
|
||||||
rotary_embeddings=image_rotary_emb
|
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 = len(self.transformer_blocks) / len(controlnet_block_samples)
|
||||||
interval_control = int(math.ceil(interval_control))
|
interval_control = int(math.ceil(interval_control))
|
||||||
hidden_states = hidden_states + controlnet_block_samples[idx // interval_control]
|
hidden_states = hidden_states + controlnet_block_samples[idx // interval_control]
|
||||||
@ -67,7 +67,7 @@ class Transformer(nn.Module):
|
|||||||
text_embeddings=text_embeddings,
|
text_embeddings=text_embeddings,
|
||||||
rotary_embeddings=image_rotary_emb
|
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 = len(self.single_transformer_blocks) / len(controlnet_single_block_samples)
|
||||||
interval_control = int(math.ceil(interval_control))
|
interval_control = int(math.ceil(interval_control))
|
||||||
hidden_states[:, encoder_hidden_states.shape[1] :, ...] = (
|
hidden_states[:, encoder_hidden_states.shape[1] :, ...] = (
|
||||||
|
|||||||
@ -90,7 +90,7 @@ class WeightHandler:
|
|||||||
return weights, quantization_level
|
return weights, quantization_level
|
||||||
|
|
||||||
@staticmethod
|
@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"]))
|
controlnet_path = Path(snapshot_download(repo_id=controlnet_id,allow_patterns=["*.safetensors","config.json"]))
|
||||||
file = next(controlnet_path.glob("diffusion_pytorch_model.safetensors"))
|
file = next(controlnet_path.glob("diffusion_pytorch_model.safetensors"))
|
||||||
quantization_level = mx.load(str(file), return_metadata=True)[1].get("quantization_level")
|
quantization_level = mx.load(str(file), return_metadata=True)[1].get("quantization_level")
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user