Skip to content

Commit dce4ded

Browse files
committed
Add put_params_on_devices helper to support multi-host SPMD parameter placement
1 parent ea2603b commit dce4ded

1 file changed

Lines changed: 13 additions & 4 deletions

File tree

src/maxdiffusion/generate_flux2klein.py

Lines changed: 13 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -129,6 +129,7 @@ def main(argv):
129129
pyconfig.initialize(default_args)
130130

131131
# Import modules after jax.distributed.initialize() has run via pyconfig.initialize()
132+
from maxdiffusion.max_utils import device_put_replicated
132133
from maxdiffusion.models.flux.util import (
133134
load_and_convert_flux_klein_weights,
134135
load_and_convert_vae_weights,
@@ -368,12 +369,20 @@ def unbox_fn(x):
368369
max_logging.log("\n" + "=" * 80)
369370
max_logging.log("🚀 Pinning all parameters to TPU HBM permanently...")
370371
max_logging.log("=" * 80 + "\n")
372+
def put_params_on_devices(params_tree, shardings_tree):
373+
def _put_leaf(param, sharding):
374+
if hasattr(sharding, "is_fully_addressable") and sharding.is_fully_addressable:
375+
return jax.device_put(param, sharding)
376+
return device_put_replicated(param, sharding)
377+
378+
return jax.tree_util.tree_map(_put_leaf, params_tree, shardings_tree)
379+
371380
max_logging.log("Putting params on TPU HBM...")
372381
with mesh, nn_partitioning.axis_rules(config.logical_axis_rules):
373382
try:
374-
params = jax.device_put(params, transformer_shardings)
383+
params = put_params_on_devices(params, transformer_shardings)
375384
except Exception as err:
376-
max_logging.log("\njax.device_put(params, transformer_shardings) FAILED!")
385+
max_logging.log("\nput_params_on_devices(params, transformer_shardings) FAILED!")
377386
flat_p = flax.traverse_util.flatten_dict(params)
378387
flat_s = flax.traverse_util.flatten_dict(transformer_shardings)
379388
k_p = set(flat_p.keys())
@@ -383,9 +392,9 @@ def unbox_fn(x):
383392
sys.stdout.flush()
384393
raise err
385394
max_logging.log("Putting vae_params on TPU HBM...")
386-
vae_params = jax.device_put(vae_params, vae_shardings)
395+
vae_params = put_params_on_devices(vae_params, vae_shardings)
387396
max_logging.log("Putting qwen3_params on TPU HBM...")
388-
qwen3_params = jax.device_put(qwen3_params, qwen3_shardings)
397+
qwen3_params = put_params_on_devices(qwen3_params, qwen3_shardings)
389398
max_logging.log("All parameters placed on TPU HBM successfully!")
390399
gc.collect()
391400
jax.effects_barrier()

0 commit comments

Comments
 (0)