2222from flax .linen import partitioning as nn_partitioning
2323
2424from maxdiffusion import pyconfig
25+ import flax .linen as nn
2526from maxdiffusion .max_utils import create_device_mesh
2627from jax .sharding import Mesh
2728
@@ -251,8 +252,8 @@ def encode_prompt(
251252
252253def encode_prompt_jax (
253254 prompt : Union [str , List [str ]],
254- qwen3_model ,
255255 qwen3_params ,
256+ jitted_qwen3_fn ,
256257 repo_id : str = "black-forest-labs/FLUX.2-klein-4B" ,
257258 max_sequence_length : int = 512 ,
258259):
@@ -295,15 +296,7 @@ def encode_prompt_jax(
295296 input_ids = jnp .array (inputs ["input_ids" ])
296297 attention_mask = jnp .array (inputs ["attention_mask" ])
297298
298- @jax .jit
299- def jitted_qwen3 (q_params , ids , mask ):
300- return qwen3_model .apply (
301- {"params" : q_params },
302- input_ids = ids ,
303- attention_mask = mask ,
304- )
305-
306- hidden_states , all_hidden_states = jitted_qwen3 (qwen3_params , input_ids , attention_mask )
299+ hidden_states , all_hidden_states = jitted_qwen3_fn (qwen3_params , input_ids , attention_mask )
307300
308301 # Extract layers 8, 17, 26 (indices 9, 18, 27 in all_hidden_states)
309302 h_9 = all_hidden_states [9 ]
@@ -678,29 +671,90 @@ def main(argv):
678671 safetensors_path = os .path .join (snapshot_dir , "transformer" , "diffusion_pytorch_model.safetensors" )
679672 vae_safetensors_path = os .path .join (snapshot_dir , "vae" , "diffusion_pytorch_model.safetensors" )
680673
681- # 7. Initialize JAX parameters and load weights for all models (Flux, VAE, Qwen3)
682- print ("Initializing JAX parameters and loading PyTorch weights..." )
674+ # 7. Evaluate shapes and extract TPU shardings using jax.eval_shape
675+ print ("Evaluating shapes and extracting TPU shardings..." )
676+
677+ # Determine sequence lengths based on resolution
678+ h_packed = height // 16
679+ w_packed = width // 16
680+ seq_len_img = h_packed * w_packed
681+ seq_len_txt = config .max_sequence_length
682+
683+ # Define dummy inputs once for reuse
684+ img_dummy = jnp .zeros ((batch_size , seq_len_img , 128 ))
685+ img_ids_dummy = jnp .zeros ((batch_size , seq_len_img , 4 ))
686+ txt_dummy = jnp .zeros ((batch_size , seq_len_txt , 7680 ))
687+ txt_ids_dummy = jnp .zeros ((batch_size , seq_len_txt , 4 ))
688+ vec_dummy = jnp .zeros ((batch_size , 768 ))
689+ t_vec_dummy = jnp .zeros ((batch_size ,))
690+ guidance_vec_dummy = jnp .zeros ((batch_size ,))
691+ dummy_img = jnp .zeros ((batch_size , 3 , 512 , 512 )) # for VAE
692+
693+ # Initialize JAX Qwen3 Config & Model (needed for shape eval)
694+ from transformers import AutoConfig
695+ text_encoder_path = os .path .join (snapshot_dir , "text_encoder" )
696+ print (f"Loading Qwen3 config from text_encoder path: { text_encoder_path } ..." )
697+ pt_config = AutoConfig .from_pretrained (text_encoder_path , local_files_only = True )
698+
699+ qwen3_config = FlaxQwen3Config (
700+ vocab_size = pt_config .vocab_size ,
701+ hidden_size = pt_config .hidden_size ,
702+ intermediate_size = pt_config .intermediate_size ,
703+ num_hidden_layers = pt_config .num_hidden_layers ,
704+ num_attention_heads = pt_config .num_attention_heads ,
705+ num_key_value_heads = pt_config .num_key_value_heads ,
706+ max_position_embeddings = pt_config .max_position_embeddings ,
707+ rms_norm_eps = pt_config .rms_norm_eps ,
708+ rope_theta = pt_config .rope_theta ,
709+ dtype = jnp .bfloat16 if config .weights_dtype == "bfloat16" else jnp .float32 ,
710+ )
711+ qwen3_model = FlaxQwen3Model (qwen3_config )
712+
713+ # Dummy inputs for Qwen3 init
714+ dummy_ids = jnp .zeros ((batch_size , seq_len_txt ), dtype = jnp .int32 )
715+ dummy_mask = jnp .zeros ((batch_size , seq_len_txt ), dtype = jnp .int32 )
716+
717+ key = jax .random .PRNGKey (0 )
718+ key , vae_key , qwen_key = jax .random .split (key , 3 )
719+
720+ def transformer_init_fn ():
721+ return transformer .init (
722+ key ,
723+ hidden_states = img_dummy ,
724+ img_ids = img_ids_dummy ,
725+ encoder_hidden_states = txt_dummy ,
726+ txt_ids = txt_ids_dummy ,
727+ pooled_projections = vec_dummy ,
728+ timestep = t_vec_dummy ,
729+ guidance = guidance_vec_dummy ,
730+ )
731+ def vae_init_fn ():
732+ return vae .init (vae_key , dummy_img )
733+ def qwen3_init_fn ():
734+ return qwen3_model .init (qwen_key , dummy_ids , dummy_mask )
735+
736+ with mesh , nn_partitioning .axis_rules (config .logical_axis_rules ):
737+ abstract_transformer_vars = jax .eval_shape (transformer_init_fn )
738+ abstract_vae_vars = jax .eval_shape (vae_init_fn )
739+ abstract_qwen3_vars = jax .eval_shape (qwen3_init_fn )
740+
741+ logical_transformer_specs = nn .get_partition_spec (abstract_transformer_vars )
742+ logical_vae_specs = nn .get_partition_spec (abstract_vae_vars )
743+ logical_qwen3_specs = nn .get_partition_spec (abstract_qwen3_vars )
744+
745+ transformer_mesh_shardings = nn .logical_to_mesh_sharding (logical_transformer_specs , mesh , config .logical_axis_rules )
746+ vae_mesh_shardings = nn .logical_to_mesh_sharding (logical_vae_specs , mesh , config .logical_axis_rules )
747+ qwen3_mesh_shardings = nn .logical_to_mesh_sharding (logical_qwen3_specs , mesh , config .logical_axis_rules )
748+
749+ transformer_shardings = flax .core .freeze (transformer_mesh_shardings ['params' ])
750+ vae_shardings = flax .core .freeze (vae_mesh_shardings ['params' ])
751+ qwen3_shardings = flax .core .freeze (qwen3_mesh_shardings ['params' ])
752+
753+ # 8. Initialize JAX parameters on CPU
754+ print ("Initializing JAX parameters on CPU..." )
683755 cpu_device = jax .devices ("cpu" )[0 ]
684756 with jax .default_device (cpu_device ):
685757 with mesh , nn_partitioning .axis_rules (config .logical_axis_rules ):
686- # Determine sequence lengths based on resolution
687- h_packed = height // 16
688- w_packed = width // 16
689- seq_len_img = h_packed * w_packed
690- seq_len_txt = config .max_sequence_length
691-
692- # Dummy inputs for transformer init
693- img_dummy = jnp .zeros ((batch_size , seq_len_img , 128 ))
694- img_ids_dummy = jnp .zeros ((batch_size , seq_len_img , 4 ))
695- txt_dummy = jnp .zeros ((batch_size , seq_len_txt , 7680 ))
696- txt_ids_dummy = jnp .zeros ((batch_size , seq_len_txt , 4 ))
697- vec_dummy = jnp .zeros ((batch_size , 768 ))
698- t_vec_dummy = jnp .zeros ((batch_size ,))
699- guidance_vec_dummy = jnp .zeros ((batch_size ,))
700-
701- key = jax .random .PRNGKey (0 )
702- key , vae_key , qwen_key = jax .random .split (key , 3 )
703-
704758 # Initialize Transformer
705759 variables = transformer .init (
706760 key ,
@@ -715,35 +769,11 @@ def main(argv):
715769 params = variables ["params" ]
716770
717771 # Initialize VAE
718- dummy_img = jnp .zeros ((batch_size , 3 , 512 , 512 ))
719772 vae_variables = vae .init (vae_key , dummy_img )
720773 vae_params = vae_variables ["params" ]
721774
722- # Initialize JAX Qwen3 Config & Model
723- from transformers import AutoConfig
724- text_encoder_path = os .path .join (snapshot_dir , "text_encoder" )
725- print (f"Loading Qwen3 config from text_encoder path: { text_encoder_path } ..." )
726- pt_config = AutoConfig .from_pretrained (text_encoder_path , local_files_only = True )
727-
728- qwen3_config = FlaxQwen3Config (
729- vocab_size = pt_config .vocab_size ,
730- hidden_size = pt_config .hidden_size ,
731- intermediate_size = pt_config .intermediate_size ,
732- num_hidden_layers = pt_config .num_hidden_layers ,
733- num_attention_heads = pt_config .num_attention_heads ,
734- num_key_value_heads = pt_config .num_key_value_heads ,
735- max_position_embeddings = pt_config .max_position_embeddings ,
736- rms_norm_eps = pt_config .rms_norm_eps ,
737- rope_theta = pt_config .rope_theta ,
738- dtype = jnp .bfloat16 if config .weights_dtype == "bfloat16" else jnp .float32 ,
739- )
740-
741- qwen3_model = FlaxQwen3Model (qwen3_config )
742-
743775 # Initialize Qwen3 parameters
744776 print ("Initializing JAX Qwen3 parameters..." )
745- dummy_ids = jnp .zeros ((batch_size , seq_len_txt ), dtype = jnp .int32 )
746- dummy_mask = jnp .zeros ((batch_size , seq_len_txt ), dtype = jnp .int32 )
747777 qwen3_variables = qwen3_model .init (qwen_key , dummy_ids , dummy_mask )
748778 qwen3_params = qwen3_variables ["params" ]
749779
@@ -789,6 +819,8 @@ def main(argv):
789819 vae_params = flax .core .freeze (vae_params )
790820 qwen3_params = flax .core .freeze (qwen3_params )
791821
822+
823+
792824 # Dynamic Offloading Auto-Detection
793825 device = jax .devices ()[0 ]
794826 device_kind = device .device_kind
@@ -814,10 +846,9 @@ def main(argv):
814846 print ("\n " + "=" * 80 )
815847 print ("🚀 Dynamic parameter offloading disabled. Moving all parameters to TPU HBM permanently..." )
816848 print ("=" * 80 + "\n " )
817- tpu_device = jax .devices ("tpu" )[0 ]
818- params = jax .device_put (params , tpu_device )
819- vae_params = jax .device_put (vae_params , tpu_device )
820- qwen3_params = jax .device_put (qwen3_params , tpu_device )
849+ params = jax .device_put (params , transformer_shardings )
850+ vae_params = jax .device_put (vae_params , vae_shardings )
851+ qwen3_params = jax .device_put (qwen3_params , qwen3_shardings )
821852
822853 import gc
823854 gc .collect ()
@@ -878,6 +909,14 @@ def jitted_vae_decode(v_params, latents_unpatched):
878909 latents = latents_unpatched ,
879910 method = vae .decode ,
880911 )
912+
913+ @jax .jit
914+ def jitted_qwen3 (q_params , ids , mask ):
915+ return qwen3_model .apply (
916+ {"params" : q_params },
917+ input_ids = ids ,
918+ attention_mask = mask ,
919+ )
881920
882921 # Define a reusable generation function
883922 def run_generation (current_prompts : List [str ], output_name : str , measure_time : bool = False ):
@@ -896,16 +935,16 @@ def run_generation(current_prompts: List[str], output_name: str, measure_time: b
896935
897936 if dynamic_offload :
898937 print (" Moving Qwen3 parameters to TPU HBM..." )
899- q_params_tpu = jax .device_put (qwen3_params , tpu_device )
938+ q_params_tpu = jax .device_put (qwen3_params , qwen3_shardings )
900939 else :
901940 q_params_tpu = qwen3_params
902941
903942 # Run JAX Qwen3 forward pass
904943 prompt_embeds_jax = encode_prompt_jax (
905944 current_prompts ,
906- qwen3_model = qwen3_model ,
907945 qwen3_params = q_params_tpu ,
908- repo_id = config .pretrained_model_name_or_path ,
946+ jitted_qwen3_fn = jitted_qwen3 ,
947+ repo_id = snapshot_dir ,
909948 max_sequence_length = config .max_sequence_length ,
910949 )
911950
@@ -931,7 +970,7 @@ def run_generation(current_prompts: List[str], output_name: str, measure_time: b
931970
932971 if dynamic_offload :
933972 print (" Moving Flux Transformer parameters to TPU HBM..." )
934- t_params_tpu = jax .device_put (params , tpu_device )
973+ t_params_tpu = jax .device_put (params , transformer_shardings )
935974 else :
936975 t_params_tpu = params
937976
@@ -1003,7 +1042,7 @@ def run_generation(current_prompts: List[str], output_name: str, measure_time: b
10031042
10041043 if dynamic_offload :
10051044 print (" Moving VAE parameters to TPU HBM..." )
1006- v_params_tpu = jax .device_put (vae_params , tpu_device )
1045+ v_params_tpu = jax .device_put (vae_params , vae_shardings )
10071046 else :
10081047 v_params_tpu = vae_params
10091048
0 commit comments