@@ -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..." )
0 commit comments