Skip to content

Commit 472fad7

Browse files
committed
Shard latents_jax and prompt_embeds_jax across data_sharding for multi-host execution
1 parent 9374611 commit 472fad7

1 file changed

Lines changed: 13 additions & 0 deletions

File tree

src/maxdiffusion/pipelines/flux/flux2klein_pipeline.py

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -21,10 +21,12 @@
2121

2222
import jax
2323
import jax.numpy as jnp
24+
from jax.sharding import PartitionSpec as P
2425
import numpy as np
2526
from flax.linen import partitioning as nn_partitioning
2627

2728
from maxdiffusion import max_logging
29+
from maxdiffusion.max_utils import device_put_replicated
2830
from ..pipeline_flax_utils import FlaxDiffusionPipeline
2931
from ...models.flux.transformers.transformer_flux_flax import Flux2KleinTransformer2DModel
3032
from ...models.vae_flax import FlaxAutoencoderKL
@@ -240,6 +242,17 @@ def __call__(
240242
trace["prompt_encoding"] = time.perf_counter() - t0
241243
max_logging.log(f" -> [TIMING] Prompt Encoding (Qwen3): {trace['prompt_encoding']:.4f} seconds ⏱️")
242244

245+
# Shard pipeline batch inputs across data_sharding for SPMD multi-host execution
246+
data_sharding = jax.sharding.NamedSharding(self.mesh, P(*self._config.data_sharding))
247+
248+
def put_data_on_devices(x, sharding):
249+
if hasattr(sharding, "is_fully_addressable") and sharding.is_fully_addressable:
250+
return jax.device_put(x, sharding)
251+
return device_put_replicated(x, sharding)
252+
253+
latents_jax = put_data_on_devices(latents_jax, data_sharding)
254+
prompt_embeds_jax = put_data_on_devices(prompt_embeds_jax, data_sharding)
255+
243256
# ---------------------------------------------------------------------
244257
# PHASE B: Denoising Loop (Flux Transformer)
245258
# ---------------------------------------------------------------------

0 commit comments

Comments
 (0)