Merge pull request #4 from Xuzzo/feature/allow_image_resolution
Generate image at different resolutions
This commit is contained in:
commit
f0ca12f075
1
.gitignore
vendored
1
.gitignore
vendored
@ -10,3 +10,4 @@
|
|||||||
.venv
|
.venv
|
||||||
*.png
|
*.png
|
||||||
*.jpg
|
*.jpg
|
||||||
|
*.pyc
|
||||||
|
|||||||
@ -129,7 +129,6 @@ Luxury food photograph of an italian Linguine pasta alle vongole dish with lots
|
|||||||
|
|
||||||
- Images are generated one by one.
|
- Images are generated one by one.
|
||||||
- Negative prompts not supported.
|
- Negative prompts not supported.
|
||||||
- Image resolution is (1024, 1024)
|
|
||||||
|
|
||||||
### TODO
|
### TODO
|
||||||
|
|
||||||
|
|||||||
2
main.py
2
main.py
@ -13,6 +13,8 @@ image = flux.generate_image(
|
|||||||
prompt="Luxury food photograph of a birthday cake. In the middle it has three candles shaped like letters spelling the word 'MLX'. It has perfect lighting and a cozy background with big bokeh and shallow depth of field. The mood is a sunset balcony in tuscany. The photo is taken from the side of the cake. The scene is complemented by a warm, inviting light that highlights the textures and colors of the ingredients, giving it an appetizing and elegant look.",
|
prompt="Luxury food photograph of a birthday cake. In the middle it has three candles shaped like letters spelling the word 'MLX'. It has perfect lighting and a cozy background with big bokeh and shallow depth of field. The mood is a sunset balcony in tuscany. The photo is taken from the side of the cake. The scene is complemented by a warm, inviting light that highlights the textures and colors of the ingredients, giving it an appetizing and elegant look.",
|
||||||
config=Config(
|
config=Config(
|
||||||
num_inference_steps=2,
|
num_inference_steps=2,
|
||||||
|
height=768,
|
||||||
|
width=1360,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@ -1,6 +1,8 @@
|
|||||||
import mlx.core as mx
|
import mlx.core as mx
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
import logging
|
||||||
|
|
||||||
|
log = logging.getLogger(__name__)
|
||||||
|
|
||||||
class Config:
|
class Config:
|
||||||
precision: mx.Dtype = mx.float16
|
precision: mx.Dtype = mx.float16
|
||||||
@ -9,7 +11,13 @@ class Config:
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
num_inference_steps: int = 4,
|
num_inference_steps: int = 4,
|
||||||
|
width: int = 1024,
|
||||||
|
height: int = 1024,
|
||||||
):
|
):
|
||||||
|
if width % 16 != 0 or height % 16 != 0:
|
||||||
|
log.warning("Width and height should be multiples of 16. Rounding down.")
|
||||||
|
self.width = 16 * (height // 16)
|
||||||
|
self.height = 16 * (width // 16)
|
||||||
base_sigmas = Config.base_sigmas(num_inference_steps)
|
base_sigmas = Config.base_sigmas(num_inference_steps)
|
||||||
self.num_inference_steps = num_inference_steps
|
self.num_inference_steps = num_inference_steps
|
||||||
self.time_steps = base_sigmas * self.num_train_steps
|
self.time_steps = base_sigmas * self.num_train_steps
|
||||||
|
|||||||
@ -31,7 +31,7 @@ class Flux1Schnell:
|
|||||||
self.clip_text_encoder = CLIPEncoder(weights.clip_encoder)
|
self.clip_text_encoder = CLIPEncoder(weights.clip_encoder)
|
||||||
|
|
||||||
def generate_image(self, seed: int, prompt: str, config: Config = Config()) -> PIL.Image.Image:
|
def generate_image(self, seed: int, prompt: str, config: Config = Config()) -> PIL.Image.Image:
|
||||||
latents = LatentCreator.create(seed)
|
latents = LatentCreator.create(config.height, config.width, seed)
|
||||||
|
|
||||||
t5_tokens = self.t5_tokenizer.tokenize(prompt)
|
t5_tokens = self.t5_tokenizer.tokenize(prompt)
|
||||||
clip_tokens = self.clip_tokenizer.tokenize(prompt)
|
clip_tokens = self.clip_tokenizer.tokenize(prompt)
|
||||||
@ -56,15 +56,15 @@ class Flux1Schnell:
|
|||||||
|
|
||||||
mx.eval(latents)
|
mx.eval(latents)
|
||||||
|
|
||||||
latents = Flux1Schnell._unpack_latents(latents)
|
latents = Flux1Schnell._unpack_latents(latents, config.height, config.width)
|
||||||
decoded = self.vae.decode(latents)
|
decoded = self.vae.decode(latents)
|
||||||
return ImageUtil.to_image(decoded)
|
return ImageUtil.to_image(decoded)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _unpack_latents(latents):
|
def _unpack_latents(latents: mx.array, width: int, height: int) -> mx.array:
|
||||||
latents = mx.reshape(latents, (1, 64, 64, 16, 2, 2))
|
latents = mx.reshape(latents, (1, height//16, width//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, 128, 128))
|
latents = mx.reshape(latents, (1, 16, height//16 *2, width//16 * 2))
|
||||||
return latents
|
return latents
|
||||||
|
|
||||||
def encode(self, path: str) -> mx.array:
|
def encode(self, path: str) -> mx.array:
|
||||||
|
|||||||
@ -4,6 +4,8 @@ import mlx.core as mx
|
|||||||
class LatentCreator:
|
class LatentCreator:
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def create(seed: int) -> (mx.array, mx.array):
|
def create(height: int, width: int, seed: int) -> (mx.array, mx.array):
|
||||||
latents = mx.random.normal(shape=[1, 4096, 64], key=mx.random.key(seed))
|
latent_height = height // 16
|
||||||
|
latent_width = width // 16
|
||||||
|
latents = mx.random.normal(shape=[1, latent_height * latent_width, 64], key=mx.random.key(seed))
|
||||||
return latents
|
return latents
|
||||||
|
|||||||
@ -39,7 +39,7 @@ class Transformer(nn.Module):
|
|||||||
text_embeddings = self.time_text_embed.forward(time_step, pooled_prompt_embeds)
|
text_embeddings = self.time_text_embed.forward(time_step, pooled_prompt_embeds)
|
||||||
encoder_hidden_states = self.context_embedder(prompt_embeds)
|
encoder_hidden_states = self.context_embedder(prompt_embeds)
|
||||||
txt_ids = Transformer._prepare_text_ids()
|
txt_ids = Transformer._prepare_text_ids()
|
||||||
img_ids = Transformer._prepare_latent_image_ids()
|
img_ids = Transformer._prepare_latent_image_ids(config.height, config.width)
|
||||||
ids = mx.concatenate((txt_ids, img_ids), axis=1)
|
ids = mx.concatenate((txt_ids, img_ids), axis=1)
|
||||||
image_rotary_emb = self.pos_embed.forward(ids)
|
image_rotary_emb = self.pos_embed.forward(ids)
|
||||||
|
|
||||||
@ -67,12 +67,14 @@ class Transformer(nn.Module):
|
|||||||
return noise
|
return noise
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _prepare_latent_image_ids() -> mx.array:
|
def _prepare_latent_image_ids(width: int, height: int) -> mx.array:
|
||||||
latent_image_ids = mx.zeros((64, 64, 3))
|
latent_height = height // 16
|
||||||
latent_image_ids = latent_image_ids.at[:, :, 1].add(mx.arange(0, 64)[:, None])
|
latent_width = width // 16
|
||||||
latent_image_ids = latent_image_ids.at[:, :, 2].add(mx.arange(0, 64)[None, :])
|
latent_image_ids = mx.zeros((latent_height, latent_width, 3))
|
||||||
|
latent_image_ids = latent_image_ids.at[:, :, 1].add(mx.arange(0, latent_height)[:, None])
|
||||||
|
latent_image_ids = latent_image_ids.at[:, :, 2].add(mx.arange(0, latent_width)[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, 4096, 3))
|
latent_image_ids = mx.reshape(latent_image_ids, (1, latent_height * latent_width, 3))
|
||||||
return latent_image_ids
|
return latent_image_ids
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user