fix type and add main_dev
This commit is contained in:
parent
d14f49f758
commit
68cbaab816
4
.gitignore
vendored
4
.gitignore
vendored
@ -11,7 +11,3 @@
|
|||||||
*.png
|
*.png
|
||||||
*.jpg
|
*.jpg
|
||||||
*.pyc
|
*.pyc
|
||||||
|
|
||||||
|
|
||||||
*.pt
|
|
||||||
trial.py
|
|
||||||
6
main.py
6
main.py
@ -6,15 +6,15 @@ 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.config.config import Config
|
||||||
from flux_1_schnell.flux import Flux1
|
from flux_1_schnell.flux import Flux1
|
||||||
|
|
||||||
flux = Flux1("black-forest-labs/FLUX.1-dev")
|
flux = Flux1("black-forest-labs/FLUX.1-dev", max_sequence_length=512)
|
||||||
|
|
||||||
image = flux.generate_image(
|
image = flux.generate_image(
|
||||||
seed=3,
|
seed=3,
|
||||||
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=20,
|
num_inference_steps=20,
|
||||||
width=256,
|
height=768,
|
||||||
height=256,
|
width=1360,
|
||||||
guidance=3.5,
|
guidance=3.5,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|||||||
@ -5,7 +5,7 @@ import logging
|
|||||||
log = logging.getLogger(__name__)
|
log = logging.getLogger(__name__)
|
||||||
|
|
||||||
class Config:
|
class Config:
|
||||||
precision: mx.Dtype = mx.float16
|
precision: mx.Dtype = mx.bfloat16
|
||||||
num_train_steps = 1000
|
num_train_steps = 1000
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@ -37,3 +37,4 @@ class Config:
|
|||||||
b = y1 - m * x1
|
b = y1 - m * x1
|
||||||
mu = m * self.width * self.height / 256 + b
|
mu = m * self.width * self.height / 256 + b
|
||||||
self.sigmas = mx.exp(mu) / (mx.exp(mu) + (1 / self.sigmas - 1))
|
self.sigmas = mx.exp(mu) / (mx.exp(mu) + (1 / self.sigmas - 1))
|
||||||
|
self.sigmas[-1] = 0
|
||||||
|
|||||||
@ -18,11 +18,10 @@ from flux_1_schnell.weights.weight_handler import WeightHandler
|
|||||||
|
|
||||||
class Flux1:
|
class Flux1:
|
||||||
|
|
||||||
def __init__(self, repo_id: str):
|
def __init__(self, repo_id: str, max_sequence_length: int = 512):
|
||||||
self.is_dev = "FLUX.1-dev" in repo_id
|
self.is_dev = "FLUX.1-dev" in repo_id
|
||||||
max_t5_length = 512 if self.is_dev else 256
|
tokenizers = TokenizerHandler.load_from_disk_or_huggingface(repo_id, max_sequence_length)
|
||||||
tokenizers = TokenizerHandler.load_from_disk_or_huggingface(repo_id, max_t5_length)
|
self.t5_tokenizer = TokenizerT5(tokenizers.t5, max_length=max_sequence_length)
|
||||||
self.t5_tokenizer = TokenizerT5(tokenizers.t5, max_length=max_t5_length)
|
|
||||||
self.clip_tokenizer = TokenizerCLIP(tokenizers.clip)
|
self.clip_tokenizer = TokenizerCLIP(tokenizers.clip)
|
||||||
|
|
||||||
weights = WeightHandler.load_from_disk_or_huggingface(repo_id)
|
weights = WeightHandler.load_from_disk_or_huggingface(repo_id)
|
||||||
|
|||||||
@ -35,9 +35,9 @@ class Transformer(nn.Module):
|
|||||||
config: Config
|
config: Config
|
||||||
) -> mx.array:
|
) -> mx.array:
|
||||||
time_step = config.sigmas[t] * config.num_train_steps
|
time_step = config.sigmas[t] * config.num_train_steps
|
||||||
time_step = mx.broadcast_to(time_step, (1,))
|
time_step = mx.broadcast_to(time_step, (1,)).astype(config.precision)
|
||||||
hidden_states = self.x_embedder(hidden_states)
|
hidden_states = self.x_embedder(hidden_states)
|
||||||
guidance = mx.broadcast_to(config.guidance * config.num_train_steps, (1,))
|
guidance = mx.broadcast_to(config.guidance * config.num_train_steps, (1,)).astype(config.precision)
|
||||||
text_embeddings = self.time_text_embed.forward(time_step, pooled_prompt_embeds, guidance)
|
text_embeddings = self.time_text_embed.forward(time_step, pooled_prompt_embeds, guidance)
|
||||||
encoder_hidden_states = self.context_embedder(prompt_embeds)
|
encoder_hidden_states = self.context_embedder(prompt_embeds)
|
||||||
txt_ids = Transformer._prepare_text_ids(seq_len = prompt_embeds.shape[1])
|
txt_ids = Transformer._prepare_text_ids(seq_len = prompt_embeds.shape[1])
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user