From 91b0c466fbf9a2ec6f8e06afbf2c279428e8821e Mon Sep 17 00:00:00 2001 From: filipstrand Date: Sun, 22 Sep 2024 21:07:28 +0200 Subject: [PATCH] Avoid circular dep. issue (maybe better way to do this) --- src/mflux/controlnet/controlnet_util.py | 7 ++++--- src/mflux/controlnet/flux_controlnet.py | 8 ++++++-- src/mflux/post_processing/generated_image.py | 3 ++- 3 files changed, 12 insertions(+), 6 deletions(-) diff --git a/src/mflux/controlnet/controlnet_util.py b/src/mflux/controlnet/controlnet_util.py index e7069d0..7f083cc 100644 --- a/src/mflux/controlnet/controlnet_util.py +++ b/src/mflux/controlnet/controlnet_util.py @@ -3,9 +3,8 @@ import os import cv2 import numpy as np -import PIL +import PIL.Image -from mflux import ImageUtil log = logging.getLogger(__name__) @@ -27,7 +26,9 @@ class ControlnetUtil: return img @staticmethod - def save_canny_image(control_image, path: str): + def save_canny_image(control_image: PIL.Image, path: str): + from mflux import ImageUtil + base, ext = os.path.splitext(path) new_filename = f"{base}_controlnet_canny{ext}" ImageUtil.save_image(control_image, new_filename) diff --git a/src/mflux/controlnet/flux_controlnet.py b/src/mflux/controlnet/flux_controlnet.py index 5b0d3e8..1483dd0 100644 --- a/src/mflux/controlnet/flux_controlnet.py +++ b/src/mflux/controlnet/flux_controlnet.py @@ -1,4 +1,5 @@ import logging +from typing import TYPE_CHECKING import mlx.core as mx from mlx import nn @@ -14,7 +15,6 @@ 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.transformer import Transformer from mflux.models.vae.vae import VAE -from mflux.post_processing.generated_image import GeneratedImage from mflux.post_processing.image_util import ImageUtil from mflux.tokenizer.clip_tokenizer import TokenizerCLIP from mflux.tokenizer.t5_tokenizer import TokenizerT5 @@ -22,6 +22,10 @@ from mflux.tokenizer.tokenizer_handler import TokenizerHandler from mflux.weights.model_saver import ModelSaver from mflux.weights.weight_handler import WeightHandler +if TYPE_CHECKING: + from mflux.post_processing.generated_image import GeneratedImage + + log = logging.getLogger(__name__) CONTROLNET_ID = "InstantX/FLUX.1-dev-Controlnet-Canny" @@ -107,7 +111,7 @@ class Flux1Controlnet: control_image_path: str, save_control_image_canny: bool = False, config: ConfigControlnet = ConfigControlnet() - ) -> GeneratedImage: # fmt: off + ) -> "GeneratedImage": # fmt: off # Create a new runtime config based on the model type and input parameters config = RuntimeConfig(config, self.model_config) time_steps = tqdm(range(config.num_inference_steps)) diff --git a/src/mflux/post_processing/generated_image.py b/src/mflux/post_processing/generated_image.py index 35019ef..add347c 100644 --- a/src/mflux/post_processing/generated_image.py +++ b/src/mflux/post_processing/generated_image.py @@ -3,7 +3,6 @@ import importlib import PIL.Image import mlx.core as mx -from mflux import ImageUtil from mflux.config.model_config import ModelConfig @@ -39,6 +38,8 @@ class GeneratedImage: self.controlnet_strength = controlnet_strength def save(self, path: str, export_json_metadata: bool = False) -> None: + from mflux import ImageUtil + ImageUtil.save_image(self.image, path, self._get_metadata(), export_json_metadata) def _get_metadata(self) -> dict: