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 cv2
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import PIL
|
import PIL.Image
|
||||||
|
|
||||||
from mflux import ImageUtil
|
|
||||||
|
|
||||||
log = logging.getLogger(__name__)
|
log = logging.getLogger(__name__)
|
||||||
|
|
||||||
@ -27,7 +26,9 @@ class ControlnetUtil:
|
|||||||
return img
|
return img
|
||||||
|
|
||||||
@staticmethod
|
@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)
|
base, ext = os.path.splitext(path)
|
||||||
new_filename = f"{base}_controlnet_canny{ext}"
|
new_filename = f"{base}_controlnet_canny{ext}"
|
||||||
ImageUtil.save_image(control_image, new_filename)
|
ImageUtil.save_image(control_image, new_filename)
|
||||||
|
|||||||
@ -1,4 +1,5 @@
|
|||||||
import logging
|
import logging
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
from mlx import nn
|
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.text_encoder.t5_encoder.t5_encoder import T5Encoder
|
||||||
from mflux.models.transformer.transformer import Transformer
|
from mflux.models.transformer.transformer import Transformer
|
||||||
from mflux.models.vae.vae import VAE
|
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.post_processing.image_util import ImageUtil
|
||||||
from mflux.tokenizer.clip_tokenizer import TokenizerCLIP
|
from mflux.tokenizer.clip_tokenizer import TokenizerCLIP
|
||||||
from mflux.tokenizer.t5_tokenizer import TokenizerT5
|
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.model_saver import ModelSaver
|
||||||
from mflux.weights.weight_handler import WeightHandler
|
from mflux.weights.weight_handler import WeightHandler
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from mflux.post_processing.generated_image import GeneratedImage
|
||||||
|
|
||||||
|
|
||||||
log = logging.getLogger(__name__)
|
log = logging.getLogger(__name__)
|
||||||
|
|
||||||
CONTROLNET_ID = "InstantX/FLUX.1-dev-Controlnet-Canny"
|
CONTROLNET_ID = "InstantX/FLUX.1-dev-Controlnet-Canny"
|
||||||
@ -107,7 +111,7 @@ class Flux1Controlnet:
|
|||||||
control_image_path: str,
|
control_image_path: str,
|
||||||
save_control_image_canny: bool = False,
|
save_control_image_canny: bool = False,
|
||||||
config: ConfigControlnet = ConfigControlnet()
|
config: ConfigControlnet = ConfigControlnet()
|
||||||
) -> GeneratedImage: # fmt: off
|
) -> "GeneratedImage": # fmt: off
|
||||||
# Create a new runtime config based on the model type and input parameters
|
# Create a new runtime config based on the model type and input parameters
|
||||||
config = RuntimeConfig(config, self.model_config)
|
config = RuntimeConfig(config, self.model_config)
|
||||||
time_steps = tqdm(range(config.num_inference_steps))
|
time_steps = tqdm(range(config.num_inference_steps))
|
||||||
|
|||||||
@ -3,7 +3,6 @@ import importlib
|
|||||||
import PIL.Image
|
import PIL.Image
|
||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
|
|
||||||
from mflux import ImageUtil
|
|
||||||
from mflux.config.model_config import ModelConfig
|
from mflux.config.model_config import ModelConfig
|
||||||
|
|
||||||
|
|
||||||
@ -39,6 +38,8 @@ class GeneratedImage:
|
|||||||
self.controlnet_strength = controlnet_strength
|
self.controlnet_strength = controlnet_strength
|
||||||
|
|
||||||
def save(self, path: str, export_json_metadata: bool = False) -> None:
|
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)
|
ImageUtil.save_image(self.image, path, self._get_metadata(), export_json_metadata)
|
||||||
|
|
||||||
def _get_metadata(self) -> dict:
|
def _get_metadata(self) -> dict:
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user