fix merge conflicts
This commit is contained in:
commit
16ee26c2e9
@ -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).
|
||||||
|
|||||||
2
main.py
2
main.py
@ -19,4 +19,4 @@ image = flux.generate_image(
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
ImageUtil.save_image(image, "image.png")
|
ImageUtil.save_image(image, "image.png")
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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:
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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}")
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user