Sharded Pipeline Quick Reference
September 16, 2026 ยท View on GitHub
| Metadata | Value |
|---|---|
| Level | Intermediate |
| Runtime | ~5 min |
| Prerequisites | Basic Datarax pipeline, JAX sharding concepts |
| Format | Python + Jupyter |
Overview
Distribute data processing across multiple JAX devices using Datarax sharding. This enables efficient utilization of multi-GPU setups for large-scale data pipelines, essential for training on large datasets.
What You'll Learn
- Create a JAX device mesh for multi-device execution
- Configure Datarax pipelines for sharded data distribution
- Verify data is properly distributed across devices
- Handle single-device fallback gracefully
Coming from PyTorch?
| PyTorch | Datarax |
|---|---|
DistributedSampler(dataset) | JAX Mesh with PartitionSpec |
DataParallel(model) | Data sharded along batch dimension |
torch.distributed.init_process_group() | DeviceMeshManager.create_data_parallel_mesh() |
sampler.set_epoch(epoch) | RNG-based shuffling per device |
Key difference: Datarax uses JAX's built-in GSPMD for transparent sharding without explicit communication.
Coming from TensorFlow?
| TensorFlow | Datarax |
|---|---|
tf.distribute.MirroredStrategy | Mesh with data axis |
strategy.experimental_distribute_dataset | jax.device_put(batch, sharding) |
tf.distribute.Strategy.scope() | jax.set_mesh(mesh) |
strategy.reduce() | JAX handles via GSPMD |
Files
- Python Script:
examples/advanced/distributed/01_sharding_quickref.py - Jupyter Notebook:
examples/advanced/distributed/01_sharding_quickref.ipynb
Quick Start
python examples/advanced/distributed/01_sharding_quickref.py
Architecture
flowchart TB
subgraph Source["Data Source"]
D[MemorySource<br/>1024 samples]
end
subgraph Pipeline["Pipeline"]
P[Pipeline<br/>batch_size=128]
end
subgraph Mesh["Device Mesh"]
direction LR
G0[GPU 0<br/>batch[0:64]]
G1[GPU 1<br/>batch[64:128]]
end
D --> P
P --> G0
P --> G1
Key Concepts
Step 1: Check Device Availability
import jax
devices = jax.devices()
use_sharding = len(devices) >= 2
print(f"JAX devices: {devices}")
print(f"Device count: {len(devices)}")
Terminal Output:
JAX devices: [cuda:0, cuda:1]
Device count: 2
Step 2: Create Data and Pipeline
Standard pipeline setup - sharding is applied at the mesh level, not by changing how you define sources or operators:
from datarax.operators import ElementOperator, ElementOperatorConfig
from datarax.pipeline import Pipeline
from datarax.sources import MemorySource, MemorySourceConfig
num_samples = 1024
data = {
"image": np.random.rand(num_samples, 32, 32, 3).astype(np.float32),
"feature": np.random.rand(num_samples, 128).astype(np.float32),
"label": np.random.randint(0, 10, (num_samples,)).astype(np.int32),
}
source = MemorySource(MemorySourceConfig(), data=data, rngs=nnx.Rngs(0))
def normalize(element, key=None):
"""Normalize image to [0, 1] range."""
return element.update_data({"image": element.data["image"] / 255.0})
normalizer = ElementOperator(
ElementOperatorConfig(stochastic=False), fn=normalize, rngs=nnx.Rngs(0)
)
pipeline = Pipeline(source=source, stages=[normalizer], batch_size=128, rngs=nnx.Rngs(0))
print("Pipeline created with batch_size=128")
Terminal Output:
Pipeline created with batch_size=128
Step 3: Create Device Mesh
from substrax.mesh import DeviceMeshManager
from substrax.spmd import create_data_parallel_sharding
# Create a mesh for data parallelism (Auto axes)
mesh = DeviceMeshManager.create_data_parallel_mesh()
# Split the batch dimension of every array across the "data" axis
batch_sharding = create_data_parallel_sharding(mesh)
print(f"Created mesh with {mesh.devices.size} devices along 'data' axis")
Terminal Output:
Created mesh with 2 devices along 'data' axis
Step 4: Process with Sharding
from substrax.spmd import place_batch_on_shards
with jax.set_mesh(mesh):
for i, batch in enumerate(pipeline):
if i >= 2:
break
# Place every array of the batch on the mesh
sharded_batch = place_batch_on_shards(batch, batch_sharding)
print(f"Batch {i}:")
print(f" Image shape: {sharded_batch['image'].shape}")
print(f" Image sharding: {sharded_batch['image'].sharding}")
print(f" Label shape: {sharded_batch['label'].shape}")
Terminal Output (multi-GPU):
Batch 0:
Image shape: (128, 32, 32, 3)
Image sharding: NamedSharding(mesh=..., spec=PartitionSpec('data',))
Label shape: (128,)
Mesh Configurations
| Pattern | Mesh Shape | Use Case |
|---|---|---|
| Data Parallel | ("data",) | Replicate model, shard data |
| Model Parallel | ("model",) | Shard model, replicate data |
| Hybrid | ("data", "model") | Large models + large batches |
Results Summary
| Feature | Value |
|---|---|
| Device Count | Depends on system |
| Mesh Shape | (N,) for N devices |
| Data Parallelism | Batch dimension sharded |
| Fallback | Single-device execution |
Sharding benefits:
- Memory efficiency: Data distributed across device memories
- Throughput: Parallel preprocessing on multiple devices
- Scalability: Easily scales with more devices
Next Steps
- Sharding Guide - Advanced sharding patterns
- Checkpointing - Save distributed state
- Performance Guide - Optimize throughput
- API Reference: Sharding - Complete API