fix merge conflicts

This commit is contained in:
Fabio 2024-08-18 23:33:22 +02:00
commit 16ee26c2e9
6 changed files with 38 additions and 11 deletions

View File

@ -51,6 +51,7 @@ sys.path.append("/path/to/mflux/src")
from flux_1_schnell.config.config import Config from flux_1_schnell.config.config import Config
from flux_1_schnell.flux import Flux1Schnell from flux_1_schnell.flux import Flux1Schnell
from flux_1_schnell.post_processing.image_util import ImageUtil
flux = Flux1Schnell("black-forest-labs/FLUX.1-schnell") flux = Flux1Schnell("black-forest-labs/FLUX.1-schnell")
@ -62,7 +63,7 @@ image = flux.generate_image(
) )
) )
image.save("image.png") ImageUtil.save_image(image, "image.png")
``` ```
If the model is not already downloaded on your machine, it will start the download process and fetch the model weights (~34GB in size for the Schnell model). If the model is not already downloaded on your machine, it will start the download process and fetch the model weights (~34GB in size for the Schnell model).

View File

@ -19,4 +19,4 @@ image = flux.generate_image(
) )
) )
ImageUtil.save_image(image, "image.png") ImageUtil.save_image(image, "image.png")

View File

@ -4,6 +4,7 @@ import logging
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
class Config: class Config:
precision: mx.Dtype = mx.bfloat16 precision: mx.Dtype = mx.bfloat16
num_train_steps = 1000 num_train_steps = 1000

View File

@ -59,10 +59,10 @@ class Flux1:
return ImageUtil.to_image(decoded) return ImageUtil.to_image(decoded)
@staticmethod @staticmethod
def _unpack_latents(latents: mx.array, width: int, height: int) -> mx.array: def _unpack_latents(latents: mx.array, height: int, width: int) -> mx.array:
latents = mx.reshape(latents, (1, height//16, width//16, 16, 2, 2)) latents = mx.reshape(latents, (1, width // 16, height // 16, 16, 2, 2))
latents = mx.transpose(latents, (0, 3, 1, 4, 2, 5)) latents = mx.transpose(latents, (0, 3, 1, 4, 2, 5))
latents = mx.reshape(latents, (1, 16, height//16 *2, width//16 * 2)) latents = mx.reshape(latents, (1, 16, width // 16 * 2, height // 16 * 2))
return latents return latents
def encode(self, path: str) -> mx.array: def encode(self, path: str) -> mx.array:

View File

@ -69,14 +69,14 @@ class Transformer(nn.Module):
return noise return noise
@staticmethod @staticmethod
def _prepare_latent_image_ids(width: int, height: int) -> mx.array: def _prepare_latent_image_ids(height: int, width: int) -> mx.array:
latent_height = height // 16
latent_width = width // 16 latent_width = width // 16
latent_image_ids = mx.zeros((latent_height, latent_width, 3)) latent_height = height // 16
latent_image_ids = latent_image_ids.at[:, :, 1].add(mx.arange(0, latent_height)[:, None]) latent_image_ids = mx.zeros((latent_width, latent_height, 3))
latent_image_ids = latent_image_ids.at[:, :, 2].add(mx.arange(0, latent_width)[None, :]) latent_image_ids = latent_image_ids.at[:, :, 1].add(mx.arange(0, latent_width)[:, None])
latent_image_ids = latent_image_ids.at[:, :, 2].add(mx.arange(0, latent_height)[None, :])
latent_image_ids = mx.repeat(latent_image_ids[None, :], 1, axis=0) latent_image_ids = mx.repeat(latent_image_ids[None, :], 1, axis=0)
latent_image_ids = mx.reshape(latent_image_ids, (1, latent_height * latent_width, 3)) latent_image_ids = mx.reshape(latent_image_ids, (1, latent_width * latent_height, 3))
return latent_image_ids return latent_image_ids
@staticmethod @staticmethod

View File

@ -1,8 +1,13 @@
import logging
from pathlib import Path
import PIL import PIL
import mlx.core as mx import mlx.core as mx
import numpy as np import numpy as np
from PIL import Image from PIL import Image
log = logging.getLogger(__name__)
class ImageUtil: class ImageUtil:
@ -53,3 +58,23 @@ class ImageUtil:
def resize(image): def resize(image):
image = image.resize((1024, 1024), resample=PIL.Image.LANCZOS) image = image.resize((1024, 1024), resample=PIL.Image.LANCZOS)
return image return image
@staticmethod
def save_image(image: Image.Image, path: str) -> None:
file_path = Path(path)
file_path.parent.mkdir(parents=True, exist_ok=True)
file_name = file_path.stem
file_extension = file_path.suffix
# If a file already exists, create a new name with a counter
counter = 1
while file_path.exists():
new_name = f"{file_name}({counter}){file_extension}"
file_path = file_path.with_name(new_name)
counter += 1
try:
image.save(file_path)
log.info(f"Image saved successfully at: {file_path}")
except Exception as e:
log.info(f"Error saving image: {e}")