WIP include flux1 dev

This commit is contained in:
Fabio 2024-08-17 22:33:11 +02:00
parent 7661622839
commit 7b1eba235d
5 changed files with 7 additions and 31 deletions

1
.gitignore vendored
View File

@ -10,3 +10,4 @@
.venv
*.png
*.jpg
*.pyc

View File

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

View File

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

View File

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