Create util class for controlnet
This commit is contained in:
parent
78bbf3d8fb
commit
ad6f4ca8af
14
src/mflux/controlnet/controlnet_util.py
Normal file
14
src/mflux/controlnet/controlnet_util.py
Normal 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)
|
||||||
@ -11,7 +11,7 @@ from tqdm import tqdm
|
|||||||
from mflux.config.config import 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.controlnet_util import ControlnetUtil
|
||||||
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.embed_nd import EmbedND
|
||||||
@ -114,7 +114,7 @@ class Flux1Controlnet:
|
|||||||
shape=[1, (config.height // 16) * (config.width // 16), 64],
|
shape=[1, (config.height // 16) * (config.width // 16), 64],
|
||||||
key=mx.random.key(seed)
|
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_cond = ImageUtil.to_array(control_image)
|
||||||
controlnet_cong = self.vae.encode(controlnet_cond)
|
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
|
# the rescaling in the next line is not in the huggingface code, but without it the images from
|
||||||
|
|||||||
@ -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)
|
|
||||||
Loading…
Reference in New Issue
Block a user