Skip to content

Commit 9ec2762

Browse files
committed
Use PartitionSpec data for input tensor batch sharding
1 parent 472fad7 commit 9ec2762

1 file changed

Lines changed: 2 additions & 2 deletions

File tree

src/maxdiffusion/pipelines/flux/flux2klein_pipeline.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -242,8 +242,8 @@ def __call__(
242242
trace["prompt_encoding"] = time.perf_counter() - t0
243243
max_logging.log(f" -> [TIMING] Prompt Encoding (Qwen3): {trace['prompt_encoding']:.4f} seconds ⏱️")
244244

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))
245+
# Shard pipeline batch inputs across data axis ("data") for SPMD multi-host execution
246+
data_sharding = jax.sharding.NamedSharding(self.mesh, P("data"))
247247

248248
def put_data_on_devices(x, sharding):
249249
if hasattr(sharding, "is_fully_addressable") and sharding.is_fully_addressable:

0 commit comments

Comments
 (0)