Avoid circular dep. issue (maybe better way to do this)
This commit is contained in:
parent
15ecaa7370
commit
91b0c466fb
@ -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)
|
||||
|
||||
@ -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))
|
||||
|
||||
@ -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:
|
||||
|
||||
Loading…
Reference in New Issue
Block a user