Data Parallelism Fundamentals
Data parallelism partitions datasets across computational nodes to accelerate machine learning workflows. Each node maintains a full model replica but processes distinct data subsets. This approach enhances efficiency in large-scale model training through distributed computation.
Synchronous vs. Asynchronous Methods
Synchronous approaches require parameter synchronization after each iteration, while asynchronous methods enable independent node updates at the cost of potential parameter inconsistencies. Implementation variants include:
- Distributed Data Parallelism
- Fully Sharded Data Parallel
- Parameter Server Architectures
- Elastic Parallelism
Data Parallel (DP) Implementation
DP splits mini-batches across devices within a single machine. The workflow consists of:
- Forward Pass: Replicate model and optimizer to all devices
- Backward Pass: Aggregate gradients to primary device for updates
Limitations include Python's Global Interpreter Lock constraint and uneven device utilization during gradient aggregation.
Distributed Datta Parallel (DDP)
DDP extends parallelism across multiple machines using multi-process execution and communication optimizations:
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
def train_process(device_id, num_devices):
dist.init_process_group("nccl", rank=device_id, world_size=num_devices)
model = torch.nn.Linear(10, 10).to(device_id)
ddp_model = DDP(model, device_ids=[device_id])
optimizer = torch.optim.SGD(ddp_model.parameters(), lr=0.001)
# Training loop
outputs = ddp_model(torch.randn(20, 10).to(device_id))
loss = torch.nn.MSELoss()(outputs, torch.randn(20, 10).to(device_id))
loss.backward()
optimizer.step()
if __name__ == "__main__":
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29500"
torch.multiprocessing.spawn(train_process, args=(2,), nprocs=2)
DDP Communication Optimization
Key mechanisms insure efficiency:
- Gradient Bucketing: Groups parameters for collective communication
- Overlap Technique: Concurrent gradient computation and communication
- Autograd Hooks: Trigger immediate communication upon gradient readiness
Distributed Data Loading
Efficient sampling across devices:
# Simplified DistributedSampler
class DistributedSampler:
def __iter__(self):
indices = list(range(len(self.dataset)))
return iter(indices[self.rank::self.world_size])
Asynchronous Parallelism
Nodes update parameters independently without synchronization barriers:
- Compute gradients locally
- Push updates to parameter server
- Pull updated parameters
This approach improves device utilization but requires techniques to handle parameter staleness.
Elastic Training
Torchelastic enables fault-tolerant distributed training:
# Elastic checkpointing example
def main():
state = load_checkpoint(checkpoint_path)
# Initialize distributed group
for epoch in range(state.epoch, total_epochs):
train_epoch(state.model)
state.epoch += 1
save_checkpoint(state)
Rendezvous Mechanism
Coordinates dynamic node participation:
- Nodes join through barrier synchronization
- Role assignment via unique rank allocation
- Shared key-value store initialization