try: # Automatically detects jax.process_index() and jax.process_count() shard_options = grain.ShardByJaxProcess(drop_remainder=True) except ImportError: # Fallback for a single process setup shard_options = grain.ShardOptions( shard_index=0, shard_count=1, drop_remainder=True ) sampler = grain.IndexSampler( num_records=len(source), shard_options=shard_options, shuffle=True, num_epochs=None, seed=42, ) __ __