diff --git a/src/mflux/controlnet/controlnet_util.py b/src/mflux/controlnet/controlnet_util.py new file mode 100644 index 0000000..ab6a9c5 --- /dev/null +++ b/src/mflux/controlnet/controlnet_util.py @@ -0,0 +1,14 @@ +import cv2 +import numpy as np +import PIL + + +class ControlnetUtil: + + @staticmethod + def preprocess_canny(img: PIL.Image) -> PIL.Image: + image_to_canny = np.array(img) + image_to_canny = cv2.Canny(image_to_canny, 100, 200) + image_to_canny = np.array(image_to_canny[:, :, None]) + image_to_canny = np.concatenate([image_to_canny, image_to_canny, image_to_canny], axis=2) + return PIL.Image.fromarray(image_to_canny) diff --git a/src/mflux/controlnet/flux_controlnet.py b/src/mflux/controlnet/flux_controlnet.py index 36e92bf..6d36e4a 100644 --- a/src/mflux/controlnet/flux_controlnet.py +++ b/src/mflux/controlnet/flux_controlnet.py @@ -11,7 +11,7 @@ from tqdm import tqdm 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 +from mflux.controlnet.controlnet_util import ControlnetUtil 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 @@ -114,7 +114,7 @@ class Flux1Controlnet: shape=[1, (config.height // 16) * (config.width // 16), 64], key=mx.random.key(seed) ) - control_image = preprocess_canny(control_image) + control_image = ControlnetUtil.preprocess_canny(control_image) controlnet_cond = ImageUtil.to_array(control_image) controlnet_cong = self.vae.encode(controlnet_cond) # the rescaling in the next line is not in the huggingface code, but without it the images from diff --git a/src/mflux/controlnet/utils_controlnet.py b/src/mflux/controlnet/utils_controlnet.py deleted file mode 100644 index 5de6b06..0000000 --- a/src/mflux/controlnet/utils_controlnet.py +++ /dev/null @@ -1,10 +0,0 @@ -import cv2 -import numpy as np -import PIL - -def preprocess_canny(img: PIL.Image) -> PIL.Image: - image_to_canny = np.array(img) - image_to_canny = cv2.Canny(image_to_canny, 100, 200) - image_to_canny = np.array(image_to_canny[:, :, None]) - image_to_canny = np.concatenate([image_to_canny, image_to_canny, image_to_canny], axis=2) - return PIL.Image.fromarray(image_to_canny)