A wrapper for optimizers in PyTorch to enable decentralized training
Project description
decent-optim-wrapper
A PyTorch optimizer wrapper for decentralized distributed training. This package enables efficient decentralized optimization by wrapping any PyTorch optimizer and managing communication between processes using various network topologies.
Features
- 🔄 Decentralized Training: Enable decentralized optimization without a central parameter server
- 🌐 Multiple Topologies: Support for Ring, Complete, and custom topologies
- 📦 Efficient Communication: Bucket-based gradient communication for reduced overhead
- ⚡ Asynchronous Operations: Non-blocking communication for improved performance
- 🎯 PyTorch Native: Seamless integration with existing PyTorch training code
- 🔧 Flexible: Works with any PyTorch optimizer (SGD, Adam, AdamW, etc.)
Installation
pip install decent-optim-wrapper
Or install with uv:
uv add decent-optim-wrapper
Or install from source:
git clone https://github.com/yourusername/decent-optim-wrapper.git
cd decent-optim-wrapper
pip install -e .
Requirements
- Python >= 3.12
- PyTorch >= 2.8.0
- torchvision >= 0.23.0
- loguru >= 0.7.3
Quick Start
Here's a basic example of using DecentOptimWrapper in a distributed training setup:
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.optim import SGD
from decent_optim_wrapper.wrapper import DecentOptimWrapper
# Initialize distributed training
dist.init_process_group(backend='nccl')
rank = dist.get_rank()
world_size = dist.get_world_size()
local_world_size = torch.cuda.device_count()
# Create your model and base optimizer
model = nn.Linear(10, 10).cuda()
base_optimizer = SGD(model.parameters(), lr=0.01)
# Wrap with DecentOptimWrapper
optimizer = DecentOptimWrapper(
optimizer=base_optimizer,
rank=rank,
world_size=world_size,
local_world_size=local_world_size,
topology='ring', # or 'complete'
bucket_cap_mb=25
)
# Training loop
for epoch in range(num_epochs):
for batch in dataloader:
optimizer.zero_grad()
inputs, targets = batch
outputs = model(inputs)
loss = criterion(outputs, targets)
loss.backward()
optimizer.step() # Handles decentralized averaging
API Reference
DecentOptimWrapper
The main wrapper class for decentralized optimization.
Parameters
- optimizer (
torch.optim.Optimizer): The base PyTorch optimizer to wrap (e.g., SGD, Adam, AdamW) - rank (
int): The rank of the current process in the distributed setup - world_size (
int): Total number of processes participating in training - local_world_size (
int): Number of processes in the local node/group - topology (
str): Network topology for communication. Options:'ring','complete' - bucket_cap_mb (
int, optional): Maximum bucket size in megabytes for gradient bucketing. Default:25
Methods
step(closure=None)
Performs a single optimization step with decentralized averaging.
optimizer.step()
Note: The closure parameter is not supported in this implementation.
zero_grad(set_to_none=True)
Clears the gradients of all optimized parameters.
optimizer.zero_grad()
Parameters:
- set_to_none (
bool): IfTrue, sets gradients toNoneinstead of zero. Default:True
global_avg(may_revert=True)
Performs a global average of parameters across all processes (centralized operation).
optimizer.global_avg()
Parameters:
- may_revert (
bool): IfTrue, allows reverting the global average. Default:True
revert_global_avg()
Reverts the last global average operation, restoring the previous parameter values.
optimizer.revert_global_avg()
Topologies
The wrapper supports different communication topologies for decentralized training:
Ring Topology
In a ring topology, each process communicates with its two neighbors in a circular arrangement. This provides a good balance between communication efficiency and convergence speed.
optimizer = DecentOptimWrapper(
optimizer=base_optimizer,
rank=rank,
world_size=world_size,
local_world_size=local_world_size,
topology='ring'
)
Requirements:
- World size must be even
- Each process alternates between communicating with left and right neighbors
Complete Topology
In a complete topology, all processes communicate with each other in every step, achieving faster convergence at the cost of higher communication overhead.
optimizer = DecentOptimWrapper(
optimizer=base_optimizer,
rank=rank,
world_size=world_size,
local_world_size=local_world_size,
topology='complete'
)
Requirements:
- World size must be greater than 1
Custom Topologies
You can implement custom topologies by extending the Topology class:
from decent_optim_wrapper.topo import Topology, TopologyFactory
class MyCustomTopology(Topology):
def assign_groups(self):
# Implement your custom group assignment logic
# Return a list of lists of lists representing groups for each rank
groups = [[] for _ in range(self._world_size)]
# ... your logic here ...
return groups
# Register the custom topology
TopologyFactory.add_topology('my_custom', MyCustomTopology)
# Use it
optimizer = DecentOptimWrapper(
optimizer=base_optimizer,
rank=rank,
world_size=world_size,
local_world_size=local_world_size,
topology='my_custom'
)
Advanced Usage
Bucket Configuration
The bucket_cap_mb parameter controls how parameters are grouped for communication. Larger buckets reduce communication overhead but increase memory usage:
# Smaller buckets (more communication ops, less memory)
optimizer = DecentOptimWrapper(
optimizer=base_optimizer,
rank=rank,
world_size=world_size,
local_world_size=local_world_size,
topology='ring',
bucket_cap_mb=10 # 10 MB buckets
)
# Larger buckets (fewer communication ops, more memory)
optimizer = DecentOptimWrapper(
optimizer=base_optimizer,
rank=rank,
world_size=world_size,
local_world_size=local_world_size,
topology='ring',
bucket_cap_mb=100 # 100 MB buckets
)
Global Averaging for Evaluation
During evaluation, you might want to perform a global average to get the true ensemble model:
# Before evaluation
optimizer.global_avg()
# Evaluate model
model.eval()
with torch.no_grad():
for batch in val_loader:
outputs = model(batch)
# ... evaluation logic ...
# Optionally revert to continue decentralized training
optimizer.revert_global_avg()
model.train()
Integration with Learning Rate Schedulers
The wrapper works seamlessly with PyTorch learning rate schedulers:
from torch.optim.lr_scheduler import StepLR
base_optimizer = SGD(model.parameters(), lr=0.1)
optimizer = DecentOptimWrapper(
optimizer=base_optimizer,
rank=rank,
world_size=world_size,
local_world_size=local_world_size,
topology='ring'
)
scheduler = StepLR(base_optimizer, step_size=10, gamma=0.1)
for epoch in range(num_epochs):
for batch in dataloader:
optimizer.zero_grad()
loss = train_step(model, batch)
loss.backward()
optimizer.step()
scheduler.step() # Update learning rate
How It Works
The DecentOptimWrapper implements decentralized optimization through the following mechanism:
- Bucketing: Parameters are grouped into buckets based on
bucket_cap_mbfor efficient communication - Local Update: Each process performs a local gradient descent step using the wrapped optimizer
- Asynchronous Communication: Parameters are averaged with neighboring processes according to the topology
- Non-blocking: Communication happens asynchronously to overlap with computation
This approach enables each process to maintain its own model while gradually converging through periodic averaging with neighbors, eliminating the need for a central parameter server.
Contributing
Contributions are welcome! Please feel free to submit issues or pull requests.
Author
Zesen Wang
Email: zesen@kth.se
License
This project is licensed under the MIT License.
Citation
If you use this package in your research, please cite:
@software{decent_optim_wrapper,
author = {Wang, Zesen},
title = {decent-optim-wrapper: A PyTorch Wrapper for Decentralized Optimization},
year = {2025},
url = {https://github.com/yourusername/decent-optim-wrapper}
}
Project details
Release history Release notifications | RSS feed
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file decent_optim_wrapper-0.1.0.tar.gz.
File metadata
- Download URL: decent_optim_wrapper-0.1.0.tar.gz
- Upload date:
- Size: 5.7 kB
- Tags: Source
- Uploaded using Trusted Publishing? Yes
- Uploaded via: uv/0.9.0
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
9a066b6229510f580c284e81bfedaeff8fa2d940b12426b9421cb51a63ac3617
|
|
| MD5 |
eaadeefb28c81c56c8aedd6faf67d523
|
|
| BLAKE2b-256 |
25896ed792b30d6ff6264ca4a93430862db4e55834179bf128cf3f9f29a46a1b
|
File details
Details for the file decent_optim_wrapper-0.1.0-py3-none-any.whl.
File metadata
- Download URL: decent_optim_wrapper-0.1.0-py3-none-any.whl
- Upload date:
- Size: 7.2 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? Yes
- Uploaded via: uv/0.9.0
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
b08bfdf3b026f16587f1a4170fd2b9af89149008f1f17c17c421dd288b97c97a
|
|
| MD5 |
928987de154ef9be1adb6944c121ba69
|
|
| BLAKE2b-256 |
67d49020e2d6818bb6b8189c86af4ee105aa4317a06e84f7f3186a93412e5d54
|