add shift to schedule
This commit is contained in:
parent
14c5b2bd0e
commit
40b40b5541
3
main.py
3
main.py
@ -12,9 +12,10 @@ 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=2,
|
num_inference_steps=20,
|
||||||
width=256,
|
width=256,
|
||||||
height=256,
|
height=256,
|
||||||
|
guidance=3.5,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@ -13,18 +13,26 @@ class Config:
|
|||||||
num_inference_steps: int = 4,
|
num_inference_steps: int = 4,
|
||||||
width: int = 1024,
|
width: int = 1024,
|
||||||
height: int = 1024,
|
height: int = 1024,
|
||||||
|
guidance: float = 4.0,
|
||||||
):
|
):
|
||||||
if width %16 != 0 or height % 16 != 0:
|
if width %16 != 0 or height % 16 != 0:
|
||||||
log.warning("Width and height should be multiples of 16. Rounding down.")
|
log.warning("Width and height should be multiples of 16. Rounding down.")
|
||||||
self.width = 16 * (height // 16)
|
self.width = 16 * (height // 16)
|
||||||
self.height = 16 * (width // 16)
|
self.height = 16 * (width // 16)
|
||||||
base_sigmas = Config.base_sigmas(num_inference_steps)
|
self.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.guidance = guidance
|
||||||
self.sigmas = mx.concatenate([base_sigmas, mx.zeros(1)])
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def base_sigmas(num_inference_steps):
|
def base_sigmas(num_inference_steps):
|
||||||
sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps)
|
sigmas = np.linspace(1.0, 0, num_inference_steps+1)
|
||||||
sigmas = mx.array(sigmas).astype(mx.float32)
|
sigmas = mx.array(sigmas).astype(mx.float32)
|
||||||
return sigmas
|
return sigmas
|
||||||
|
|
||||||
|
def shift_sigmas(self):
|
||||||
|
y1 = 0.5
|
||||||
|
x1 = 256
|
||||||
|
m = (1.15 - y1) / (4096 - x1)
|
||||||
|
b = y1 - m * x1
|
||||||
|
mu = m + b
|
||||||
|
self.sigmas = mx.exp(mu) / (mx.exp(mu) + (1 / self.sigmas - 1))
|
||||||
|
|||||||
@ -19,8 +19,9 @@ 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):
|
||||||
is_dev = "FLUX.1-dev" in repo_id
|
self.is_dev = "FLUX.1-dev" in repo_id
|
||||||
max_t5_length = 512 if is_dev else 256
|
# max_t5_length = 512 if self.is_dev else 256
|
||||||
|
max_t5_length = 256
|
||||||
tokenizers = TokenizerHandler.load_from_disk_or_huggingface(repo_id, max_t5_length)
|
tokenizers = TokenizerHandler.load_from_disk_or_huggingface(repo_id, max_t5_length)
|
||||||
self.t5_tokenizer = TokenizerT5(tokenizers.t5, max_length=max_t5_length)
|
self.t5_tokenizer = TokenizerT5(tokenizers.t5, max_length=max_t5_length)
|
||||||
self.clip_tokenizer = TokenizerCLIP(tokenizers.clip)
|
self.clip_tokenizer = TokenizerCLIP(tokenizers.clip)
|
||||||
@ -32,6 +33,8 @@ class Flux1:
|
|||||||
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:
|
||||||
|
if self.is_dev:
|
||||||
|
config.shift_sigmas()
|
||||||
latents = LatentCreator.create(config.height, config.width, seed)
|
latents = LatentCreator.create(config.height, config.width, seed)
|
||||||
|
|
||||||
t5_tokens = self.t5_tokenizer.tokenize(prompt)
|
t5_tokens = self.t5_tokenizer.tokenize(prompt)
|
||||||
|
|||||||
@ -15,15 +15,14 @@ class TimeTextEmbed(nn.Module):
|
|||||||
self.text_embedder = TextEmbedder()
|
self.text_embedder = TextEmbedder()
|
||||||
self.with_guidance_embed = with_guidance_embed
|
self.with_guidance_embed = with_guidance_embed
|
||||||
if self.with_guidance_embed:
|
if self.with_guidance_embed:
|
||||||
self.guidance = mx.broadcast_to(4.0, (1,))
|
|
||||||
self.guidance_embedder = GuidanceEmbedder()
|
self.guidance_embedder = GuidanceEmbedder()
|
||||||
self.timestep_embedder = TimestepEmbedder()
|
self.timestep_embedder = TimestepEmbedder()
|
||||||
|
|
||||||
def forward(self, time_step: mx.array, pooled_projection: mx.array) -> mx.array:
|
def forward(self, time_step: mx.array, pooled_projection: mx.array, guidance: mx.array) -> mx.array:
|
||||||
time_steps_proj = self._time_proj(time_step)
|
time_steps_proj = self._time_proj(time_step)
|
||||||
time_steps_emb = self.timestep_embedder.forward(time_steps_proj)
|
time_steps_emb = self.timestep_embedder.forward(time_steps_proj)
|
||||||
if self.with_guidance_embed:
|
if self.with_guidance_embed:
|
||||||
time_steps_emb += self.guidance_embedder.forward(self._time_proj(self.guidance))
|
time_steps_emb += self.guidance_embedder.forward(self._time_proj(guidance))
|
||||||
pooled_projections = self.text_embedder.forward(pooled_projection)
|
pooled_projections = self.text_embedder.forward(pooled_projection)
|
||||||
conditioning = time_steps_emb + pooled_projections
|
conditioning = time_steps_emb + pooled_projections
|
||||||
return conditioning.astype(Config.precision)
|
return conditioning.astype(Config.precision)
|
||||||
|
|||||||
@ -34,10 +34,11 @@ class Transformer(nn.Module):
|
|||||||
hidden_states: mx.array,
|
hidden_states: mx.array,
|
||||||
config: Config
|
config: Config
|
||||||
) -> mx.array:
|
) -> mx.array:
|
||||||
time_step = config.time_steps[t]
|
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,))
|
||||||
hidden_states = self.x_embedder(hidden_states)
|
hidden_states = self.x_embedder(hidden_states)
|
||||||
text_embeddings = self.time_text_embed.forward(time_step, pooled_prompt_embeds)
|
guidance = mx.broadcast_to(config.guidance, (1,))
|
||||||
|
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])
|
||||||
img_ids = Transformer._prepare_latent_image_ids(config.width, config.height)
|
img_ids = Transformer._prepare_latent_image_ids(config.width, config.height)
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user