Skip to content

Commit 39c2171

Browse files
committed
Validate converted Flax PyTree state dict structure against expected initialization layout
1 parent aceb903 commit 39c2171

1 file changed

Lines changed: 66 additions & 26 deletions

File tree

  • src/maxdiffusion/models/flux

src/maxdiffusion/models/flux/util.py

Lines changed: 66 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)