diff --git a/scripts/compute_norm_stats.py b/scripts/compute_norm_stats.py index c8aef87222..7f4067798b 100644 --- a/scripts/compute_norm_stats.py +++ b/scripts/compute_norm_stats.py @@ -14,6 +14,7 @@ import openpi.training.config as _config import openpi.training.data_loader as _data_loader import openpi.transforms as transforms +import jax class RemoveStrings(transforms.DataTransformFn): @@ -47,9 +48,12 @@ def create_torch_dataloader( else: num_batches = len(dataset) // batch_size shuffle = False + # Single-device sharding: batch size is often 1 for stats; multi-GPU default + # mesh sharding cannot split size-1 batches (see openpi#438). data_loader = _data_loader.TorchDataLoader( dataset, local_batch_size=batch_size, + sharding=jax.sharding.SingleDeviceSharding(jax.devices()[0]), num_workers=num_workers, shuffle=shuffle, num_batches=num_batches,