Avoid circular dep. issue (maybe better way to do this)

This commit is contained in:
filipstrand 2024-09-22 21:07:28 +02:00
parent 15ecaa7370
commit 91b0c466fb
3 changed files with 12 additions and 6 deletions

View File

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

View File

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

View File

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