This repository provides a flexible and modular implementation of the Transformer architecture,
Project description
FlexiTransformers
A modular transformer framework for educational purposes, enabling flexible experimentation with encoder-decoder, encoder-only (BERT-style), and decoder-only (GPT-style) architectures.
Note: This library is primarily designed for educational purposes and research experimentation. For production use cases, consider mature frameworks like Hugging Face Transformers.
Features
| Feature | Support |
|---|---|
| Model Types | Encoder-Decoder, Encoder-Only, Decoder-Only |
| Attention Mechanisms | Absolute, ALiBi, Relative (Transformer-XL), Rotary (RoFormer) |
| Positional Encodings | Absolute (sinusoidal), ALiBi, Rotary |
| Normalization | Pre-norm, Post-norm |
| Training Utilities | Built-in Trainer, Callbacks, Learning rate scheduling |
| Custom Architectures | Full configuration control |
Installation
Requirements:
- Python 3.11+
- PyTorch 2.0.1+
Via pip
pip install flexitransformers
From source
git clone https://github.com/A-Elshahawy/flexitransformers.git
cd flexitransformers
pip install -e .
Import the library as flexit in your code.
Quick Start
1. Encoder-Decoder (Seq2Seq Translation)
import torch
from flexit.models import FlexiTransformer
from flexit.utils import subsequent_mask
# Define model configuration
model = FlexiTransformer(
model_type='encoder-decoder',
src_vocab=10000,
tgt_vocab=10000,
d_model=512,
n_heads=8,
n_layers=6,
dropout=0.1,
pe_type='absolute' # or 'alibi', 'rotary'
)
# Create sample data
batch_size, seq_len = 32, 64
src = torch.randint(0, 10000, (batch_size, seq_len))
tgt = torch.randint(0, 10000, (batch_size, seq_len))
# Create masks (assuming 0 is padding)
src_mask = (src != 0).unsqueeze(-2)
tgt_mask = (tgt != 0).unsqueeze(-2) & subsequent_mask(tgt.size(-1))
# Forward pass
output = model(src, tgt, src_mask, tgt_mask)
print(f"Output shape: {output.shape}") # [32, 64, 512]
2. Encoder-Only (BERT-style Classification)
from flexit.models import FlexiBERT
# BERT-style model for binary classification
model = FlexiBERT(
src_vocab=30000,
num_classes=2,
d_model=768,
n_heads=12,
n_layers=12,
pe_type='alibi', # ALiBi works well for BERT-style models
dropout=0.1
)
# Input data
input_ids = torch.randint(0, 30000, (32, 128))
attention_mask = (input_ids != 0).unsqueeze(-2)
# Get classification logits
logits = model(input_ids, attention_mask)
print(f"Logits shape: {logits.shape}") # [32, 2]
3. Decoder-Only (GPT-style Language Model)
from flexit.models import FlexiGPT
# GPT-style autoregressive model
model = FlexiGPT(
tgt_vocab=50000,
d_model=768,
n_heads=12,
n_layers=12,
pe_type='rotary', # Rotary embeddings work well for GPT-style
dropout=0.1
)
# Input sequence
input_ids = torch.randint(0, 50000, (32, 128))
tgt_mask = subsequent_mask(input_ids.size(-1))
# Forward pass
output = model(input_ids, tgt_mask)
print(f"Output shape: {output.shape}") # [32, 128, 768]
Training
Basic Training Loop
import torch.optim as optim
from torch.utils.data import DataLoader
from flexit.train import Trainer, Batch
from flexit.loss import LossCompute
from flexit.callbacks import CheckpointCallback, EarlyStoppingCallback
# Prepare your data
train_loader = DataLoader(your_dataset, batch_size=64, shuffle=True)
val_loader = DataLoader(your_val_dataset, batch_size=64)
# Setup training components
criterion = torch.nn.CrossEntropyLoss(ignore_index=0)
loss_compute = LossCompute(
generator=model.generator,
criterion=criterion,
model=model,
grad_clip=1.0
)
optimizer = optim.Adam(model.parameters(), lr=1e-4, betas=(0.9, 0.98))
scheduler = optim.lr_scheduler.LambdaLR(
optimizer,
lr_lambda=lambda step: min((step + 1) ** -0.5, (step + 1) * 4000 ** -1.5)
)
# Initialize trainer with callbacks
trainer = Trainer(
model=model,
optimizer=optimizer,
scheduler=scheduler,
loss_fn=loss_compute,
train_dataloader=train_loader,
val_dataloader=val_loader,
callbacks=[
CheckpointCallback(save_best=True, keep_last=3),
EarlyStoppingCallback(patience=5, min_delta=0.001)
]
)
# Train the model
metrics = trainer.fit(epochs=20)
print(metrics.to_dict())
Custom Batch Handling
from flexit.train import Batch
# For decoder-only models (GPT-style)
batch = Batch(
tgt=sequence_tensor, # [batch_size, seq_len]
model_type='decoder-only',
pad=0
)
# For encoder-only models (BERT-style)
batch = Batch(
src=input_tensor,
labels=label_tensor,
model_type='encoder-only',
pad=0
)
# For encoder-decoder models
batch = Batch(
src=source_tensor,
tgt=target_tensor,
model_type='encoder-decoder',
pad=0
)
Advanced Configuration
Comparing Attention Mechanisms
# Experiment with different attention types
configs = {
'absolute': {'pe_type': 'absolute'},
'alibi': {'pe_type': 'alibi'},
'rotary': {'pe_type': 'rotary', 'rope_percentage': 0.5},
'relative': {'pe_type': 'relative', 'max_len': 1024}
}
for name, config in configs.items():
model = FlexiTransformer(
model_type='decoder-only',
tgt_vocab=10000,
d_model=512,
n_heads=8,
n_layers=6,
**config
)
# Train and evaluate each variant
Asymmetric Encoder-Decoder
# Different layer counts for encoder/decoder
model = FlexiTransformer(
model_type='encoder-decoder',
src_vocab=10000,
tgt_vocab=10000,
d_model=512,
n_heads=8,
n_layers=(12, 6), # 12 encoder layers, 6 decoder layers
dropout=0.1
)
Custom Initialization
model = FlexiTransformer(
src_vocab=10000,
tgt_vocab=10000,
init_method='kaiming_uniform', # or 'xavier_uniform', 'orthogonal'
ff_activation='gelu', # or 'relu', 'silu'
pre_norm=True # Pre-layer normalization (like GPT)
)
Architecture Variants
Available Model Classes
FlexiTransformer: Base class, fully customizableFlexiBERT: Encoder-only, optimized for classificationFlexiGPT: Decoder-only, optimized for generationTransformerModel: Standard encoder-decoder
Configuration Options
from flexit.configs import ModelConfig
config = ModelConfig(
model_type='encoder-decoder', # or 'encoder-only', 'decoder-only'
src_vocab=10000,
tgt_vocab=10000,
d_model=512, # Model dimension
d_ff=2048, # Feed-forward dimension
n_heads=8, # Attention heads
n_layers=6, # Number of layers (or tuple for asymmetric)
dropout=0.1,
pe_type='absolute', # 'absolute', 'alibi', 'relative', 'rotary'
pre_norm=True, # Pre-norm vs post-norm
ff_activation='relu', # 'relu', 'gelu', 'silu', etc.
init_method='xavier_uniform'
)
API Reference
Full documentation: https://a-elshahawy.github.io/FlexiTransformers/
Key Modules
flexit.models: Model classes (FlexiTransformer,FlexiBERT,FlexiGPT)flexit.attention: Attention mechanisms (Absolute, ALiBi, Relative, Rotary)flexit.train: Training utilities (Trainer,Batch,LossCompute)flexit.callbacks: Training callbacks (CheckpointCallback,EarlyStoppingCallback)flexit.configs: Configuration classes (ModelConfig)flexit.loss: Loss functions (LabelSmoothing,BertLoss)
Contributing
Contributions are welcome. Please:
- Fork the repository
- Create a feature branch (
git checkout -b feature/improvement) - Make your changes with tests
- Run tests and type checking (
mypy,ruff) - Submit a pull request
For major changes, open an issue first to discuss the proposed changes.
License
MIT License - see LICENSE file for details.
Citation
If you use FlexiTransformers in your research, please cite:
@software{flexitransformers2024,
author = {Elshahawy, Ahmed},
title = {FlexiTransformers: A Modular Transformer Framework},
year = {2024},
url = {https://github.com/A-Elshahawy/flexitransformers}
}
References
This library implements concepts from:
- Vaswani et al. (2017) - "Attention is All You Need"
- Press et al. (2021) - "Train Short, Test Long: Attention with Linear Biases" (ALiBi)
- Su et al. (2021) - "RoFormer: Enhanced Transformer with Rotary Position Embedding"
- Dai et al. (2019) - "Transformer-XL: Attentive Language Models Beyond a Fixed-Length Context"
Contact
Ahmed Elshahawy
- GitHub: @A-Elshahawy
- LinkedIn: Ahmed Elshahawy
- Email: ahmedelshahawy078@gmail.com
Links:
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 flexitransformers-0.2.0.tar.gz.
File metadata
- Download URL: flexitransformers-0.2.0.tar.gz
- Upload date:
- Size: 42.1 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.2.0 CPython/3.11.13
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
5f470be8616db9184564ab490e90c88c045332b02c1ae55aea50af4336207786
|
|
| MD5 |
961e83a05110d11e68b9650058a53c06
|
|
| BLAKE2b-256 |
4a539da482de3900f5f345b1ac0f9627e1555b772c012f0b225ccea02a24dcb6
|
File details
Details for the file flexitransformers-0.2.0-py3-none-any.whl.
File metadata
- Download URL: flexitransformers-0.2.0-py3-none-any.whl
- Upload date:
- Size: 43.9 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.2.0 CPython/3.11.13
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
26cbd985cfed6000f62ea8235f18561a3ce4e931be11ca9a034a10f138b1842d
|
|
| MD5 |
ba0f63c4c685a08a357f783c3b531119
|
|
| BLAKE2b-256 |
2c07a9c5e5f43ebad6aca8358c364488f07140091ae776d209652554b4ca4d26
|