fix architecture for dev
This commit is contained in:
parent
7b1eba235d
commit
14c5b2bd0e
2
main.py
2
main.py
@ -6,7 +6,7 @@ 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.flux import Flux1
|
||||
|
||||
flux = Flux1("black-forest-labs/FLUX.1-schnell")
|
||||
flux = Flux1("black-forest-labs/FLUX.1-dev")
|
||||
|
||||
image = flux.generate_image(
|
||||
seed=3,
|
||||
|
||||
@ -19,8 +19,10 @@ from flux_1_schnell.weights.weight_handler import WeightHandler
|
||||
class Flux1:
|
||||
|
||||
def __init__(self, repo_id: str):
|
||||
tokenizers = TokenizerHandler.load_from_disk_or_huggingface(repo_id)
|
||||
self.t5_tokenizer = TokenizerT5(tokenizers.t5)
|
||||
is_dev = "FLUX.1-dev" in repo_id
|
||||
max_t5_length = 512 if is_dev else 256
|
||||
tokenizers = TokenizerHandler.load_from_disk_or_huggingface(repo_id, max_t5_length)
|
||||
self.t5_tokenizer = TokenizerT5(tokenizers.t5, max_length=max_t5_length)
|
||||
self.clip_tokenizer = TokenizerCLIP(tokenizers.clip)
|
||||
|
||||
weights = WeightHandler.load_from_disk_or_huggingface(repo_id)
|
||||
|
||||
@ -19,7 +19,8 @@ class T5SelfAttention(nn.Module):
|
||||
key_states = T5SelfAttention.shape(self.k(hidden_states))
|
||||
value_states = T5SelfAttention.shape(self.v(hidden_states))
|
||||
scores = mx.matmul(query_states, mx.transpose(key_states, (0, 1, 3, 2)))
|
||||
position_bias = self._compute_bias()
|
||||
seq_length = hidden_states.shape[1]
|
||||
position_bias = self._compute_bias(seq_length=seq_length)
|
||||
scores += position_bias
|
||||
attn_weights = nn.softmax(scores, axis=-1)
|
||||
attn_output = T5SelfAttention.un_shape(mx.matmul(attn_weights, value_states))
|
||||
@ -34,9 +35,9 @@ class T5SelfAttention(nn.Module):
|
||||
def un_shape(states):
|
||||
return mx.reshape(mx.transpose(states, (0, 2, 1, 3)), (1, -1, 4096))
|
||||
|
||||
def _compute_bias(self):
|
||||
context_position = mx.arange(start=0, stop=256, step=1)[:, None]
|
||||
memory_position = mx.arange(start=0, stop=256, step=1)[None, :]
|
||||
def _compute_bias(self, seq_length):
|
||||
context_position = mx.arange(start=0, stop=seq_length, step=1)[:, None]
|
||||
memory_position = mx.arange(start=0, stop=seq_length, step=1)[None, :]
|
||||
relative_position = memory_position - context_position
|
||||
relative_position_bucket = T5SelfAttention._relative_position_bucket(relative_position)
|
||||
values = self.relative_attention_bias(relative_position_bucket)
|
||||
|
||||
16
src/flux_1_schnell/models/transformer/guidance_embedder.py
Normal file
16
src/flux_1_schnell/models/transformer/guidance_embedder.py
Normal file
@ -0,0 +1,16 @@
|
||||
from mlx import nn
|
||||
import mlx.core as mx
|
||||
|
||||
|
||||
class GuidanceEmbedder(nn.Module):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.linear_1 = nn.Linear(256, 3072)
|
||||
self.linear_2 = nn.Linear(3072, 3072)
|
||||
|
||||
def forward(self, sample: mx.array) -> mx.array:
|
||||
sample = self.linear_1(sample)
|
||||
sample = nn.silu(sample)
|
||||
sample = self.linear_2(sample)
|
||||
return sample
|
||||
@ -5,18 +5,25 @@ import mlx.core as mx
|
||||
from flux_1_schnell.config.config import Config
|
||||
from flux_1_schnell.models.transformer.text_embedder import TextEmbedder
|
||||
from flux_1_schnell.models.transformer.timestep_embedder import TimestepEmbedder
|
||||
from flux_1_schnell.models.transformer.guidance_embedder import GuidanceEmbedder
|
||||
|
||||
|
||||
class TimeTextEmbed(nn.Module):
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self, with_guidance_embed: bool = False):
|
||||
super().__init__()
|
||||
self.text_embedder = TextEmbedder()
|
||||
self.with_guidance_embed = with_guidance_embed
|
||||
if self.with_guidance_embed:
|
||||
self.guidance = mx.broadcast_to(4.0, (1,))
|
||||
self.guidance_embedder = GuidanceEmbedder()
|
||||
self.timestep_embedder = TimestepEmbedder()
|
||||
|
||||
def forward(self, time_step: mx.array, pooled_projection: mx.array) -> mx.array:
|
||||
time_steps_proj = TimeTextEmbed._time_proj(time_step)
|
||||
time_steps_proj = self._time_proj(time_step)
|
||||
time_steps_emb = self.timestep_embedder.forward(time_steps_proj)
|
||||
if self.with_guidance_embed:
|
||||
time_steps_emb += self.guidance_embedder.forward(self._time_proj(self.guidance))
|
||||
pooled_projections = self.text_embedder.forward(pooled_projection)
|
||||
conditioning = time_steps_emb + pooled_projections
|
||||
return conditioning.astype(Config.precision)
|
||||
|
||||
@ -15,7 +15,8 @@ class Transformer(nn.Module):
|
||||
super().__init__()
|
||||
self.pos_embed = EmbedND()
|
||||
self.x_embedder = nn.Linear(64, 3072)
|
||||
self.time_text_embed = TimeTextEmbed()
|
||||
with_guidance_embed = "guidance_embedder" in weights["time_text_embed"].keys()
|
||||
self.time_text_embed = TimeTextEmbed(with_guidance_embed = with_guidance_embed)
|
||||
self.context_embedder = nn.Linear(4096, 3072)
|
||||
self.transformer_blocks = [JointTransformerBlock(i) for i in range(19)]
|
||||
self.single_transformer_blocks = [SingleTransformerBlock(i) for i in range(38)]
|
||||
@ -38,7 +39,7 @@ class Transformer(nn.Module):
|
||||
hidden_states = self.x_embedder(hidden_states)
|
||||
text_embeddings = self.time_text_embed.forward(time_step, pooled_prompt_embeds)
|
||||
encoder_hidden_states = self.context_embedder(prompt_embeds)
|
||||
txt_ids = Transformer._prepare_text_ids()
|
||||
txt_ids = Transformer._prepare_text_ids(seq_len = prompt_embeds.shape[1])
|
||||
img_ids = Transformer._prepare_latent_image_ids(config.width, config.height)
|
||||
ids = mx.concatenate((txt_ids, img_ids), axis=1)
|
||||
image_rotary_emb = self.pos_embed.forward(ids)
|
||||
@ -78,5 +79,5 @@ class Transformer(nn.Module):
|
||||
return latent_image_ids
|
||||
|
||||
@staticmethod
|
||||
def _prepare_text_ids() -> mx.array:
|
||||
return mx.zeros((1, 256, 3))
|
||||
def _prepare_text_ids(seq_len) -> mx.array:
|
||||
return mx.zeros((1, seq_len, 3))
|
||||
|
||||
@ -3,16 +3,16 @@ from transformers import T5Tokenizer
|
||||
|
||||
|
||||
class TokenizerT5:
|
||||
MAX_TOKEN_LENGTH = 256
|
||||
|
||||
def __init__(self, tokenizer: T5Tokenizer):
|
||||
def __init__(self, tokenizer: T5Tokenizer, max_length: int = 256):
|
||||
self.tokenizer = tokenizer
|
||||
self.max_length = max_length
|
||||
|
||||
def tokenize(self, prompt: str) -> mx.array:
|
||||
return self.tokenizer(
|
||||
[prompt],
|
||||
padding="max_length",
|
||||
max_length=TokenizerT5.MAX_TOKEN_LENGTH,
|
||||
max_length=self.max_length,
|
||||
truncation=True,
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
|
||||
@ -9,7 +9,7 @@ from flux_1_schnell.tokenizer.t5_tokenizer import TokenizerT5
|
||||
|
||||
class TokenizerHandler:
|
||||
|
||||
def __init__(self, repo_id: str):
|
||||
def __init__(self, repo_id: str, max_t5_length: int = 256):
|
||||
root_path = TokenizerHandler._download_or_get_cached_tokenizers(repo_id)
|
||||
|
||||
self.clip = transformers.CLIPTokenizer.from_pretrained(
|
||||
@ -20,12 +20,12 @@ class TokenizerHandler:
|
||||
self.t5 = transformers.T5Tokenizer.from_pretrained(
|
||||
pretrained_model_name_or_path=root_path / "tokenizer_2",
|
||||
local_files_only=True,
|
||||
max_length=TokenizerT5.MAX_TOKEN_LENGTH
|
||||
max_length=max_t5_length
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def load_from_disk_or_huggingface(repo_id: str) -> "TokenizerHandler":
|
||||
return TokenizerHandler(repo_id)
|
||||
def load_from_disk_or_huggingface(repo_id: str, max_t5_length: int = 256) -> "TokenizerHandler":
|
||||
return TokenizerHandler(repo_id, max_t5_length)
|
||||
|
||||
@staticmethod
|
||||
def _download_or_get_cached_tokenizers(repo_id: str) -> Path:
|
||||
|
||||
Loading…
Reference in New Issue
Block a user