remedy missed ruff check/formats in 0.6.0 release

This commit is contained in:
Anthony Wu 2025-03-07 20:27:41 -08:00
parent e13181e176
commit ce7dfeabfd
8 changed files with 19 additions and 27 deletions

View File

@ -1,4 +1,4 @@
from mflux.callbacks.callback import BeforeLoopCallback, InLoopCallback, InterruptCallback, AfterLoopCallback
from mflux.callbacks.callback import AfterLoopCallback, BeforeLoopCallback, InLoopCallback, InterruptCallback
class CallbackRegistry:

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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