Skip to main content

simpleT5_TRL is built on top of HuggingFace Transformers and TRL that lets you quickly train encoder-decoder models (T5, BART, T5Gemma, etc.) with full finetuning, LoRA, QLoRA, DPO, SimPO, and RFT support.

Project description

simpleT5_TRL

PyPI version Downloads License: MIT

Train encoder-decoder models (T5, BART, T5Gemma, etc.) with Full Finetuning, LoRA, QLoRA, DPO, SimPO & RFT in just a few lines of code.

Built on top of HuggingFace Transformers and TRL.

Supported models: T5, MT5, ByT5, CodeT5, BART, mBART, Pegasus, LED, T5Gemma, and more!


Installation

From PyPI

pip install simplet5_trl

From Source

git clone https://github.com/siddharth0112358/simpleT5_TRL.git
cd simpleT5_TRL
pip install -e .

Quick Start

from simplet5_trl import SimpleT5_TRL
import pandas as pd

# Prepare data
train_df = pd.DataFrame({
    "source_text": ["summarize: This is a long article..."],
    "target_text": ["Short summary"]
})

# Train
model = SimpleT5_TRL()
model.from_pretrained("t5-base")
model.train(train_df=train_df, eval_df=train_df, max_epochs=3)

# Predict
model.load_model("outputs/checkpoint-xxx", use_gpu=True)
print(model.predict("summarize: Your text here"))

Training Methods

Method Function Data Columns Use When
SFT train() source_text, target_text Standard supervised finetuning
DPO train_dpo() prompt, chosen, rejected You have preference pairs
SimPO train_simpo() prompt, chosen, rejected Limited GPU memory (no ref model)
RFT train_rft() source_text, target_text Curated high-quality samples

Choosing the Right Method

Have preference pairs (chosen vs rejected)?
├── Yes → Have enough GPU memory?
│         ├── Yes → DPO
│         └── No  → SimPO
└── No  → RFT (or standard train)

Finetuning Modes

All training methods support three finetuning modes via finetuning= parameter:

Mode Memory Quality Use Case
"full" High Best Small models, enough VRAM
"lora" Low Good Large models, limited VRAM
"qlora" Very Low Good Very large models, consumer GPUs

Hyperparameters

Common Parameters (All Methods)
# Training
max_epochs=3                    # Number of epochs
max_steps=-1                    # Max steps (-1 = use epochs)
batch_size=8                    # Batch size
learning_rate=1e-4              # Learning rate
precision="32"                  # "32", "16", "bf16"
seed=42                         # Random seed

# Optimizer
optim="adamw_torch"             # "adamw_torch", "sgd", "adafactor"
weight_decay=0.0                # Weight decay
warmup_steps=0                  # Warmup steps
warmup_ratio=0.0                # Warmup ratio
lr_scheduler_type="linear"      # "linear", "cosine", "constant", "polynomial"

# Gradient
gradient_accumulation_steps=1   # Gradient accumulation
gradient_checkpointing=False    # Memory-saving checkpointing

# Saving
outputdir="outputs"             # Output directory
save_strategy="epoch"           # "epoch", "steps", "no"
save_steps=500                  # Save every N steps
save_total_limit=None           # Max checkpoints to keep

# Evaluation
eval_strategy="epoch"           # "epoch", "steps", "no"
eval_steps=500                  # Eval every N steps

# Logging
logging_steps=1                 # Log every N steps
report_to=["tensorboard"]       # "tensorboard", "wandb"

# Finetuning
finetuning="full"               # "full", "lora", "qlora"
LoRA Parameters
lora_r=16                       # LoRA rank
lora_alpha=32                   # LoRA alpha
lora_dropout=0.05               # LoRA dropout
lora_target_modules=None        # Auto-detected if None
QLoRA Parameters
quantization="4bit"             # "4bit" or "8bit"
bnb_4bit_compute_dtype="float16"
bnb_4bit_quant_type="nf4"       # "nf4" or "fp4"
bnb_4bit_use_double_quant=True
DPO Parameters
model.train_dpo(
    beta=0.1,                   # Deviation from reference (lower = more deviation)
    loss_type="sigmoid",        # "sigmoid", "hinge", "ipo", "robust"
    label_smoothing=0.0,        # Label smoothing
    max_length=512,             # Max sequence length
    max_prompt_length=256,      # Max prompt length
)
SimPO Parameters
model.train_simpo(
    beta=2.0,                   # SimPO beta (higher than DPO)
    simpo_gamma=0.5,            # Target reward margin
    label_smoothing=0.0,
    max_length=512,
    max_prompt_length=256,
)
RFT Parameters
model.train_rft(
    max_seq_length=512,         # Max sequence length
    packing=False,              # Pack multiple examples
    dataset_text_field="text",  # Column for text data
)
Prediction Parameters
model.predict(
    source_text="input",
    max_length=512,
    num_beams=2,
    do_sample=True,
    temperature=1.0,
    top_k=50,
    top_p=0.95,
    repetition_penalty=2.5,
)

Loading Models

# Full finetuned
model.load_model("outputs/checkpoint-xxx", use_gpu=True)

# LoRA
model.load_model("outputs/checkpoint-xxx", finetuning="lora", base_model_name="t5-base")

# QLoRA
model.load_model("outputs/checkpoint-xxx", finetuning="qlora", base_model_name="t5-large", quantization="4bit")

Supported Models

Family Examples
T5 t5-small, t5-base, t5-large, t5-3b
MT5 google/mt5-small, google/mt5-base
BART facebook/bart-base, facebook/bart-large
CodeT5 Salesforce/codet5-base
Pegasus google/pegasus-xsum
LongT5 google/long-t5-local-base
T5Gemma google/t5gemma-2-270m-270m, google/t5gemma-2b-2b-ul2

Troubleshooting

Missing embed_positions.weight Warning

When loading BART or Pegasus models, you may see:

model.decoder.embed_positions.weight | MISSING
model.encoder.embed_positions.weight | MISSING

This is expected and harmless. These models use sinusoidal positional embeddings computed at runtime, not learned weights.

LongT5-tglobal NaN Issues

The google/long-t5-tglobal-* models may produce NaN values during training. Use the local variant instead:

model.from_pretrained("google/long-t5-local-base")  # Stable
# instead of "google/long-t5-tglobal-base"          # May produce NaN

Acknowledgements


License

MIT License

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

simplet5_trl-0.1.0.tar.gz (25.3 kB view details)

Uploaded Source

Built Distribution

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

simplet5_trl-0.1.0-py3-none-any.whl (22.3 kB view details)

Uploaded Python 3

File details

Details for the file simplet5_trl-0.1.0.tar.gz.

File metadata

  • Download URL: simplet5_trl-0.1.0.tar.gz
  • Upload date:
  • Size: 25.3 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.2.0 CPython/3.11.14

File hashes

Hashes for simplet5_trl-0.1.0.tar.gz
Algorithm Hash digest
SHA256 c5a44e194ed07ce66b4b5b8119331223133f8a6af9cc15d0e25da5bb949070c8
MD5 9012380a49cea8408c654fd1a735fec8
BLAKE2b-256 464abf9a76b4db8e13abb20a970884080f33682180fa0db4e89a000b7a6082f8

See more details on using hashes here.

File details

Details for the file simplet5_trl-0.1.0-py3-none-any.whl.

File metadata

  • Download URL: simplet5_trl-0.1.0-py3-none-any.whl
  • Upload date:
  • Size: 22.3 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.2.0 CPython/3.11.14

File hashes

Hashes for simplet5_trl-0.1.0-py3-none-any.whl
Algorithm Hash digest
SHA256 51b443607b74c20fd415655514d1b8227f3087a4b1d7fcab849307a750cf0bfd
MD5 62514223a589d411770a26a1ef758987
BLAKE2b-256 02d912f74a9a799a2ee98ee6827b38f5b965b77a6927dbfa3b2ae4cdb899c022

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