diff --git a/.gitignore b/.gitignore index fd92c93..ee9ae73 100644 --- a/.gitignore +++ b/.gitignore @@ -10,3 +10,4 @@ .venv *.png *.jpg +*.pyc diff --git a/main.py b/main.py index 60273f3..c5a0469 100644 --- a/main.py +++ b/main.py @@ -4,9 +4,9 @@ import sys sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), 'src'))) from flux_1_schnell.config.config import Config -from flux_1_schnell.flux import Flux1Schnell +from flux_1_schnell.flux import Flux1 -flux = Flux1Schnell("black-forest-labs/FLUX.1-schnell") +flux = Flux1("black-forest-labs/FLUX.1-schnell") image = flux.generate_image( seed=3, diff --git a/src/flux_1_schnell/flux.py b/src/flux_1_schnell/flux.py index d2d3299..1fd564f 100644 --- a/src/flux_1_schnell/flux.py +++ b/src/flux_1_schnell/flux.py @@ -10,14 +10,13 @@ from flux_1_schnell.models.text_encoder.t5_encoder.t5_encoder import T5Encoder from flux_1_schnell.models.transformer.transformer import Transformer from flux_1_schnell.models.vae.vae import VAE from flux_1_schnell.post_processing.image_util import ImageUtil -from flux_1_schnell.scheduler.scheduler import FlowMatchEulerDiscreteNoiseScheduler from flux_1_schnell.tokenizer.clip_tokenizer import TokenizerCLIP from flux_1_schnell.tokenizer.t5_tokenizer import TokenizerT5 from flux_1_schnell.tokenizer.tokenizer_handler import TokenizerHandler from flux_1_schnell.weights.weight_handler import WeightHandler -class Flux1Schnell: +class Flux1: def __init__(self, repo_id: str): tokenizers = TokenizerHandler.load_from_disk_or_huggingface(repo_id) @@ -47,16 +46,12 @@ class Flux1Schnell: config=config ) - latents = FlowMatchEulerDiscreteNoiseScheduler.denoise( - t=t, - noise=noise, - latent=latents, - config=config - ) + dt = config.sigmas[t + 1] - config.sigmas[t] + latents += noise * dt mx.eval(latents) - latents = Flux1Schnell._unpack_latents(latents, config.width, config.height) + latents = Flux1._unpack_latents(latents, config.width, config.height) decoded = self.vae.decode(latents) return ImageUtil.to_image(decoded) diff --git a/src/flux_1_schnell/scheduler/__init__.py b/src/flux_1_schnell/scheduler/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/src/flux_1_schnell/scheduler/scheduler.py b/src/flux_1_schnell/scheduler/scheduler.py deleted file mode 100644 index 3fb3e7b..0000000 --- a/src/flux_1_schnell/scheduler/scheduler.py +++ /dev/null @@ -1,20 +0,0 @@ -import mlx.core as mx - -from flux_1_schnell.config.config import Config - - -class FlowMatchEulerDiscreteNoiseScheduler: - - @staticmethod - def denoise( - t: int, - noise: mx.array, - latent: mx.array, - config: Config, - ) -> mx.array: - sigma = config.sigmas[t] - denoised = latent - noise * sigma - derivative = (latent - denoised) / sigma - dt = config.sigmas[t + 1] - sigma - prev_sample = latent + derivative * dt - return prev_sample