@@ -118,21 +118,28 @@ def validate_flax_state_dict(expected_pytree: dict, new_pytree: dict):
118118 expected_pytree: dict - a pytree that comes from initializing the model.
119119 new_pytree: dict - a pytree that has been created from pytorch weights.
120120 """
121- expected_pytree = flatten_dict (expected_pytree )
122- if len (expected_pytree .keys ()) != len (new_pytree .keys ()):
123- set1 = set (expected_pytree .keys ())
124- set2 = set (new_pytree .keys ())
121+ expected_flat = (
122+ flatten_dict (expected_pytree ) if not isinstance (next (iter (expected_pytree .keys ()), None ), tuple ) else expected_pytree
123+ )
124+ new_flat = flatten_dict (new_pytree ) if not isinstance (next (iter (new_pytree .keys ()), None ), tuple ) else new_pytree
125+
126+ if len (expected_flat .keys ()) != len (new_flat .keys ()):
127+ set1 = set (expected_flat .keys ())
128+ set2 = set (new_flat .keys ())
125129 missing_keys = set1 ^ set2
126- max_logging .log (f"missing keys : { missing_keys } " )
127- for key in expected_pytree .keys ():
128- if key in new_pytree .keys ():
130+ max_logging .log (
131+ f"Missing or extra parameter keys count mismatch ({ len (expected_flat )} expected vs { len (new_flat )} converted): { missing_keys } "
132+ )
133+
134+ for key in expected_flat .keys ():
135+ if key in new_flat .keys ():
129136 try :
130- expected_pytree_shape = expected_pytree [key ].shape
137+ expected_pytree_shape = expected_flat [key ].shape
131138 except Exception :
132- expected_pytree_shape = expected_pytree [key ]. value .shape
133- if expected_pytree_shape != new_pytree [key ].shape :
139+ expected_pytree_shape = getattr ( expected_flat [key ], " value" , expected_flat [ key ]) .shape
140+ if expected_pytree_shape != new_flat [key ].shape :
134141 max_logging .log (
135- f"shape mismatch, expected shape of { expected_pytree [ key ]. shape } , but got shape of { new_pytree [key ].shape } "
142+ f"shape mismatch for key ' { key } ': expected shape of { expected_pytree_shape } , but got shape of { new_flat [key ].shape } "
136143 )
137144 else :
138145 max_logging .log (f"key: { key } not found..." )
@@ -266,6 +273,7 @@ def load_flow_model(name: str, eval_shapes: dict, device: str, hf_download: bool
266273 jax .clear_caches ()
267274 return flax_state_dict
268275
276+
269277# -----------------------------------------------------------------------------
270278# Latent Packing & Unpacking Helpers
271279# -----------------------------------------------------------------------------
@@ -431,6 +439,8 @@ def load_and_convert_flux_klein_weights(safetensors_path, params, num_double_lay
431439
432440 max_logging .log ("Mapping weights to JAX parameters..." )
433441
442+ expected_pytree = jax .tree_util .tree_map (lambda leaf : leaf , params )
443+
434444 first_leaf = jax .tree_util .tree_leaves (params )[0 ]
435445 target_dtype = first_leaf .dtype
436446
@@ -440,7 +450,9 @@ def convert_and_transpose_tensor(tensor, transpose=False):
440450 return jnp .array (tensor , dtype = target_dtype )
441451
442452 # Global layers
443- params ["context_embedder" ]["kernel" ] = convert_and_transpose_tensor (pt_state_dict .pop ("context_embedder.weight" ), transpose = True )
453+ params ["context_embedder" ]["kernel" ] = convert_and_transpose_tensor (
454+ pt_state_dict .pop ("context_embedder.weight" ), transpose = True
455+ )
444456 params ["x_embedder" ]["kernel" ] = convert_and_transpose_tensor (pt_state_dict .pop ("x_embedder.weight" ), transpose = True )
445457 params ["double_stream_modulation_img" ]["kernel" ] = convert_and_transpose_tensor (
446458 pt_state_dict .pop ("double_stream_modulation_img.linear.weight" ), transpose = True
@@ -454,7 +466,9 @@ def convert_and_transpose_tensor(tensor, transpose=False):
454466 params ["proj_out" ]["kernel" ] = convert_and_transpose_tensor (pt_state_dict .pop ("proj_out.weight" ), transpose = True )
455467
456468 # norm_out
457- params ["norm_out" ]["linear" ]["kernel" ] = convert_and_transpose_tensor (pt_state_dict .pop ("norm_out.linear.weight" ), transpose = True )
469+ params ["norm_out" ]["linear" ]["kernel" ] = convert_and_transpose_tensor (
470+ pt_state_dict .pop ("norm_out.linear.weight" ), transpose = True
471+ )
458472
459473 # time_text_embed (Timestep Embedding)
460474 if "time_guidance_embed.timestep_embedder.linear_1.weight" in pt_state_dict :
@@ -492,18 +506,30 @@ def convert_and_transpose_tensor(tensor, transpose=False):
492506 jax_db ["attn" ]["e_qkv" ]["kernel" ] = jnp .array (np .concatenate ([add_q , add_k , add_v ], axis = 1 ), dtype = target_dtype )
493507
494508 # Projections out
495- jax_db ["attn" ]["i_proj" ]["kernel" ] = convert_and_transpose_tensor (pt_state_dict .pop (prefix + "attn.to_out.0.weight" ), transpose = True )
496- jax_db ["attn" ]["e_proj" ]["kernel" ] = convert_and_transpose_tensor (pt_state_dict .pop (prefix + "attn.to_add_out.weight" ), transpose = True )
509+ jax_db ["attn" ]["i_proj" ]["kernel" ] = convert_and_transpose_tensor (
510+ pt_state_dict .pop (prefix + "attn.to_out.0.weight" ), transpose = True
511+ )
512+ jax_db ["attn" ]["e_proj" ]["kernel" ] = convert_and_transpose_tensor (
513+ pt_state_dict .pop (prefix + "attn.to_add_out.weight" ), transpose = True
514+ )
497515
498516 # Norm scales
499517 jax_db ["attn" ]["query_norm" ]["scale" ] = convert_and_transpose_tensor (pt_state_dict .pop (prefix + "attn.norm_q.weight" ))
500518 jax_db ["attn" ]["key_norm" ]["scale" ] = convert_and_transpose_tensor (pt_state_dict .pop (prefix + "attn.norm_k.weight" ))
501- jax_db ["attn" ]["encoder_query_norm" ]["scale" ] = convert_and_transpose_tensor (pt_state_dict .pop (prefix + "attn.norm_added_q.weight" ))
502- jax_db ["attn" ]["encoder_key_norm" ]["scale" ] = convert_and_transpose_tensor (pt_state_dict .pop (prefix + "attn.norm_added_k.weight" ))
519+ jax_db ["attn" ]["encoder_query_norm" ]["scale" ] = convert_and_transpose_tensor (
520+ pt_state_dict .pop (prefix + "attn.norm_added_q.weight" )
521+ )
522+ jax_db ["attn" ]["encoder_key_norm" ]["scale" ] = convert_and_transpose_tensor (
523+ pt_state_dict .pop (prefix + "attn.norm_added_k.weight" )
524+ )
503525
504526 # SwiGLU MLPs
505- jax_db ["ff" ]["linear_in" ]["kernel" ] = convert_and_transpose_tensor (pt_state_dict .pop (prefix + "ff.linear_in.weight" ), transpose = True )
506- jax_db ["ff" ]["linear_out" ]["kernel" ] = convert_and_transpose_tensor (pt_state_dict .pop (prefix + "ff.linear_out.weight" ), transpose = True )
527+ jax_db ["ff" ]["linear_in" ]["kernel" ] = convert_and_transpose_tensor (
528+ pt_state_dict .pop (prefix + "ff.linear_in.weight" ), transpose = True
529+ )
530+ jax_db ["ff" ]["linear_out" ]["kernel" ] = convert_and_transpose_tensor (
531+ pt_state_dict .pop (prefix + "ff.linear_out.weight" ), transpose = True
532+ )
507533 jax_db ["ff_context" ]["linear_in" ]["kernel" ] = convert_and_transpose_tensor (
508534 pt_state_dict .pop (prefix + "ff_context.linear_in.weight" ), transpose = True
509535 )
@@ -518,8 +544,12 @@ def convert_and_transpose_tensor(tensor, transpose=False):
518544 s_prefix = f"single_transformer_blocks.{ block_idx } ."
519545
520546 # Joint projections
521- jax_sb ["linear1" ]["kernel" ] = convert_and_transpose_tensor (pt_state_dict .pop (s_prefix + "attn.to_qkv_mlp_proj.weight" ), transpose = True )
522- jax_sb ["linear2" ]["kernel" ] = convert_and_transpose_tensor (pt_state_dict .pop (s_prefix + "attn.to_out.weight" ), transpose = True )
547+ jax_sb ["linear1" ]["kernel" ] = convert_and_transpose_tensor (
548+ pt_state_dict .pop (s_prefix + "attn.to_qkv_mlp_proj.weight" ), transpose = True
549+ )
550+ jax_sb ["linear2" ]["kernel" ] = convert_and_transpose_tensor (
551+ pt_state_dict .pop (s_prefix + "attn.to_out.weight" ), transpose = True
552+ )
523553
524554 # Norm scales
525555 jax_sb ["attn" ]["query_norm" ]["scale" ] = convert_and_transpose_tensor (pt_state_dict .pop (s_prefix + "attn.norm_q.weight" ))
@@ -530,7 +560,9 @@ def convert_and_transpose_tensor(tensor, transpose=False):
530560 )
531561 del pt_state_dict
532562 gc .collect ()
533- max_logging .log ("Weight conversion complete!" )
563+ max_logging .log ("Validating converted Flax PyTree state dict structure..." )
564+ validate_flax_state_dict (expected_pytree , params )
565+ max_logging .log ("Weight conversion complete & verified!" )
534566 return params
535567
536568
@@ -553,11 +585,15 @@ def get_pytorch_weight_tensor(key):
553585 max_logging .log ("Mapping VAE decoder weights to JAX parameters..." )
554586
555587 # post_quant_conv
556- jax_params ["post_quant_conv" ]["kernel" ] = jnp .array (get_pytorch_weight_tensor ("post_quant_conv.weight" ).transpose (2 , 3 , 1 , 0 ))
588+ jax_params ["post_quant_conv" ]["kernel" ] = jnp .array (
589+ get_pytorch_weight_tensor ("post_quant_conv.weight" ).transpose (2 , 3 , 1 , 0 )
590+ )
557591 jax_params ["post_quant_conv" ]["bias" ] = jnp .array (get_pytorch_weight_tensor ("post_quant_conv.bias" ))
558592
559593 # decoder.conv_in
560- jax_params ["decoder" ]["conv_in" ]["kernel" ] = jnp .array (get_pytorch_weight_tensor ("decoder.conv_in.weight" ).transpose (2 , 3 , 1 , 0 ))
594+ jax_params ["decoder" ]["conv_in" ]["kernel" ] = jnp .array (
595+ get_pytorch_weight_tensor ("decoder.conv_in.weight" ).transpose (2 , 3 , 1 , 0 )
596+ )
561597 jax_params ["decoder" ]["conv_in" ]["bias" ] = jnp .array (get_pytorch_weight_tensor ("decoder.conv_in.bias" ))
562598
563599 # decoder.mid_block
@@ -621,13 +657,17 @@ def get_pytorch_weight_tensor(key):
621657 upsampler_jax = up_block_jax ["upsamplers_0" ]
622658 upsampler_pt = f"{ up_block_pt } .upsamplers.0"
623659
624- upsampler_jax ["conv" ]["kernel" ] = jnp .array (get_pytorch_weight_tensor (f"{ upsampler_pt } .conv.weight" ).transpose (2 , 3 , 1 , 0 ))
660+ upsampler_jax ["conv" ]["kernel" ] = jnp .array (
661+ get_pytorch_weight_tensor (f"{ upsampler_pt } .conv.weight" ).transpose (2 , 3 , 1 , 0 )
662+ )
625663 upsampler_jax ["conv" ]["bias" ] = jnp .array (get_pytorch_weight_tensor (f"{ upsampler_pt } .conv.bias" ))
626664
627665 # decoder.conv_norm_out & conv_out
628666 jax_params ["decoder" ]["conv_norm_out" ]["scale" ] = jnp .array (get_pytorch_weight_tensor ("decoder.conv_norm_out.weight" ))
629667 jax_params ["decoder" ]["conv_norm_out" ]["bias" ] = jnp .array (get_pytorch_weight_tensor ("decoder.conv_norm_out.bias" ))
630- jax_params ["decoder" ]["conv_out" ]["kernel" ] = jnp .array (get_pytorch_weight_tensor ("decoder.conv_out.weight" ).transpose (2 , 3 , 1 , 0 ))
668+ jax_params ["decoder" ]["conv_out" ]["kernel" ] = jnp .array (
669+ get_pytorch_weight_tensor ("decoder.conv_out.weight" ).transpose (2 , 3 , 1 , 0 )
670+ )
631671 jax_params ["decoder" ]["conv_out" ]["bias" ] = jnp .array (get_pytorch_weight_tensor ("decoder.conv_out.bias" ))
632672
633673 jax_params = jax .tree_util .tree_map (
0 commit comments