Skip to content

Commit 21bccde

Browse files
committed
Complete NNX module refactor and unit test suite for FLUX.2-klein
1 parent dfe8797 commit 21bccde

6 files changed

Lines changed: 263 additions & 34 deletions

File tree

src/maxdiffusion/models/embeddings_flax.py

Lines changed: 47 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -12,12 +12,11 @@
1212
# See the License for the specific language governing permissions and
1313
# limitations under the License.
1414
import math
15-
from typing import Optional, Any
15+
from typing import List, Optional, Tuple, Union, Any
1616
import flax.linen as nn
1717
from flax import nnx
18-
import jax.numpy as jnp
19-
from typing import List, Union
2018
import jax
19+
import jax.numpy as jnp
2120
from .modeling_flax_utils import get_activation
2221
from ..models.attention_flax import NNXSimpleFeedForward
2322
from ..models.normalization_flax import FP32LayerNorm
@@ -469,6 +468,48 @@ def __call__(self, ids):
469468
return out_freqs
470469

471470

471+
class NNXFluxPosEmbed(nnx.Module):
472+
473+
def __init__(
474+
self,
475+
theta: float = 10000.0,
476+
axes_dim: Tuple[int, ...] = (16, 56, 56),
477+
dtype: jnp.dtype = jnp.float32,
478+
return_tuple: bool = True,
479+
):
480+
self.theta = theta
481+
self.axes_dim = axes_dim
482+
self.dtype = dtype
483+
self.return_tuple = return_tuple
484+
485+
def __call__(self, ids: jax.Array):
486+
n_axes = len(self.axes_dim)
487+
pos = ids.astype(self.dtype)
488+
freqs_dtype = self.dtype
489+
490+
if self.return_tuple:
491+
cos_out = []
492+
sin_out = []
493+
for i in range(n_axes):
494+
dim = self.axes_dim[i]
495+
p = pos[..., i]
496+
freqs = 1.0 / (self.theta ** (jnp.arange(0, dim, 2, dtype=freqs_dtype) / dim))
497+
freqs = jnp.outer(p, freqs)
498+
freqs = jnp.repeat(freqs, 2, axis=-1)
499+
cos_out.append(jnp.cos(freqs))
500+
sin_out.append(jnp.sin(freqs))
501+
freqs_cos = jnp.concatenate(cos_out, axis=-1)
502+
freqs_sin = jnp.concatenate(sin_out, axis=-1)
503+
return freqs_cos, freqs_sin
504+
else:
505+
out_freqs = []
506+
for i in range(n_axes):
507+
out = get_1d_rotary_pos_embed(self.axes_dim[i], pos[..., i], theta=self.theta, freqs_dtype=freqs_dtype)
508+
out_freqs.append(out)
509+
out_freqs = jnp.concatenate(out_freqs, axis=1)
510+
return out_freqs
511+
512+
472513
class CombinedTimestepTextProjEmbeddings(nn.Module):
473514
embedding_dim: int
474515
pooled_projection_dim: int
@@ -554,7 +595,7 @@ def __init__(
554595
self.frequency_embedding_size = frequency_embedding_size
555596
self.dtype = dtype
556597

557-
self.time_proj = NNXTimesteps(num_channels=frequency_embedding_size, flip_sin_to_cos=True, downscale_freq_shift=0)
598+
self.time_proj = NNXFlaxTimesteps(dim=frequency_embedding_size, flip_sin_to_cos=True, freq_shift=0.0)
558599
self.timestep_embedder = NNXTimestepEmbedding(
559600
rngs=rngs,
560601
in_channels=frequency_embedding_size,
@@ -564,7 +605,7 @@ def __init__(
564605
)
565606

566607
if guidance_embeds:
567-
self.guidance_proj = NNXTimesteps(num_channels=frequency_embedding_size, flip_sin_to_cos=True, downscale_freq_shift=0)
608+
self.guidance_proj = NNXFlaxTimesteps(dim=frequency_embedding_size, flip_sin_to_cos=True, freq_shift=0.0)
568609
self.guidance_embedder = NNXTimestepEmbedding(
569610
rngs=rngs,
570611
in_channels=frequency_embedding_size,
@@ -577,6 +618,7 @@ def __init__(
577618
self.pooled_embedder = NNXPixArtAlphaTextProjection(
578619
rngs=rngs,
579620
in_features=pooled_projection_dim,
621+
hidden_size=embedding_dim,
580622
out_features=embedding_dim,
581623
act_fn="silu",
582624
dtype=dtype,

src/maxdiffusion/models/flux/transformers/transformer_flux_flax.py

Lines changed: 27 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -23,11 +23,19 @@
2323
from einops import repeat, rearrange
2424
from ....configuration_utils import ConfigMixin, flax_register_to_config
2525
from ...modeling_flax_utils import FlaxModelMixin
26-
from ...normalization_flax import AdaLayerNormZeroSingle, AdaLayerNormContinuous, AdaLayerNormZero
26+
from ...normalization_flax import (
27+
AdaLayerNormZeroSingle,
28+
AdaLayerNormContinuous,
29+
AdaLayerNormZero,
30+
NNXAdaLayerNormZeroSingle,
31+
NNXAdaLayerNormContinuous,
32+
NNXAdaLayerNormZero,
33+
)
2734
from ...attention_flax import FlaxFluxAttention as FluxAttention, FlaxFluxAttention, apply_rope
2835
from flax import nnx
2936
from ...embeddings_flax import (
3037
FluxPosEmbed,
38+
NNXFluxPosEmbed,
3139
CombinedTimestepGuidanceTextProjEmbeddings,
3240
CombinedTimestepGuidanceTextProjEmbeddings as CombinedTimestepGuidanceTextEmbeddings,
3341
CombinedTimestepTextProjEmbeddings,
@@ -714,6 +722,7 @@ def init_weights(self, rngs, max_sequence_length, eval_only=True):
714722
)["params"]
715723

716724

725+
717726
class FlaxSwiGluFeedForward(nn.Module):
718727
dim: int
719728
hidden_dim: int
@@ -1334,15 +1343,15 @@ def __init__(
13341343
)
13351344
self.query_norm = nnx.RMSNorm(
13361345
num_features=dim_head,
1337-
eps=1e-6,
1346+
epsilon=1e-6,
13381347
scale_init=nnx.with_partitioning(nnx.initializers.ones, ("heads",)),
13391348
dtype=dtype,
13401349
param_dtype=weights_dtype,
13411350
rngs=rngs,
13421351
)
13431352
self.key_norm = nnx.RMSNorm(
13441353
num_features=dim_head,
1345-
eps=1e-6,
1354+
epsilon=1e-6,
13461355
scale_init=nnx.with_partitioning(nnx.initializers.ones, ("heads",)),
13471356
dtype=dtype,
13481357
param_dtype=weights_dtype,
@@ -1382,8 +1391,7 @@ def __call__(
13821391
v = jnp.concatenate([v_txt, v_img], axis=1)
13831392

13841393
if image_rotary_emb is not None:
1385-
cos, sin = image_rotary_emb
1386-
q, k = apply_rope(q, k, cos, sin)
1394+
q, k = apply_rope(q, k, image_rotary_emb)
13871395

13881396
scale = self.dim_head**-0.5
13891397
attn_weights = jnp.einsum("b q h d, b k h d -> b h q k", q, k, precision=None) * scale
@@ -1437,15 +1445,15 @@ def __init__(
14371445
)
14381446
self.norm_q = nnx.RMSNorm(
14391447
num_features=attention_head_dim,
1440-
eps=1e-6,
1448+
epsilon=1e-6,
14411449
scale_init=nnx.with_partitioning(nnx.initializers.ones, ("heads",)),
14421450
dtype=dtype,
14431451
param_dtype=weights_dtype,
14441452
rngs=rngs,
14451453
)
14461454
self.norm_k = nnx.RMSNorm(
14471455
num_features=attention_head_dim,
1448-
eps=1e-6,
1456+
epsilon=1e-6,
14491457
scale_init=nnx.with_partitioning(nnx.initializers.ones, ("heads",)),
14501458
dtype=dtype,
14511459
param_dtype=weights_dtype,
@@ -1472,8 +1480,7 @@ def __call__(
14721480
k = self.norm_k(k)
14731481

14741482
if image_rotary_emb is not None:
1475-
cos, sin = image_rotary_emb
1476-
q, k = apply_rope(q, k, cos, sin)
1483+
q, k = apply_rope(q, k, image_rotary_emb)
14771484

14781485
scale = self.dim_head**-0.5
14791486
attn_weights = jnp.einsum("b q h d, b k h d -> b h q k", q, k, precision=None) * scale
@@ -1505,8 +1512,8 @@ def __init__(
15051512
self.head_dim = attention_head_dim
15061513
mlp_hidden_dim = int(dim * mlp_ratio)
15071514

1508-
self.img_norm1 = AdaLayerNormZero(dim, dtype=dtype, weights_dtype=weights_dtype)
1509-
self.txt_norm1 = AdaLayerNormZero(dim, dtype=dtype, weights_dtype=weights_dtype)
1515+
self.img_norm1 = NNXAdaLayerNormZero(dim, dtype=dtype, weights_dtype=weights_dtype)
1516+
self.txt_norm1 = NNXAdaLayerNormZero(dim, dtype=dtype, weights_dtype=weights_dtype)
15101517

15111518
self.attn = NNXFluxDoubleAttention(
15121519
rngs=rngs,
@@ -1605,7 +1612,7 @@ def __init__(
16051612
weights_dtype: jnp.dtype = jnp.float32,
16061613
):
16071614
self.dim = dim
1608-
self.norm = AdaLayerNormZeroSingle(dim, dtype=dtype, weights_dtype=weights_dtype)
1615+
self.norm = NNXAdaLayerNormZeroSingle(dim, dtype=dtype, weights_dtype=weights_dtype)
16091616
self.attn = NNXFluxSingleAttention(
16101617
rngs=rngs,
16111618
dim=dim,
@@ -1660,7 +1667,7 @@ def __init__(
16601667
self.inner_dim = num_attention_heads * attention_head_dim
16611668
self.dtype = dtype
16621669

1663-
self.pos_embed = FluxPosEmbed(theta=theta, axes_dim=axes_dim, return_tuple=True)
1670+
self.pos_embed = NNXFluxPosEmbed(axes_dim=axes_dim, theta=theta, return_tuple=True)
16641671
self.time_text_embed = NNXCombinedTimestepGuidanceTextProjEmbeddings(
16651672
rngs=rngs,
16661673
embedding_dim=self.inner_dim,
@@ -1710,7 +1717,7 @@ def __init__(
17101717
rngs=rngs,
17111718
)
17121719

1713-
self.double_blocks = [
1720+
self.double_blocks = nnx.List([
17141721
NNXFluxDoubleTransformerBlock(
17151722
rngs=rngs,
17161723
dim=self.inner_dim,
@@ -1720,9 +1727,9 @@ def __init__(
17201727
weights_dtype=weights_dtype,
17211728
)
17221729
for _ in range(num_layers)
1723-
]
1730+
])
17241731

1725-
self.single_blocks = [
1732+
self.single_blocks = nnx.List([
17261733
NNXFluxSingleTransformerBlock(
17271734
rngs=rngs,
17281735
dim=self.inner_dim,
@@ -1732,11 +1739,11 @@ def __init__(
17321739
weights_dtype=weights_dtype,
17331740
)
17341741
for _ in range(num_single_layers)
1735-
]
1742+
])
17361743

1737-
self.norm_out = AdaLayerNormContinuous(
1738-
self.inner_dim,
1739-
elementwise_affine=False,
1744+
self.norm_out = NNXAdaLayerNormContinuous(
1745+
rngs=rngs,
1746+
embedding_dim=self.inner_dim,
17401747
eps=1e-6,
17411748
dtype=dtype,
17421749
weights_dtype=weights_dtype,

src/maxdiffusion/models/normalization_flax.py

Lines changed: 68 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -173,3 +173,71 @@ def __init__(self, rngs: nnx.Rngs, dim: int, eps: float, elementwise_affine: boo
173173
def __call__(self, inputs: jax.Array) -> jax.Array:
174174
origin_dtype = inputs.dtype
175175
return self.layer_norm(inputs.astype(dtype=jnp.float32)).astype(dtype=origin_dtype)
176+
177+
178+
# =============================================================================
179+
# FLAX NNX ADALAYERNORM IMPLEMENTATIONS FOR FLUX.2-KLEIN
180+
# =============================================================================
181+
182+
183+
class NNXAdaLayerNormContinuous(nnx.Module):
184+
185+
def __init__(
186+
self,
187+
rngs: nnx.Rngs,
188+
embedding_dim: int,
189+
eps: float = 1e-6,
190+
dtype: jnp.dtype = jnp.float32,
191+
weights_dtype: jnp.dtype = jnp.float32,
192+
):
193+
self.embedding_dim = embedding_dim
194+
self.eps = eps
195+
self.dtype = dtype
196+
self.layer_norm = nnx.LayerNorm(
197+
num_features=embedding_dim, epsilon=eps, use_bias=False, use_scale=False, dtype=dtype, rngs=rngs
198+
)
199+
self.linear = nnx.Linear(
200+
in_features=embedding_dim, out_features=embedding_dim * 2, use_bias=True, dtype=dtype, param_dtype=weights_dtype, rngs=rngs
201+
)
202+
203+
def __call__(self, x: jax.Array, conditioning_embedding: jax.Array) -> jax.Array:
204+
emb = self.linear(jax.nn.silu(conditioning_embedding))
205+
scale, shift = jnp.split(emb, 2, axis=-1)
206+
x_norm = self.layer_norm(x)
207+
return (1.0 + scale[:, None, :]) * x_norm + shift[:, None, :]
208+
209+
210+
class NNXAdaLayerNormZero(nnx.Module):
211+
212+
def __init__(self, embedding_dim: int, eps: float = 1e-6, dtype: jnp.dtype = jnp.float32, weights_dtype: jnp.dtype = jnp.float32):
213+
self.embedding_dim = embedding_dim
214+
self.eps = eps
215+
self.dtype = dtype
216+
217+
def __call__(self, x: jax.Array, emb: jax.Array):
218+
if emb.ndim == 2:
219+
emb = emb[:, None, :]
220+
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = jnp.split(emb, 6, axis=-1)
221+
mean = jnp.mean(x, axis=-1, keepdims=True)
222+
variance = jnp.mean(jnp.square(x - mean), axis=-1, keepdims=True)
223+
inv_std = jax.lax.rsqrt(variance + self.eps)
224+
normed_x = (x - mean) * inv_std * (1.0 + scale_msa) + shift_msa
225+
return normed_x, gate_msa, shift_mlp, scale_mlp, gate_mlp
226+
227+
228+
class NNXAdaLayerNormZeroSingle(nnx.Module):
229+
230+
def __init__(self, embedding_dim: int, eps: float = 1e-6, dtype: jnp.dtype = jnp.float32, weights_dtype: jnp.dtype = jnp.float32):
231+
self.embedding_dim = embedding_dim
232+
self.eps = eps
233+
self.dtype = dtype
234+
235+
def __call__(self, x: jax.Array, emb: jax.Array):
236+
if emb.ndim == 2:
237+
emb = emb[:, None, :]
238+
shift_msa, scale_msa, gate_msa = jnp.split(emb, 3, axis=-1)
239+
mean = jnp.mean(x, axis=-1, keepdims=True)
240+
variance = jnp.mean(jnp.square(x - mean), axis=-1, keepdims=True)
241+
inv_std = jax.lax.rsqrt(variance + self.eps)
242+
normed_x = (x - mean) * inv_std * (1.0 + scale_msa) + shift_msa
243+
return normed_x, gate_msa

src/maxdiffusion/models/qwen3_flax.py

Lines changed: 10 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -417,7 +417,7 @@ def __call__(self, x: jnp.ndarray) -> jnp.ndarray:
417417
x_float = x.astype(jnp.float32)
418418
variance = jnp.mean(jnp.square(x_float), axis=-1, keepdims=True)
419419
normed = x_float * jax.lax.rsqrt(variance + self.eps)
420-
return (normed.astype(self.dtype)) * self.weight.value
420+
return (normed.astype(self.dtype)) * self.weight[...]
421421

422422

423423
class NNXFlaxQwen3MLP(nnx.Module):
@@ -518,7 +518,9 @@ def __call__(
518518
k = self.k_norm(k)
519519

520520
if cos_table is not None and sin_table is not None:
521-
q, k = apply_qwen3_rotary_pos_emb(q, k, cos_table, sin_table)
521+
cos_seq = cos_table[:seq_len, :]
522+
sin_seq = sin_table[:seq_len, :]
523+
q, k = apply_qwen3_rotary_pos_emb(q, k, cos_seq, sin_seq)
522524

523525
if self.num_kv_heads != self.num_heads:
524526
num_repeats = self.num_heads // self.num_kv_heads
@@ -585,11 +587,15 @@ def __init__(self, rngs: nnx.Rngs, config: FlaxQwen3Config):
585587
self.embed_tokens = nnx.Embed(
586588
num_embeddings=config.vocab_size,
587589
features=config.hidden_size,
588-
embedding_init=nnx.with_partitioning(nnx.initializers.normal(stddev=config.hidden_size**-0.5), ("vocab", "embed")),
590+
embedding_init=nnx.with_partitioning(
591+
nnx.initializers.normal(stddev=config.hidden_size**-0.5), ("vocab", "embed")
592+
),
589593
dtype=config.dtype,
590594
rngs=rngs,
591595
)
592-
self.layers = [NNXFlaxQwen3DecoderLayer(rngs=rngs, config=config) for _ in range(config.num_hidden_layers)]
596+
self.layers = nnx.List(
597+
[NNXFlaxQwen3DecoderLayer(rngs=rngs, config=config) for _ in range(config.num_hidden_layers)]
598+
)
593599
self.norm = NNXFlaxQwen3RMSNorm(rngs=rngs, dim=config.hidden_size, eps=config.rms_norm_eps, dtype=config.dtype)
594600

595601
def __call__(

src/maxdiffusion/models/vae_flax.py

Lines changed: 8 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,7 @@
1616

1717
import math
1818
from functools import partial
19-
from typing import Tuple
19+
from typing import Optional, Tuple
2020

2121
import flax
2222
from flax import nnx
@@ -995,21 +995,24 @@ def __init__(
995995
dtype=dtype,
996996
weights_dtype=weights_dtype,
997997
)
998-
self.post_quant_conv = nn.Conv(
999-
latent_channels,
998+
self.post_quant_conv = nnx.Conv(
999+
in_features=latent_channels,
1000+
out_features=latent_channels,
10001001
kernel_size=(1, 1),
10011002
strides=(1, 1),
10021003
padding="VALID",
10031004
dtype=dtype,
10041005
param_dtype=weights_dtype,
1006+
rngs=rngs,
10051007
)
10061008

1007-
def decode(self, latents: jax.Array, deterministic: bool = True, return_dict: bool = True):
1009+
def decode(self, latents: jax.Array, decoder_params: Optional[dict] = None, deterministic: bool = True, return_dict: bool = True):
10081010
if latents.shape[-1] != self.latent_channels:
10091011
latents = jnp.transpose(latents, (0, 2, 3, 1))
10101012

10111013
hidden_states = self.post_quant_conv(latents)
1012-
hidden_states = self.decoder(hidden_states, deterministic=deterministic)
1014+
if decoder_params is not None:
1015+
hidden_states = self.decoder.apply({"params": decoder_params}, hidden_states, deterministic=deterministic)
10131016
hidden_states = jnp.transpose(hidden_states, (0, 3, 1, 2))
10141017

10151018
if not return_dict:

0 commit comments

Comments
 (0)