Skip to content
WodlfvllfPublic

About

QuintNet is a research-oriented PyTorch framework designed to explore and implement multi-dimensional parallelism strategies for distributed deep learning.

Topics

Resources

Stars

19 stars

Watchers

0 watching

Forks

Latest commit

Β 

History

454 Commits

Folders and files

NameName
Last commit message
Last commit date
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 
Β 

Repository files navigation

πŸš€ QuintNet

A PyTorch Framework for 3D Distributed Deep Learning

Data Parallel β€’ Pipeline Parallel β€’ Tensor Parallel


✨ Overview

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)  β”‚         β”‚
β”‚  β””β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”˜  β””β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”˜  β””β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”˜         β”‚
β””β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”€β”˜

🎯 Key Features

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

πŸ“¦ Installation

Prerequisites

  • Python 3.8+
  • PyTorch 2.0+ with CUDA support
  • NCCL backend for distributed training

Quick Install

# 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.txt

PyTorch with CUDA (Recommended)

conda install pytorch pytorch-cuda=12.1 -c pytorch -c nvidia

πŸš€ Quick Start

Training with 3D Parallelism

from 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()

Running Examples

# 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.py

πŸ“ Project Structure

QuintNet/
β”œβ”€β”€ 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

βš™οΈ 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'

πŸ”§ Parallelism Strategies

Data Parallelism

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_dp

Pipeline Parallelism

Splits model layers across GPUs. Uses micro-batching with 1F1B schedule for efficiency.

torchrun --nproc_per_node=4 -m QuintNet.examples.simple_pp

Tensor Parallelism

Splits individual layer weights across GPUs. Useful for very large layers (e.g., LLM attention/FFN).

torchrun --nproc_per_node=2 -m QuintNet.examples.simple_tp

3D Hybrid Parallelism

Combines 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_3d

πŸ“Š Results

Training 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)

Training Configuration

  • 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

πŸ€– GPT-2 Text Summarization

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.

Training Results

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)

Training Progress Visualization

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

Key Observations

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

Training Analysis

  • 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

GPT-2 Training Configuration

  • 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

πŸ§ͺ Testing

# Run all tests
pytest

# Run specific test
pytest tests/test_data_parallel.py -v

πŸ› οΈ Development

Adding a New Strategy

  1. Create a new strategy in strategy/strategies/
  2. Inherit from BaseParallelismStrategy
  3. Implement apply() method
  4. Register in strategy/__init__.py
class MyStrategy(BaseParallelismStrategy):
    def apply(self, model: nn.Module) -> nn.Module:
        # Your parallelism logic here
        return wrapped_model

πŸ“š Documentation

πŸ“– 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

Key Source Files:

  • parallelism/pipeline_parallel/schedule.py - 1F1B schedule implementation
  • core/communication.py - Distributed primitives with autograd support
  • parallelism/data_parallel/core/ddp.py - DDP implementation details
  • parallelism/tensor_parallel/layers.py - Column/Row parallel layers

License: MIT

Built for learning and scaling deep learning 🧠

About

QuintNet is a research-oriented PyTorch framework designed to explore and implement multi-dimensional parallelism strategies for distributed deep learning.

Topics

Resources

Stars

19 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages