This repository contains a comprehensive PyTorch Lightning trainer for finetuning various models on the ImageNet dataset.
pip install -r requirements.txt
For training DiT models with checkpointing and wandb logging:
# Install additional dependencies
pip install wandb python-dotenv datasets
# Set up wandb (first time only)
wandb login
# Run training
python train.py
The updated trainer includes:
The trainer can be used directly from the command line:
python trainers/tstrq_trainer.py \
--model resnet50 \
--data_dir /path/to/imagenet \
--batch_size 32 \
--epochs 100 \
--lr 1e-4 \
--pretrained \
--freeze_backbone
You can also use the trainer programmatically:
from trainers.tstrq_trainer import train_model
# Train a ResNet50 model
model, trainer = train_model(
model_name="resnet50",
data_dir="/path/to/imagenet",
batch_size=32,
num_epochs=100,
learning_rate=1e-4,
pretrained=True,
freeze_backbone=False
)
Use the provided example script:
python scripts/train_example.py
resnet50 - ResNet-50resnet101 - ResNet-101resnet152 - ResNet-152vgg16 - VGG-16densenet121 - DenseNet-121efficientnet_b0 - EfficientNet-B0dit-xl-2-256 - DiT-XL-2-256 (requires special handling)The ImageNet dataset should be organized as follows:
imagenet/
├── train/
│ ├── n01440764/
│ ├── n01443537/
│ └── ...
└── val/
├── n01440764/
├── n01443537/
└── ...
model_name: Name of the model to trainnum_classes: Number of output classes (default: 1000 for ImageNet)pretrained: Whether to use pretrained weightsfreeze_backbone: Whether to freeze backbone layers for transfer learninglearning_rate: Initial learning rateweight_decay: Weight decay for optimizerbatch_size: Batch size for trainingnum_epochs: Number of training epochsscheduler_gamma: Learning rate decay factorscheduler_step_size: Epochs between learning rate decaydata_dir: Path to ImageNet datasetimage_size: Input image sizenum_workers: Number of data loading workersaccelerator: Training accelerator ("auto", "gpu", "cpu")devices: Number of devices to useprecision: Training precision ("16-mixed", "32", "bf16-mixed")The trainer includes comprehensive data augmentation for training:
The trainer creates the following outputs:
checkpoints/: Saved model checkpointslightning_logs/: TensorBoard logspython trainers/tstrq_trainer.py \
--model resnet50 \
--data_dir /path/to/imagenet \
--batch_size 64 \
--epochs 50 \
--lr 1e-3 \
--pretrained \
--freeze_backbone
python trainers/tstrq_trainer.py \
--model resnet101 \
--data_dir /path/to/imagenet \
--batch_size 32 \
--epochs 100 \
--lr 1e-4 \
--pretrained
python trainers/tstrq_trainer.py \
--model resnet50 \
--data_dir /path/to/imagenet \
--batch_size 128 \
--epochs 100 \
--lr 1e-4 \
--pretrained \
--devices 4 \
--accelerator gpu
Use TensorBoard to monitor training progress:
tensorboard --logdir lightning_logs
--precision 16-mixed)This project is licensed under the MIT License.
67 commits
Python
96.1%
Shell
3.9%
This repository contains a comprehensive PyTorch Lightning trainer for finetuning various models on the ImageNet dataset.
pip install -r requirements.txt
For training DiT models with checkpointing and wandb logging:
# Install additional dependencies
pip install wandb python-dotenv datasets
# Set up wandb (first time only)
wandb login
# Run training
python train.py
The updated trainer includes:
The trainer can be used directly from the command line:
python trainers/tstrq_trainer.py \
--model resnet50 \
--data_dir /path/to/imagenet \
--batch_size 32 \
--epochs 100 \
--lr 1e-4 \
--pretrained \
--freeze_backbone
You can also use the trainer programmatically:
from trainers.tstrq_trainer import train_model
# Train a ResNet50 model
model, trainer = train_model(
model_name="resnet50",
data_dir="/path/to/imagenet",
batch_size=32,
num_epochs=100,
learning_rate=1e-4,
pretrained=True,
freeze_backbone=False
)
Use the provided example script:
python scripts/train_example.py
resnet50 - ResNet-50resnet101 - ResNet-101resnet152 - ResNet-152vgg16 - VGG-16densenet121 - DenseNet-121efficientnet_b0 - EfficientNet-B0dit-xl-2-256 - DiT-XL-2-256 (requires special handling)The ImageNet dataset should be organized as follows:
imagenet/
├── train/
│ ├── n01440764/
│ ├── n01443537/
│ └── ...
└── val/
├── n01440764/
├── n01443537/
└── ...
model_name: Name of the model to trainnum_classes: Number of output classes (default: 1000 for ImageNet)pretrained: Whether to use pretrained weightsfreeze_backbone: Whether to freeze backbone layers for transfer learninglearning_rate: Initial learning rateweight_decay: Weight decay for optimizerbatch_size: Batch size for trainingnum_epochs: Number of training epochsscheduler_gamma: Learning rate decay factorscheduler_step_size: Epochs between learning rate decaydata_dir: Path to ImageNet datasetimage_size: Input image sizenum_workers: Number of data loading workersaccelerator: Training accelerator ("auto", "gpu", "cpu")devices: Number of devices to useprecision: Training precision ("16-mixed", "32", "bf16-mixed")The trainer includes comprehensive data augmentation for training:
The trainer creates the following outputs:
checkpoints/: Saved model checkpointslightning_logs/: TensorBoard logspython trainers/tstrq_trainer.py \
--model resnet50 \
--data_dir /path/to/imagenet \
--batch_size 64 \
--epochs 50 \
--lr 1e-3 \
--pretrained \
--freeze_backbone
python trainers/tstrq_trainer.py \
--model resnet101 \
--data_dir /path/to/imagenet \
--batch_size 32 \
--epochs 100 \
--lr 1e-4 \
--pretrained
python trainers/tstrq_trainer.py \
--model resnet50 \
--data_dir /path/to/imagenet \
--batch_size 128 \
--epochs 100 \
--lr 1e-4 \
--pretrained \
--devices 4 \
--accelerator gpu
Use TensorBoard to monitor training progress:
tensorboard --logdir lightning_logs
--precision 16-mixed)This project is licensed under the MIT License.
67 commits
Python
96.1%
Shell
3.9%