From 100892f1a09564d51a438ff3bd26c7d49178a179 Mon Sep 17 00:00:00 2001 From: Bartok9 Date: Fri, 10 Jul 2026 20:01:05 -0400 Subject: [PATCH] fix(scripts): single-device sharding for compute_norm_stats torch path MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Salvage of Physical-Intelligence/openpi#438 by @tlpss — rebased to main. Size-1 batches cannot shard across multi-GPU meshes; pin SingleDeviceSharding for norm-stats collection convenience. --- scripts/compute_norm_stats.py | 4 ++++ 1 file changed, 4 insertions(+) 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,