Merge pull request #140 from anthonywu/restore-keyboard-interrupt
Fix: image gen should stop even if no callbacks are registered
This commit is contained in:
commit
a1325ccba3
@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
|||||||
|
|
||||||
[project]
|
[project]
|
||||||
name = "mflux"
|
name = "mflux"
|
||||||
version = "0.6.0"
|
version = "0.6.1"
|
||||||
description = "A MLX port of FLUX based on the Huggingface Diffusers implementation."
|
description = "A MLX port of FLUX based on the Huggingface Diffusers implementation."
|
||||||
readme = "README.md"
|
readme = "README.md"
|
||||||
keywords = ["diffusers", "flux", "mlx"]
|
keywords = ["diffusers", "flux", "mlx"]
|
||||||
|
|||||||
@ -4,7 +4,6 @@ import mlx.core as mx
|
|||||||
import PIL.Image
|
import PIL.Image
|
||||||
import tqdm
|
import tqdm
|
||||||
|
|
||||||
from mflux import StopImageGenerationException
|
|
||||||
from mflux.callbacks.callback import BeforeLoopCallback, InLoopCallback, InterruptCallback
|
from mflux.callbacks.callback import BeforeLoopCallback, InLoopCallback, InterruptCallback
|
||||||
from mflux.config.runtime_config import RuntimeConfig
|
from mflux.config.runtime_config import RuntimeConfig
|
||||||
from mflux.post_processing.array_util import ArrayUtil
|
from mflux.post_processing.array_util import ArrayUtil
|
||||||
@ -69,7 +68,6 @@ class StepwiseHandler(BeforeLoopCallback, InLoopCallback, InterruptCallback):
|
|||||||
time_steps: tqdm
|
time_steps: tqdm
|
||||||
) -> None: # fmt: off
|
) -> None: # fmt: off
|
||||||
self._save_composite(seed=seed)
|
self._save_composite(seed=seed)
|
||||||
raise StopImageGenerationException(f"Stopping image generation at step {t + 1}/{len(time_steps)}")
|
|
||||||
|
|
||||||
def _save_image(
|
def _save_image(
|
||||||
self,
|
self,
|
||||||
|
|||||||
@ -6,6 +6,7 @@ from mflux.callbacks.callbacks import Callbacks
|
|||||||
from mflux.config.config import Config
|
from mflux.config.config import Config
|
||||||
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.error.exceptions import StopImageGenerationException
|
||||||
from mflux.flux.flux_initializer import FluxInitializer
|
from mflux.flux.flux_initializer import FluxInitializer
|
||||||
from mflux.latent_creator.latent_creator import Img2Img, LatentCreator
|
from mflux.latent_creator.latent_creator import Img2Img, LatentCreator
|
||||||
from mflux.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder
|
from mflux.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder
|
||||||
@ -136,6 +137,7 @@ class Flux1InContextLoRA(nn.Module):
|
|||||||
config=config,
|
config=config,
|
||||||
time_steps=time_steps,
|
time_steps=time_steps,
|
||||||
)
|
)
|
||||||
|
raise StopImageGenerationException(f"Stopping image generation at step {t + 1}/{len(time_steps)}")
|
||||||
|
|
||||||
# (Optional) Call subscribers after loop
|
# (Optional) Call subscribers after loop
|
||||||
Callbacks.after_loop(
|
Callbacks.after_loop(
|
||||||
|
|||||||
@ -8,6 +8,7 @@ from mflux.config.model_config import ModelConfig
|
|||||||
from mflux.config.runtime_config import RuntimeConfig
|
from mflux.config.runtime_config import RuntimeConfig
|
||||||
from mflux.controlnet.controlnet_util import ControlnetUtil
|
from mflux.controlnet.controlnet_util import ControlnetUtil
|
||||||
from mflux.controlnet.transformer_controlnet import TransformerControlnet
|
from mflux.controlnet.transformer_controlnet import TransformerControlnet
|
||||||
|
from mflux.error.exceptions import StopImageGenerationException
|
||||||
from mflux.flux.flux_initializer import FluxInitializer
|
from mflux.flux.flux_initializer import FluxInitializer
|
||||||
from mflux.latent_creator.latent_creator import LatentCreator
|
from mflux.latent_creator.latent_creator import LatentCreator
|
||||||
from mflux.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder
|
from mflux.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder
|
||||||
@ -141,6 +142,7 @@ class Flux1Controlnet(nn.Module):
|
|||||||
config=config,
|
config=config,
|
||||||
time_steps=time_steps,
|
time_steps=time_steps,
|
||||||
)
|
)
|
||||||
|
raise StopImageGenerationException(f"Stopping image generation at step {t + 1}/{len(time_steps)}")
|
||||||
|
|
||||||
# (Optional) Call subscribers after loop
|
# (Optional) Call subscribers after loop
|
||||||
Callbacks.after_loop(
|
Callbacks.after_loop(
|
||||||
|
|||||||
@ -6,6 +6,7 @@ from mflux.callbacks.callbacks import Callbacks
|
|||||||
from mflux.config.config import Config
|
from mflux.config.config import Config
|
||||||
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.error.exceptions import StopImageGenerationException
|
||||||
from mflux.flux.flux_initializer import FluxInitializer
|
from mflux.flux.flux_initializer import FluxInitializer
|
||||||
from mflux.latent_creator.latent_creator import Img2Img, LatentCreator
|
from mflux.latent_creator.latent_creator import Img2Img, LatentCreator
|
||||||
from mflux.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder
|
from mflux.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder
|
||||||
@ -121,6 +122,7 @@ class Flux1(nn.Module):
|
|||||||
config=config,
|
config=config,
|
||||||
time_steps=time_steps,
|
time_steps=time_steps,
|
||||||
)
|
)
|
||||||
|
raise StopImageGenerationException(f"Stopping image generation at step {t + 1}/{len(time_steps)}")
|
||||||
|
|
||||||
# (Optional) Call subscribers after loop
|
# (Optional) Call subscribers after loop
|
||||||
Callbacks.after_loop(
|
Callbacks.after_loop(
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user