Skip to content

Commit abfe781

Browse files
committed
feat(flux2klein): eliminate CPU parameter recasting during single-pass safetensors conversion
1 parent f88d1f3 commit abfe781

3 files changed

Lines changed: 22 additions & 25 deletions

File tree

src/maxdiffusion/generate_flux2klein.py

Lines changed: 6 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -360,20 +360,12 @@ def unbox_fn(x):
360360
max_logging.log(f" -> [SUB-TIMING 1/3] PyTree unboxing template setup: {time.time() - t_sub0:.2f}s")
361361
t_sub1 = time.time()
362362

363-
params = load_and_convert_flux_klein_weights(safetensors_path, params, num_double_layers, depth)
364-
vae_params, vae_bn_mean, vae_bn_std = load_and_convert_vae_weights(vae_safetensors_path, vae_params)
365-
qwen3_params = load_and_convert_qwen3_weights(text_encoder_path, qwen3_params, qwen3_config)
366-
max_logging.log(f" -> [SUB-TIMING 2/3] Safetensors loading & key mapping: {time.time() - t_sub1:.2f}s")
367-
368-
t_sub2 = time.time()
369-
if config.weights_dtype == "bfloat16":
370-
max_logging.log("Casting JAX parameters to bfloat16 in-place...")
371-
cast_dict_to_bfloat16_inplace(params, exclude_keywords=("norm",))
372-
cast_dict_to_bfloat16_inplace(vae_params, exclude_keywords=("norm",))
373-
cast_dict_to_bfloat16_inplace(qwen3_params, exclude_keywords=("norm",))
374-
vae_bn_mean = vae_bn_mean.astype(jnp.bfloat16)
375-
vae_bn_std = vae_bn_std.astype(jnp.bfloat16)
376-
max_logging.log(f" -> [SUB-TIMING 2b/3] bfloat16 in-place casting: {time.time() - t_sub2:.2f}s")
363+
weight_dtype = jnp.bfloat16 if config.weights_dtype == "bfloat16" else jnp.float32
364+
365+
params = load_and_convert_flux_klein_weights(safetensors_path, params, num_double_layers, depth, dtype=weight_dtype)
366+
vae_params, vae_bn_mean, vae_bn_std = load_and_convert_vae_weights(vae_safetensors_path, vae_params, dtype=weight_dtype)
367+
qwen3_params = load_and_convert_qwen3_weights(text_encoder_path, qwen3_params, qwen3_config, dtype=weight_dtype)
368+
max_logging.log(f" -> [SUB-TIMING 2/3] Safetensors loading & key mapping (in target dtype): {time.time() - t_sub1:.2f}s")
377369

378370
params = flax.core.freeze(params)
379371
vae_params = flax.core.freeze(vae_params)

src/maxdiffusion/models/flux/util.py

Lines changed: 11 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -411,7 +411,7 @@ def cast_dict_to_bfloat16_inplace(d, device=None, exclude_keywords=None, parent_
411411
# -----------------------------------------------------------------------------
412412

413413

414-
def load_and_convert_flux_klein_weights(safetensors_path, params, num_double_layers, num_single_layers):
414+
def load_and_convert_flux_klein_weights(safetensors_path, params, num_double_layers, num_single_layers, dtype=None):
415415
"""
416416
Loads weights from safetensors via zero-copy safetensors.numpy and converts them to JAX parameter dictionary.
417417
Supports dynamic layer counts (double and single stream blocks) and sharded safetensors directories.
@@ -439,12 +439,13 @@ def load_and_convert_flux_klein_weights(safetensors_path, params, num_double_lay
439439
expected_pytree = jax.tree_util.tree_map(lambda leaf: leaf, params)
440440

441441
first_leaf = jax.tree_util.tree_leaves(params)[0]
442-
target_dtype = first_leaf.dtype
442+
target_dtype = dtype if dtype is not None else first_leaf.dtype
443443

444-
def convert_and_transpose_tensor(tensor, transpose=False):
444+
def convert_and_transpose_tensor(tensor, transpose=False, is_norm=False):
445445
if transpose and len(tensor.shape) == 2:
446446
tensor = tensor.T
447-
return jnp.array(tensor, dtype=target_dtype)
447+
leaf_dtype = jnp.float32 if is_norm else target_dtype
448+
return jnp.array(tensor, dtype=leaf_dtype)
448449

449450
# Global layers
450451
params["context_embedder"]["kernel"] = convert_and_transpose_tensor(
@@ -563,7 +564,7 @@ def convert_and_transpose_tensor(tensor, transpose=False):
563564
return params
564565

565566

566-
def load_and_convert_vae_weights(safetensors_path, jax_params):
567+
def load_and_convert_vae_weights(safetensors_path, jax_params, dtype=None):
567568
"""Loads VAE weights from safetensors via zero-copy safetensors.numpy, maps them to JAX, and extracts BN stats."""
568569
from safetensors.numpy import load_file
569570
import flax
@@ -576,11 +577,13 @@ def load_and_convert_vae_weights(safetensors_path, jax_params):
576577
jax_params = flax.core.unfreeze(jax_params)
577578

578579
first_leaf = jax.tree_util.tree_leaves(jax_params)[0]
579-
target_dtype = first_leaf.dtype
580+
target_dtype = dtype if dtype is not None else first_leaf.dtype
580581

581-
def get_pytorch_weight_tensor(key, dtype=target_dtype):
582+
def get_pytorch_weight_tensor(key, dtype_val=target_dtype):
582583
tensor = pt_state_dict[key]
583-
return jnp.array(tensor, dtype=dtype)
584+
is_norm = any(kw in key.lower() for kw in ("norm", "layernorm", "rmsnorm", "groupnorm"))
585+
leaf_dtype = jnp.float32 if is_norm else dtype_val
586+
return jnp.array(tensor, dtype=leaf_dtype)
584587

585588
# Map weights
586589
max_logging.log("Mapping VAE decoder weights to JAX parameters...")

src/maxdiffusion/models/qwen3_flax.py

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -626,7 +626,7 @@ def __call__(
626626
# -----------------------------------------------------------------------------
627627

628628

629-
def load_and_convert_qwen3_weights(safetensors_path: str, jax_params: dict, config: FlaxQwen3Config) -> dict:
629+
def load_and_convert_qwen3_weights(safetensors_path: str, jax_params: dict, config: FlaxQwen3Config, dtype=None) -> dict:
630630
"""
631631
Loads weights from safetensors via zero-copy safetensors.numpy and converts them to JAX parameter dictionary.
632632
"""
@@ -649,7 +649,7 @@ def load_and_convert_qwen3_weights(safetensors_path: str, jax_params: dict, conf
649649
max_logging.log("Safetensors weights loaded successfully. Starting JAX parameter mapping...")
650650

651651
first_leaf = jax.tree_util.tree_leaves(jax_params)[0]
652-
target_dtype = first_leaf.dtype
652+
target_dtype = dtype if dtype is not None else first_leaf.dtype
653653

654654
# Helper to transpose and cast weight directly to target_dtype
655655
def get_w(name: str, transpose: bool = True):
@@ -659,7 +659,9 @@ def get_w(name: str, transpose: bool = True):
659659
t = torch_weights[name]
660660
if len(t.shape) == 2 and transpose:
661661
t = t.T
662-
return jnp.array(t, dtype=target_dtype)
662+
is_norm = any(kw in name.lower() for kw in ("norm", "layernorm", "rmsnorm", "groupnorm"))
663+
leaf_dtype = jnp.float32 if is_norm else target_dtype
664+
return jnp.array(t, dtype=leaf_dtype)
663665

664666
# Create mutable copy of JAX params to populate
665667
import flax

0 commit comments

Comments
 (0)