From ce7dfeabfdd7c4e45487dbc8ad8c5a2309fa9fbe Mon Sep 17 00:00:00 2001 From: Anthony Wu Date: Fri, 7 Mar 2025 20:27:41 -0800 Subject: [PATCH] remedy missed ruff check/formats in 0.6.0 release --- src/mflux/callbacks/callback_registry.py | 2 +- src/mflux/callbacks/callbacks.py | 7 +----- src/mflux/callbacks/instances/memory_saver.py | 22 ++++++++----------- src/mflux/dreambooth/state/training_state.py | 4 ++-- src/mflux/generate.py | 2 +- src/mflux/generate_controlnet.py | 3 ++- src/mflux/ui/cli/parsers.py | 2 +- tools/rename_images.py | 4 ++-- 8 files changed, 19 insertions(+), 27 deletions(-) diff --git a/src/mflux/callbacks/callback_registry.py b/src/mflux/callbacks/callback_registry.py index c2caadd..7429075 100644 --- a/src/mflux/callbacks/callback_registry.py +++ b/src/mflux/callbacks/callback_registry.py @@ -1,4 +1,4 @@ -from mflux.callbacks.callback import BeforeLoopCallback, InLoopCallback, InterruptCallback, AfterLoopCallback +from mflux.callbacks.callback import AfterLoopCallback, BeforeLoopCallback, InLoopCallback, InterruptCallback class CallbackRegistry: diff --git a/src/mflux/callbacks/callbacks.py b/src/mflux/callbacks/callbacks.py index ba30e76..f17af70 100644 --- a/src/mflux/callbacks/callbacks.py +++ b/src/mflux/callbacks/callbacks.py @@ -51,12 +51,7 @@ class Callbacks: config: RuntimeConfig ): # fmt: off for subscriber in CallbackRegistry.after_loop_callbacks(): - subscriber.call_after_loop( - seed=seed, - prompt=prompt, - latents=latents, - config=config - ) + subscriber.call_after_loop(seed=seed, prompt=prompt, latents=latents, config=config) @staticmethod def interruption( diff --git a/src/mflux/callbacks/instances/memory_saver.py b/src/mflux/callbacks/instances/memory_saver.py index 2adabd6..bd928b9 100644 --- a/src/mflux/callbacks/instances/memory_saver.py +++ b/src/mflux/callbacks/instances/memory_saver.py @@ -4,27 +4,23 @@ import mlx.core as mx import PIL.Image from tqdm import tqdm -from mflux.callbacks.callback import ( - AfterLoopCallback, - BeforeLoopCallback, - InLoopCallback -) +from mflux.callbacks.callback import AfterLoopCallback, BeforeLoopCallback, InLoopCallback from mflux.config.runtime_config import RuntimeConfig class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback): """ - Optimizes memory usage by clearing caches and removing unused model + Optimizes memory usage by clearing caches and removing unused model components at strategic points in the execution cycle. """ - + def __init__(self, flux, cache_limit_bytes: int = 1000**3): self.flux = flux self.peak_memory: int = 0 mx.metal.set_cache_limit(cache_limit_bytes) mx.metal.clear_cache() mx.metal.reset_peak_memory() - + def call_before_loop( self, seed: int, @@ -35,7 +31,7 @@ class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback): ) -> None: self.peak_memory = mx.metal.get_peak_memory() self._delete_encoders() - + def call_in_loop( self, t: int, @@ -46,7 +42,7 @@ class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback): time_steps: tqdm, ) -> None: self.peak_memory = mx.metal.get_peak_memory() - + def call_after_loop( self, seed: int, @@ -56,16 +52,16 @@ class MemorySaver(BeforeLoopCallback, InLoopCallback, AfterLoopCallback): ) -> None: self.peak_memory = mx.metal.get_peak_memory() self._delete_transformer() - + def _delete_encoders(self) -> None: self.flux.clip_text_encoder = None self.flux.t5_text_encoder = None gc.collect() mx.metal.clear_cache() - + def _delete_transformer(self) -> None: self.flux.transformer = None - if hasattr(self.flux, 'transformer_controlnet'): + if hasattr(self.flux, "transformer_controlnet"): self.flux.transformer_controlnet = None gc.collect() mx.metal.clear_cache() diff --git a/src/mflux/dreambooth/state/training_state.py b/src/mflux/dreambooth/state/training_state.py index da2aca7..a9102de 100644 --- a/src/mflux/dreambooth/state/training_state.py +++ b/src/mflux/dreambooth/state/training_state.py @@ -108,13 +108,13 @@ class TrainingState: def get_current_validation_image_path(self, training_spec: TrainingSpec) -> Path: output_path = Path(training_spec.saver.output_path) / DREAMBOOTH_PATH_VALIDATION_IMAGES output_path.mkdir(parents=True, exist_ok=True) - path = output_path / Path(f"{self.iterator.num_iterations :07d}_{DREAMBOOTH_FILE_NAME_VALIDATION_IMAGE}.png") + path = output_path / Path(f"{self.iterator.num_iterations:07d}_{DREAMBOOTH_FILE_NAME_VALIDATION_IMAGE}.png") return path def get_current_loss_plot_path(self, training_spec: TrainingSpec) -> Path: output_path = Path(training_spec.saver.output_path) / DREAMBOOTH_PATH_VALIDATION_PLOT output_path.mkdir(parents=True, exist_ok=True) - path = output_path / Path(f"{self.iterator.num_iterations :07d}_{DREAMBOOTH_FILE_NAME_VALIDATION_LOSS}.pdf") + path = output_path / Path(f"{self.iterator.num_iterations:07d}_{DREAMBOOTH_FILE_NAME_VALIDATION_LOSS}.pdf") return path @staticmethod diff --git a/src/mflux/generate.py b/src/mflux/generate.py index 5a443e1..49508da 100644 --- a/src/mflux/generate.py +++ b/src/mflux/generate.py @@ -1,7 +1,7 @@ from mflux import Config, Flux1, ModelConfig, StopImageGenerationException from mflux.callbacks.callback_registry import CallbackRegistry -from mflux.callbacks.instances.stepwise_handler import StepwiseHandler from mflux.callbacks.instances.memory_saver import MemorySaver +from mflux.callbacks.instances.stepwise_handler import StepwiseHandler from mflux.ui.cli.parsers import CommandLineParser diff --git a/src/mflux/generate_controlnet.py b/src/mflux/generate_controlnet.py index 9df5e83..90efa40 100644 --- a/src/mflux/generate_controlnet.py +++ b/src/mflux/generate_controlnet.py @@ -1,8 +1,8 @@ from mflux import Config, Flux1Controlnet, ModelConfig, StopImageGenerationException from mflux.callbacks.callback_registry import CallbackRegistry from mflux.callbacks.instances.canny_saver import CannyImageSaver -from mflux.callbacks.instances.stepwise_handler import StepwiseHandler from mflux.callbacks.instances.memory_saver import MemorySaver +from mflux.callbacks.instances.stepwise_handler import StepwiseHandler from mflux.ui.cli.parsers import CommandLineParser @@ -66,5 +66,6 @@ def main(): if memory_saver: print(memory_saver.memory_stats()) + if __name__ == "__main__": main() diff --git a/src/mflux/ui/cli/parsers.py b/src/mflux/ui/cli/parsers.py index b62cd88..ad898df 100644 --- a/src/mflux/ui/cli/parsers.py +++ b/src/mflux/ui/cli/parsers.py @@ -17,7 +17,7 @@ class ModelSpecAction(argparse.Action): if values.count("/") != 1: raise argparse.ArgumentError( - self, 'Value must be either "dev", "schnell", or "' f'in format "org/model". Got: {values}' + self, f'Value must be either "dev", "schnell", or "in format "org/model". Got: {values}' ) # If we got here, values contains exactly one slash diff --git a/tools/rename_images.py b/tools/rename_images.py index 89e737b..160d471 100644 --- a/tools/rename_images.py +++ b/tools/rename_images.py @@ -51,10 +51,10 @@ def parse_metadata_from_image(image: Image, insecure=False) -> dict: metadata = obj else: raise UnsupportedMetadata( - "The metadata is not a dict output recognized by this tool. " f"Metadata: {value.decode()}" + f"The metadata is not a dict output recognized by this tool. Metadata: {value.decode()}" ) except (ValueError, SyntaxError): - raise UnsupportedMetadata("The metadata is not parseable by this tool. " f"Metadata: {value.decode()}") + raise UnsupportedMetadata(f"The metadata is not parseable by this tool. Metadata: {value.decode()}") return metadata