Create util class for controlnet

This commit is contained in:
filipstrand 2024-09-17 06:43:13 +02:00
parent 78bbf3d8fb
commit ad6f4ca8af
3 changed files with 16 additions and 12 deletions

View File

@ -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)

View File

@ -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

View File

@ -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)