Skip to content

Commit a86a6fe

Browse files
committed
Fixes to support FSDP
1 parent 379c1a0 commit a86a6fe

1 file changed

Lines changed: 102 additions & 63 deletions

File tree

src/maxdiffusion/generate_flux2klein.py

Lines changed: 102 additions & 63 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,7 @@
2222
from flax.linen import partitioning as nn_partitioning
2323

2424
from maxdiffusion import pyconfig
25+
import flax.linen as nn
2526
from maxdiffusion.max_utils import create_device_mesh
2627
from jax.sharding import Mesh
2728

@@ -251,8 +252,8 @@ def encode_prompt(
251252

252253
def 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

Comments
 (0)