A PyTorch Framework for 3D Distributed Deep Learning
Data Parallel β’ Pipeline Parallel β’ Tensor Parallel
QuintNet is an educational and production-ready PyTorch library that implements 3D parallelism for training large-scale deep learning models across multiple GPUs. It provides clean, well-documented implementations of:
- Data Parallelism (DP) - Replicate model, split data
- Pipeline Parallelism (PP) - Split model layers across GPUs
- Tensor Parallelism (TP) - Split individual layers across GPUs
- Hybrid 3D Parallelism - Combine all three for maximum scalability
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β 3D Parallelism β
β βββββββββββββββ βββββββββββββββ βββββββββββββββ β
β β Data β β Pipeline β β Tensor β β
β β Parallel ββββ Parallel ββββ Parallel β β
β β (Batch) β β (Layers) β β (Weights) β β
β βββββββββββββββ βββββββββββββββ βββββββββββββββ β
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
| Feature | Description |
|---|---|
| Modular Design | Each parallelism strategy is independent and composable |
| 1F1B Schedule | Efficient pipeline schedule minimizing memory footprint |
| Gradient Bucketing | Optimized gradient synchronization for DP |
| Device Mesh | Flexible N-dimensional device topology |
| Zero Boilerplate | Simple strategy-based API for applying parallelism |
- Python 3.8+
- PyTorch 2.0+ with CUDA support
- NCCL backend for distributed training
# Clone the repository
git clone https://github.com/yourusername/QuintNet.git
cd QuintNet
# Install in development mode
pip install -e .
# Install dependencies
pip install -r requirements.txtconda install pytorch pytorch-cuda=12.1 -c pytorch -c nvidiafrom QuintNet import Trainer, get_strategy, init_process_groups
# Initialize distributed environment
pg_manager = init_process_groups(
mesh_dim=[2, 2, 2], # [DP, TP, PP] dimensions
mesh_name=['dp', 'tp', 'pp']
)
# Apply 3D parallelism strategy
strategy = get_strategy('3d', pg_manager, config)
parallel_model = strategy.apply(model)
# Train with the Trainer
trainer = Trainer(parallel_model, train_loader, val_loader, config, pg_manager)
trainer.fit()# Single-node, 8 GPUs with 3D parallelism
torchrun --nproc_per_node=8 -m QuintNet.examples.full_3d --config QuintNet/examples/config.yaml
# Or using Modal for cloud training
modal run train_modal_run.pyQuintNet/
βββ core/ # Core distributed primitives
β βββ communication.py # Send, Recv, AllGather, AllReduce
β βββ device_mesh.py # N-dimensional device topology
β βββ process_groups.py # Process group management
β
βββ parallelism/
β βββ data_parallel/ # Data Parallelism (DDP)
β β βββ core/ddp.py # DataParallel wrapper
β β βββ components/ # Gradient reducer, parameter broadcaster
β β
β βββ pipeline_parallel/ # Pipeline Parallelism
β β βββ wrapper.py # PipelineParallelWrapper
β β βββ schedule.py # 1F1B and AFAB schedules
β β βββ trainer.py # PipelineTrainer
β β
β βββ tensor_parallel/ # Tensor Parallelism
β βββ layers.py # ColumnParallelLinear, RowParallelLinear
β βββ model_wrapper.py # Automatic layer replacement
β
βββ coordinators/ # Multi-strategy coordinators
β βββ hybrid_3d_coordinator.py
β
βββ strategy/ # High-level strategy API
β βββ base.py
β βββ strategies/ # DP, PP, TP, 3D strategies
β
βββ trainer.py # Main Trainer class
β
βββ docs/
β βββ TRAINING_GUIDE.md # π Complete training workflow guide
β
βββ examples/
βββ full_3d.py # Complete 3D training example
βββ simple_dp.py # Data Parallel example
βββ simple_pp.py # Pipeline Parallel example
βββ simple_tp.py # Tensor Parallel example
βββ config.yaml # Training configuration
Create a config.yaml file:
# Training
dataset_path: /path/to/dataset
batch_size: 32
num_epochs: 10
learning_rate: 1e-4
grad_acc_steps: 2
# Model
img_size: 28
patch_size: 4
hidden_dim: 64
depth: 8
n_heads: 4
# Parallelism
mesh_dim: [2, 2, 2] # [DP, TP, PP]
mesh_name: ['dp', 'tp', 'pp']
strategy_name: '3d'
schedule: '1f1b'Replicates the full model on each GPU. Each GPU processes a different batch, gradients are synchronized via AllReduce.
torchrun --nproc_per_node=4 -m QuintNet.examples.simple_dpSplits model layers across GPUs. Uses micro-batching with 1F1B schedule for efficiency.
torchrun --nproc_per_node=4 -m QuintNet.examples.simple_ppSplits individual layer weights across GPUs. Useful for very large layers (e.g., LLM attention/FFN).
torchrun --nproc_per_node=2 -m QuintNet.examples.simple_tpCombines all three strategies. Requires DP Γ TP Γ PP GPUs.
# 8 GPUs: 2 DP Γ 2 TP Γ 2 PP
torchrun --nproc_per_node=8 -m QuintNet.examples.full_3dTraining a Vision Transformer on MNIST with 8 GPUs (2Γ2Γ2 mesh):
| Epoch | Train Loss | Train Acc | Val Loss | Val Acc |
|---|---|---|---|---|
| 1 | 1.3817 | 50.46% | 0.8921 | 69.30% |
| 2 | 0.6662 | 77.72% | 0.5135 | 84.52% |
| 3 | 0.4219 | 86.33% | 0.3477 | 89.24% |
| 4 | 0.3214 | 90.02% | 0.2883 | 91.16% |
| 5 | 0.2728 | 91.86% | 0.2509 | 92.06% |
| 6 | 0.2477 | 92.96% | 0.2510 | 92.50% |
| 7 | 0.2364 | 93.78% | 0.2464 | 92.76% |
| 8 | 0.2355 | 94.36% | 0.2372 | 93.18% |
| 9 | 0.2450 | 94.46% | 0.2726 | 93.16% |
| 10 | 0.2573 | 94.80% | 0.3190 | 93.24% |
Final Accuracy: 93.24% | Training Time: 1120.72 seconds (~18.7 minutes)
- Model: Vision Transformer (64 hidden dim, 8 blocks, 4 heads)
- Dataset: MNIST (60,000 train, 10,000 test)
- Batch Size: 32 (effective: 32 Γ 2 DP = 64)
- Parallelism: 2 Data Γ 2 Tensor Γ 2 Pipeline
We also benchmarked QuintNet on GPT-2 fine-tuning for text summarization using Modal cloud infrastructure, demonstrating the framework's capability for large language model training.
| Epoch | Train Loss | Train PPL | Val Loss | Val PPL |
|---|---|---|---|---|
| 1 | 8.2477 | 3818.76 | 7.3738 | 1593.61 |
| 2 | 6.9159 | 1008.13 | 5.6890 | 295.59 |
| 3 | 5.3281 | 206.05 | 3.3037 | 27.21 |
π Perplexity Reduction: 3818.76 β 206.05 (Train) | 1593.61 β 27.21 (Val)
Validation Perplexity
β
1600 β€ β
β β²
1200 β€ β²
β β²
800 β€ β²
β β²
400 β€ β²
β β
100 β€ β²
β β
0 βΌββββββββββββββββββ
1 2 3 Epoch
Train Loss
β
8.5 β€ β
β β²
7.5 β€ β²
β β
6.5 β€ β²
β β²
5.5 β€ β²
β β
5.0 βΌββββββββββββββββββ
1 2 3 Epoch
| Metric | Epoch 1 β 3 Improvement | Analysis |
|---|---|---|
| Train Loss | 8.25 β 5.33 | 35.4% reduction |
| Train PPL | 3819 β 206 | 94.6% reduction |
| Val Loss | 7.37 β 3.30 | 55.2% reduction |
| Val PPL | 1594 β 27 | 98.3% reduction |
- Strong Generalization: Validation perplexity improved more dramatically (98.3%) than training perplexity (94.6%), indicating effective generalization without overfitting
- Rapid Convergence: Significant improvements within just 3 epochs showcase efficient distributed training
- Low Final PPL: Achieving a validation perplexity of 27.21 indicates the model has learned meaningful text summarization patterns
- Model: GPT-2 (124M parameters)
- Task: Text Summarization
- Infrastructure: Modal Cloud (Multi-GPU)
- Framework: QuintNet with 3D Parallelism
# Running GPT-2 training on Modal
modal run QuintNet/gpt2_train_modal_run.py::main# Run all tests
pytest
# Run specific test
pytest tests/test_data_parallel.py -v- Create a new strategy in
strategy/strategies/ - Inherit from
BaseParallelismStrategy - Implement
apply()method - Register in
strategy/__init__.py
class MyStrategy(BaseParallelismStrategy):
def apply(self, model: nn.Module) -> nn.Module:
# Your parallelism logic here
return wrapped_modelπ Complete Training Guide - Detailed walkthrough with diagrams explaining:
- Device Mesh and Process Groups
- Model Wrapping Pipeline (TP β PP β DP)
- Data Flow Architecture
- 1F1B Pipeline Schedule
- Gradient Synchronization
parallelism/pipeline_parallel/schedule.py- 1F1B schedule implementationcore/communication.py- Distributed primitives with autograd supportparallelism/data_parallel/core/ddp.py- DDP implementation detailsparallelism/tensor_parallel/layers.py- Column/Row parallel layers
Built for learning and scaling deep learning π§