Skip to main content

Scalable GPT implementation using Distributed Data Parallel for efficient, multi-GPU training of transformer models.

Project description

GPT From Scratch

GPT From Scratch is an open-source implementation of a GPT-style transformer model, designed for scalability and efficiency. By leveraging Distributed Data Parallel (DDP) training, this project allows for efficient multi-GPU training, making it easier to train large-scale models on extensive datasets.

Features

  • Built from scratch GPT-style transformer model
  • Efficient multi-GPU training using PyTorch's Distributed Data Parallel (DDP)
  • Scalable and optimized for large datasets
  • Easy-to-use interface for training and inference

Results

Our model outperforms the GPT-2 checkpoint values on the HellaSwag benchmark, demonstrating the effectiveness of our implementation.

Model Output

Model Before Training:

On a clear day indis DNSヘラ ignore Happ Ce Croatian mugVAavorable303 wayomb prom bartender surmia pass standingotoshanMore intensely Lent loaf

Model After Training:

On a clear day, our community is growing and we are enjoying the beauty and tranquility of our natural surroundings. While it may not be a

Training Loss and HellaSwag Evaluation

Here is a figure showing the training/validation loss and model's performance on the HellaSwag benchmark:

Model's Performance

Installation

Install from PyPI

To install the package from PyPI, run the following command:

pip install gpt_from_scratch

Install from Github

To install the latest version directly from GitHub, use:

pip install git+https://github.com/MatinKhajavi/GPT-from-scratch.git

Usage

Tokenizing Data

Before training, you need to tokenize your data and prepare it in shards. Use the provided tokenize_documents_parallel.py script to process your data efficiently with multiprocessing.

Command:

python scripts/tokenize_documents_parallel.py --local_dir "edu_fineweb10B"

Parameters:

  • --local_dir: Directory where tokenized data will be stored.
  • --remote_name: Identifier for the remote dataset to download.
  • --shard_size: Number of tokens each data shard will contain.
  • --dataset_path: Full path to the dataset on the data hosting platform (e.g., Hugging Face).
  • --tokenizer_model: The tokenizer model to use.

Training the Model

You can train the model using a single GPU or multiple GPUs with Distributed Data Parallel (DDP).

Single GPU Training

Command:

python scripts/train_gpt.py

Multi-GPU Training

For multi-GPU training, ensure that your environment variables (WORLD_SIZE, RANK, LOCAL_RANK) are set correctly to facilitate distributed training. The torchrun command simplifies this setup.

Command:

torchrun --standalone --nproc_per_node=8 scripts/train_gpt.py

Parameters:

  • --n_batches: Number of batches for training or validation per iteration.
  • --n_tokens: Number of tokens to process per batch.
  • --data_root: Directory where tokenized data is stored.
  • --vocab_size: Total number of unique tokens in the model's vocabulary.
  • --emb_dim: Dimension of the embedding layer.
  • --context_length: The length of the input sequences.
  • --drop_rate: Dropout rate to use within the model to prevent overfitting.
  • --n_layers: Number of layers in the transformer model.
  • --n_heads: Number of attention heads in each transformer layer.
  • --qkv_bias: Enable bias in the query, key, and value projections within attention layers.
  • --monitor: Toggle to enable performance monitoring during training.
  • --torch_matmul_precision: Precision setting for matrix multiplications in PyTorch.
  • --log_dir: Directory to store training logs.
  • --n_epochs: Total number of training epochs.
  • --warmup_iters: Number of iterations to linearly increase the learning rate from zero to the initial rate.
  • --max_iters: Maximum number of iterations to perform during training.
  • --total_batch_size: Total batch size across all distributed training instances.
  • --metrics: Metrics used to evaluate the model's performance.
  • --max_lr: Maximum learning rate used in the learning rate scheduler.
  • --min_lr: Minimum learning rate as part of the cyclical learning rate schedule.

Example Commands

For a complete training session on 8 GPUs with specific parameters:

torchrun --standalone --nproc_per_node=8 scripts/train_gpt.py --data_root "data/edu_fineweb10B" --n_batches 16 --n_tokens 1024 --vocab_size 50304 --emb_dim 768 --context_length 1024 --n_layers 12 --n_heads 12

License

This project is licensed under the MIT License.

Project details


Release history Release notifications | RSS feed

This version

1.1

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

gpt_from_scratch-1.1.tar.gz (14.8 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

gpt_from_scratch-1.1-py3-none-any.whl (18.2 kB view details)

Uploaded Python 3

File details

Details for the file gpt_from_scratch-1.1.tar.gz.

File metadata

  • Download URL: gpt_from_scratch-1.1.tar.gz
  • Upload date:
  • Size: 14.8 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/5.1.1 CPython/3.12.2

File hashes

Hashes for gpt_from_scratch-1.1.tar.gz
Algorithm Hash digest
SHA256 8b0c4980dbb7d9b359c571f473141813d864e31bc245a6f21134366e73172166
MD5 c6e549aaed9527022bbe6b1be57dbda2
BLAKE2b-256 d0c24d44650e495b220518a74155d8ed408a9d575bcdd637847e1b0aef731bc0

See more details on using hashes here.

File details

Details for the file gpt_from_scratch-1.1-py3-none-any.whl.

File metadata

  • Download URL: gpt_from_scratch-1.1-py3-none-any.whl
  • Upload date:
  • Size: 18.2 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/5.1.1 CPython/3.12.2

File hashes

Hashes for gpt_from_scratch-1.1-py3-none-any.whl
Algorithm Hash digest
SHA256 9b91cabbd0be11930696e92e796981f8e3d827a763584d8e295d0ca2e719b3d1
MD5 79aaf2a7bb46af6a22bad902c389194c
BLAKE2b-256 31663a187b63fd04191524af804442d1e8a47f464cd2b63b40ec5c7588ce508b

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page