From 1cecc60e5e64323f5809cd45ad47adc75aaba792 Mon Sep 17 00:00:00 2001 From: filipstrand Date: Mon, 12 Aug 2024 21:37:41 +0200 Subject: [PATCH] Add progress bar --- requirements.txt | 3 ++- src/flux_1_schnell/models/flux.py | 5 ++++- src/flux_1_schnell/models/transformer/transformer.py | 2 -- 3 files changed, 6 insertions(+), 4 deletions(-) diff --git a/requirements.txt b/requirements.txt index bf4f080..dbc01f7 100644 --- a/requirements.txt +++ b/requirements.txt @@ -3,4 +3,5 @@ numpy>=2.0.1 pillow>=10.4.0 transformers>=4.44.0 sentencepiece>=0.2.0 -torch>=2.3.1 \ No newline at end of file +torch>=2.3.1 +tqdm>=4.66.5 \ No newline at end of file diff --git a/src/flux_1_schnell/models/flux.py b/src/flux_1_schnell/models/flux.py index b11b602..3a7ab9d 100644 --- a/src/flux_1_schnell/models/flux.py +++ b/src/flux_1_schnell/models/flux.py @@ -1,6 +1,7 @@ import PIL import mlx.core as mx from PIL import Image +from tqdm import tqdm from flux_1_schnell.models.text_encoder.clip_encoder.clip_encoder import CLIPEncoder from flux_1_schnell.tokenizer.clip_tokenizer import TokenizerCLIP @@ -41,7 +42,7 @@ class Flux1Schnell: t5_text_encoder=self.t5_text_encoder ) - for t in range(config.num_inference_steps): + for t in tqdm(range(config.num_inference_steps)): noise = self.transformer.predict( t=t, prompt_embeds=prompt_embeds, @@ -57,6 +58,8 @@ class Flux1Schnell: config=config ) + mx.eval(latents) + latents = Flux1Schnell._unpack_latents(latents) decoded = self.vae.decode(latents) return ImageUtil.to_image(decoded) diff --git a/src/flux_1_schnell/models/transformer/transformer.py b/src/flux_1_schnell/models/transformer/transformer.py index 67c099e..94ca0ce 100644 --- a/src/flux_1_schnell/models/transformer/transformer.py +++ b/src/flux_1_schnell/models/transformer/transformer.py @@ -64,8 +64,6 @@ class Transformer(nn.Module): hidden_states = self.norm_out.forward(hidden_states, text_embeddings) hidden_states = self.proj_out(hidden_states) noise = hidden_states - mx.eval(noise) - print(t) return noise @staticmethod