Skip to content

Commit 13d1475

Browse files
committed
Handle existing TPU JAX arrays safely in put_data_on_devices
1 parent 9ec2762 commit 13d1475

1 file changed

Lines changed: 2 additions & 0 deletions

File tree

src/maxdiffusion/pipelines/flux/flux2klein_pipeline.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -246,6 +246,8 @@ def __call__(
246246
data_sharding = jax.sharding.NamedSharding(self.mesh, P("data"))
247247

248248
def put_data_on_devices(x, sharding):
249+
if isinstance(x, jax.Array) and hasattr(x, "sharding") and not x.sharding.is_fully_addressable:
250+
return x
249251
if hasattr(sharding, "is_fully_addressable") and sharding.is_fully_addressable:
250252
return jax.device_put(x, sharding)
251253
return device_put_replicated(x, sharding)

0 commit comments

Comments
 (0)