Skip to main content

This repository provides a flexible and modular implementation of the Transformer architecture,

Project description

FlexiTransformers

FlexiTransformers Logo

License: MIT PyPI version Python 3.10+ PyTorch DocumentationCode style: Ruff mypy pre-commit

License: MIT PyPI version Python 3.11+ PyTorch Documentation Status Code style: Ruff mypy pre-commit

Note: This library is primarily designed for educational purposes and research experimentation. For production use cases, consider mature frameworks like Hugging Face Transformers.

Build, experiment, and innovate with transformers. FlexiTransformers is a modular Python library for constructing and training transformer models. Choose from encoder-decoder, encoder-only (BERT-style), and decoder-only (GPT-style) architectures, plug in any of 6 positional encoding schemes, and extend the library with your own components.

This library is primarily designed for educational purposes — flexible enough to understand and easy enough to extend.


Features

  • 3 architectures — Encoder-Decoder (T5/BART), Encoder-Only (BERT), Decoder-Only (GPT)
  • 6 positional encodings — Sinusoidal, Learned, Rotary (RoPE), ALiBi, Relative, Relative+Bias
  • Pluggable PE systemregister_pe("name", MyPE) to add custom encodings at runtime
  • KV cache — efficient autoregressive inference with per-layer key/value caching
  • Sampling strategies — greedy, temperature, top-k, nucleus (top-p)
  • Training utilitiesTrainer, LossCompute, LabelSmoothing, run_epoch, callbacks
  • Fully typed — mypy-checked, ruff-formatted, pre-commit enforced

Installation

pip install flexitransformers

Latest development version:

pip install git+https://github.com/A-Elshahawy/flexitransformers.git

Import the library as flexit.

Architecture Variants

Quick Start

Decoder-Only (GPT-style)

from flexit import FlexiGPT, greedy_decode
import torch

model = FlexiGPT(vocab_size=32000, d_model=512, n_heads=8, n_layers=6, d_ff=2048)

src = torch.randint(0, 32000, (1, 10))
out = greedy_decode(model, src, src_mask=None, max_len=50, start_symbol=1)
print(out.shape)  # [1, 59]  (10 prompt + 49 generated tokens)

Encoder-Only (BERT-style)

from flexit import FlexiBERT
import torch

model = FlexiBERT(vocab_size=32000, d_model=512, n_heads=8, n_layers=6,
                  d_ff=2048, num_classes=2)
x = torch.randint(0, 32000, (4, 64))
mask = (x != 0).unsqueeze(1).unsqueeze(2)
logits = model(x, mask)
print(logits.shape)  # [4, 2]

Advanced Config (Encoder-Decoder)

from flexit import ModelConfig, create_model

config = ModelConfig(
    model_type='encoder-decoder',
    src_vocab_size=32000,
    tgt_vocab_size=32000,
    d_model=512,
    n_heads=8,
    n_layers=6,
    d_ff=2048,
    pe_type='rotary',   # absolute | learned | rotary | alibi | relative | relative_bias | none
    dropout=0.1,
)
model = create_model(config)

Architectures

Architecture Convenience Class Config model_type Typical Use
Encoder-Decoder FlexiTransformer "encoder-decoder" Translation, summarization
Encoder-Only FlexiBERT "encoder-only" Classification, NER, embeddings
Decoder-Only FlexiGPT "decoder-only" Language modeling, generation

All three are also accessible via create_model(config) through TransformerFactory.


Positional Encodings

Name pe_type Injected at Representative models
Sinusoidal "absolute" Embedding Vaswani et al. 2017
Learned "learned" Embedding BERT, GPT-2
Rotary (RoPE) "rotary" Q/K projections LLaMA, GPT-NeoX
ALiBi "alibi" Attention scores MPT, BLOOM
Relative "relative" Attention scores Transformer-XL
Relative+Bias "relative_bias" Attention scores T5
None "none" Ablations

Custom PE Plugin

from flexit import register_pe
from flexit.attention.positional import PositionalEncoding

class NoPE(PositionalEncoding):
    @property
    def injection_point(self):
        return "embedding"

    def apply_to_embedding(self, x):
        return x  # pass-through

register_pe("nope", NoPE)

config = ModelConfig(..., pe_type="nope")
model = create_model(config)

Inference

from flexit import greedy_decode, sample_decode

# Greedy decoding
out = greedy_decode(model, src, src_mask=None, max_len=50, start_symbol=1)

# Nucleus (top-p) sampling with temperature
out = sample_decode(model, src, src_mask=None, max_len=50, start_symbol=1,
                    temperature=0.8, top_p=0.9)

Standalone samplers operate on logit tensors directly:

from flexit import temperature_sample, top_k_sample, top_p_sample

next_token = temperature_sample(logits, temperature=0.7)
next_token = top_k_sample(logits, k=50, temperature=1.0)
next_token = top_p_sample(logits, p=0.9, temperature=0.8)

Training

from flexit import (
    ModelConfig, create_model,
    Trainer, LossCompute, LabelSmoothing,
)
import torch

config = ModelConfig(model_type='decoder-only', vocab_size=32000, d_model=512,
                     n_heads=8, n_layers=6, d_ff=2048)
model = create_model(config)

criterion = LabelSmoothing(size=32000, padding_idx=0, smoothing=0.1)
loss_fn = LossCompute(generator=model.generator, criterion=criterion, model=model)
optimizer = torch.optim.Adam(model.parameters(), lr=3e-4, betas=(0.9, 0.98))

trainer = Trainer(
    model=model,
    optimizer=optimizer,
    loss_fn=loss_fn,
    train_dataloader=train_loader,
    val_dataloader=val_loader,       # optional
)
metrics = trainer.fit(epochs=10)
print(metrics.to_dict())

Callbacks

from flexit import CheckpointCallback, EarlyStoppingCallback

trainer = Trainer(
    ...,
    callbacks=[
        CheckpointCallback(save_dir="checkpoints/", monitor="val_loss"),
        EarlyStoppingCallback(patience=3, monitor="val_loss"),
    ],
)

Examples

See the examples/ directory for end-to-end runnable scripts:

File What it demonstrates
01_quick_start.py FlexiGPT, FlexiBERT, FlexiTransformer in 3 lines each
02_manual_config.py ModelConfig + create_model for all 3 architectures
03_positional_encodings.py All 6 PE types end-to-end + create_pe factory
04_encoder_only.py Classification heads (BERT-style)
05_encoder_decoder.py Seq2seq training loop
06_decoder_only.py LM training + greedy decoding
07_output_heads.py LMHead weight tying, all head types
08_save_load.py model.save() / Model.load() round-trip
09_custom_pe.py register_pe plugin with custom PE class

API Reference

Models

Symbol Description
FlexiGPT Decoder-only convenience constructor
FlexiBERT Encoder-only convenience constructor
FlexiTransformer Encoder-decoder convenience constructor
DecoderOnlyModel Full decoder-only model class
EncoderOnlyModel Full encoder-only model class
EncoderDecoderModel Full encoder-decoder model class
BaseModel Abstract base with save() / load()
TransformerModel Generic model via factory

Config & Factory

Symbol Description
ModelConfig Dataclass for all model hyperparameters
create_model(config) Build a model from ModelConfig
TransformerFactory Lower-level factory class

Attention

Symbol Description
MultiHeadAttention Unified MHA with pluggable PE + KV cache

PE Classes

Symbol Description
SinusoidalPE Fixed sinusoidal (embedding-level)
LearnedPE Learned position embeddings
RotaryPE RoPE applied to Q/K
ALiBiPE Linear bias on attention scores
RelativePE Relative position representations
RelativePEWithBias T5-style scalar relative bias
create_pe(config) Factory: config → PE instance
register_pe(name, cls) Register a custom PE class

Core Components

Symbol Description
Embeddings Token embedding with scaling
EmbeddingWithPE Token embedding + positional encoding
FeedForward Standard FFN (ReLU / GELU / SiLU)
GLUFeedForward Gated linear unit FFN (SwiGLU style)
Generator Final linear + softmax projection
LayerNorm Standard layer normalization
RMSNorm Root-mean-square normalization

Layers & Blocks

Symbol Description
EncoderLayer Self-attention + FFN encoder layer
CausalDecoderLayer Masked self-attention + FFN decoder layer
CrossAttentionDecoderLayer Self-attention + cross-attention + FFN layer
SublayerConnection Residual + norm wrapper
Encoder Stack of EncoderLayer
CausalDecoder Stack of CausalDecoderLayer
CrossAttentionDecoder Stack of CrossAttentionDecoderLayer

Output Heads

Symbol Description
LMHead Language model head (supports weight tying)
BertHead BERT pooler + classifier
SequenceClassificationHead CLS-token classification
TokenClassificationHead Per-token classification (NER)

Inference API

Symbol Description
greedy_decode Greedy autoregressive generation
sample_decode Generation with temperature + top-k/p sampling
temperature_sample Temperature-scaled sampling on logits
top_k_sample Top-k sampling on logits
top_p_sample Nucleus (top-p) sampling on logits

Training API

Symbol Description
Trainer Full training loop with callbacks and metrics
Batch Batching + masking utility
LabelSmoothing Label smoothing loss
LossCompute Loss wrapper with gradient step
BertLoss MLM-style loss for encoder-only
run_epoch Single-epoch training/eval loop
Callback Base callback class
CheckpointCallback Save best/latest checkpoints
EarlyStoppingCallback Stop when metric plateaus
TrainerMetrics Metrics container returned by fit()

Utilities

Symbol Description
subsequent_mask Causal (autoregressive) mask
create_causal_mask Causal mask from sequence length
create_padding_mask Padding mask from token ids
create_combined_mask Causal + padding combined
count_parameters Count trainable parameters

Contributing

  1. Fork the repository on GitHub
  2. Create a feature branch: git checkout -b feature/your-feature
  3. Develop your changes — follow the existing code style (ruff, mypy, type annotations)
  4. Write tests for new functionality
  5. Submit a pull request with a clear description of what changed and why

For significant architectural changes, open an issue first to discuss the approach.


License & Credits

Released under the MIT License.

Built on PyTorch. Inspired by the original "Attention Is All You Need" paper and the Hugging Face Transformers library.

Developed and maintained by Ahmed Elshahawy.

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/](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:

1. Fork the repository
2. Create a feature branch (`git checkout -b feature/improvement`)
3. Make your changes with tests
4. Run tests and type checking (`mypy`, `ruff`)
5. Submit a pull request

For major changes, open an issue first to discuss the proposed changes.

## License

MIT License - see [LICENSE](LICENSE) file for details.

## Citation

If you use FlexiTransformers in your research, please cite:

```bibtex
@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

LinkedIn Gmail

Issues and feature requests: GitHub Issues

Project details


Download files

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

Source Distribution

flexitransformers-0.3.0.tar.gz (51.0 kB view details)

Uploaded Source

Built Distribution

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

flexitransformers-0.3.0-py3-none-any.whl (57.7 kB view details)

Uploaded Python 3

File details

Details for the file flexitransformers-0.3.0.tar.gz.

File metadata

  • Download URL: flexitransformers-0.3.0.tar.gz
  • Upload date:
  • Size: 51.0 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.2.0 CPython/3.11.15

File hashes

Hashes for flexitransformers-0.3.0.tar.gz
Algorithm Hash digest
SHA256 7620246240847c5294a13fbb15d5163081bd6f554b88c83662a58fa21be5c9cc
MD5 fa5b2cc9cc314cabd24d05e890fc5b79
BLAKE2b-256 7de06a9b3f6f2a6808a55ee5bac704f620a6b972d4870b28d9f5c337f632d9ca

See more details on using hashes here.

File details

Details for the file flexitransformers-0.3.0-py3-none-any.whl.

File metadata

File hashes

Hashes for flexitransformers-0.3.0-py3-none-any.whl
Algorithm Hash digest
SHA256 79f04b42e0e4d54fb5e6a2d0875ec6052510ce90ae1de604c28cf00b3794ce9e
MD5 b5bb5d9543e412a78a76f625e96edf84
BLAKE2b-256 0e996428ab89983310d0cbbd4e0413bfcd019d6608133ce5a31f92b7621f4d1f

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