Skip to content

Commit 8b2f1c7

Browse files
committed
Apply pyink 23.10.0 formatting to 5 FLUX.2-klein source files
1 parent 3fecc8f commit 8b2f1c7

5 files changed

Lines changed: 51 additions & 40 deletions

File tree

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

Lines changed: 27 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -722,7 +722,6 @@ def init_weights(self, rngs, max_sequence_length, eval_only=True):
722722
)["params"]
723723

724724

725-
726725
class FlaxSwiGluFeedForward(nn.Module):
727726
dim: int
728727
hidden_dim: int
@@ -1717,29 +1716,33 @@ def __init__(
17171716
rngs=rngs,
17181717
)
17191718

1720-
self.double_blocks = nnx.List([
1721-
NNXFluxDoubleTransformerBlock(
1722-
rngs=rngs,
1723-
dim=self.inner_dim,
1724-
num_attention_heads=num_attention_heads,
1725-
attention_head_dim=attention_head_dim,
1726-
dtype=dtype,
1727-
weights_dtype=weights_dtype,
1728-
)
1729-
for _ in range(num_layers)
1730-
])
1731-
1732-
self.single_blocks = nnx.List([
1733-
NNXFluxSingleTransformerBlock(
1734-
rngs=rngs,
1735-
dim=self.inner_dim,
1736-
num_attention_heads=num_attention_heads,
1737-
attention_head_dim=attention_head_dim,
1738-
dtype=dtype,
1739-
weights_dtype=weights_dtype,
1740-
)
1741-
for _ in range(num_single_layers)
1742-
])
1719+
self.double_blocks = nnx.List(
1720+
[
1721+
NNXFluxDoubleTransformerBlock(
1722+
rngs=rngs,
1723+
dim=self.inner_dim,
1724+
num_attention_heads=num_attention_heads,
1725+
attention_head_dim=attention_head_dim,
1726+
dtype=dtype,
1727+
weights_dtype=weights_dtype,
1728+
)
1729+
for _ in range(num_layers)
1730+
]
1731+
)
1732+
1733+
self.single_blocks = nnx.List(
1734+
[
1735+
NNXFluxSingleTransformerBlock(
1736+
rngs=rngs,
1737+
dim=self.inner_dim,
1738+
num_attention_heads=num_attention_heads,
1739+
attention_head_dim=attention_head_dim,
1740+
dtype=dtype,
1741+
weights_dtype=weights_dtype,
1742+
)
1743+
for _ in range(num_single_layers)
1744+
]
1745+
)
17431746

17441747
self.norm_out = NNXAdaLayerNormContinuous(
17451748
rngs=rngs,

src/maxdiffusion/models/normalization_flax.py

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -197,7 +197,12 @@ def __init__(
197197
num_features=embedding_dim, epsilon=eps, use_bias=False, use_scale=False, dtype=dtype, rngs=rngs
198198
)
199199
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
200+
in_features=embedding_dim,
201+
out_features=embedding_dim * 2,
202+
use_bias=True,
203+
dtype=dtype,
204+
param_dtype=weights_dtype,
205+
rngs=rngs,
201206
)
202207

203208
def __call__(self, x: jax.Array, conditioning_embedding: jax.Array) -> jax.Array:
@@ -209,7 +214,9 @@ def __call__(self, x: jax.Array, conditioning_embedding: jax.Array) -> jax.Array
209214

210215
class NNXAdaLayerNormZero(nnx.Module):
211216

212-
def __init__(self, embedding_dim: int, eps: float = 1e-6, dtype: jnp.dtype = jnp.float32, weights_dtype: jnp.dtype = jnp.float32):
217+
def __init__(
218+
self, embedding_dim: int, eps: float = 1e-6, dtype: jnp.dtype = jnp.float32, weights_dtype: jnp.dtype = jnp.float32
219+
):
213220
self.embedding_dim = embedding_dim
214221
self.eps = eps
215222
self.dtype = dtype
@@ -227,7 +234,9 @@ def __call__(self, x: jax.Array, emb: jax.Array):
227234

228235
class NNXAdaLayerNormZeroSingle(nnx.Module):
229236

230-
def __init__(self, embedding_dim: int, eps: float = 1e-6, dtype: jnp.dtype = jnp.float32, weights_dtype: jnp.dtype = jnp.float32):
237+
def __init__(
238+
self, embedding_dim: int, eps: float = 1e-6, dtype: jnp.dtype = jnp.float32, weights_dtype: jnp.dtype = jnp.float32
239+
):
231240
self.embedding_dim = embedding_dim
232241
self.eps = eps
233242
self.dtype = dtype

src/maxdiffusion/models/qwen3_flax.py

Lines changed: 2 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -587,15 +587,11 @@ def __init__(self, rngs: nnx.Rngs, config: FlaxQwen3Config):
587587
self.embed_tokens = nnx.Embed(
588588
num_embeddings=config.vocab_size,
589589
features=config.hidden_size,
590-
embedding_init=nnx.with_partitioning(
591-
nnx.initializers.normal(stddev=config.hidden_size**-0.5), ("vocab", "embed")
592-
),
590+
embedding_init=nnx.with_partitioning(nnx.initializers.normal(stddev=config.hidden_size**-0.5), ("vocab", "embed")),
593591
dtype=config.dtype,
594592
rngs=rngs,
595593
)
596-
self.layers = nnx.List(
597-
[NNXFlaxQwen3DecoderLayer(rngs=rngs, config=config) for _ in range(config.num_hidden_layers)]
598-
)
594+
self.layers = nnx.List([NNXFlaxQwen3DecoderLayer(rngs=rngs, config=config) for _ in range(config.num_hidden_layers)])
599595
self.norm = NNXFlaxQwen3RMSNorm(rngs=rngs, dim=config.hidden_size, eps=config.rms_norm_eps, dtype=config.dtype)
600596

601597
def __call__(

src/maxdiffusion/models/vae_flax.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1006,7 +1006,9 @@ def __init__(
10061006
rngs=rngs,
10071007
)
10081008

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

src/maxdiffusion/pipelines/flux/flux2klein_pipeline.py

Lines changed: 7 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -215,10 +215,7 @@ def __call__(
215215
tokenizer = Qwen2TokenizerFast.from_pretrained(tokenizer_path, subfolder="tokenizer", local_files_only=True)
216216

217217
# Tokenize using deterministic explicit template string (version-agnostic across transformers versions)
218-
templated_texts = [
219-
f"<|im_start|>user\n{p}<|im_end|>\n<|im_start|>assistant\n<think>\n\n</think>\n\n"
220-
for p in prompts
221-
]
218+
templated_texts = [f"<|im_start|>user\n{p}<|im_end|>\n<|im_start|>assistant\n<think>\n\n</think>\n\n" for p in prompts]
222219
max_logging.log(f"DEBUG: templated_texts[0] = {repr(templated_texts[0])}")
223220
inputs = tokenizer(templated_texts, return_tensors="np", padding="max_length", truncation=True, max_length=seq_len_txt)
224221
prompt_ids = jnp.array(inputs["input_ids"])
@@ -246,7 +243,9 @@ def __call__(
246243
# ---------------------------------------------------------------------
247244
# PHASE B: Denoising Loop (Flux Transformer)
248245
# ---------------------------------------------------------------------
249-
max_logging.log(f"[PHASE B] Running {num_inference_steps}-step E2E Denoising Loop on a batch of {batch_size} images...")
246+
max_logging.log(
247+
f"[PHASE B] Running {num_inference_steps}-step E2E Denoising Loop on a batch of {batch_size} images..."
248+
)
250249
t0 = time.perf_counter()
251250

252251
guidance_vec_val = None
@@ -271,7 +270,9 @@ def __call__(
271270

272271
# Print progress
273272
sigma_val = scheduler_state.sigmas[step_idx]
274-
max_logging.log(f" -> Step {step_idx}: Timestep = {scheduler_state.timesteps[step_idx]:.4f}, Sigma = {sigma_val:.4f}")
273+
max_logging.log(
274+
f" -> Step {step_idx}: Timestep = {scheduler_state.timesteps[step_idx]:.4f}, Sigma = {sigma_val:.4f}"
275+
)
275276

276277
latents_jax.block_until_ready()
277278

0 commit comments

Comments
 (0)