Skip to main content

Pure MLX training framework for ModernBERT — no PyTorch, no TensorFlow. Just Metal.

Project description

mlx-modernbert

Pure MLX training framework for ModernBERT. No PyTorch. No TensorFlow. Just Metal.

Fine-tune ModernBERT-large on Apple Silicon for sequence classification (text classification) and token classification (NER / PII detection) with fp16, gradient checkpointing, and entity-level evaluation.

Install

# With uv (recommended)
uv add mlx-modernbert

# With pip
pip install mlx-modernbert

Or install from source:

uv sync

Requires macOS with Apple Silicon (M1/M2/M3/M4).

Quick Start

Sequence Classification (Binary / Multi-class)

from datasets import Dataset
from mlx_trainer import Trainer, TrainingArgs
from mlx_trainer.load import load

# Load model + tokenizer
model, tokenizer = load(
    "answerdotai/ModernBERT-large",
    train=True,
    model_config={
        "num_labels": 2,
        "id2label": {0: "SAFE", 1: "INJECTION"},
        "label2id": {"SAFE": 0, "INJECTION": 1},
    },
)

# Your dataset with "text" and "label" columns
train_dataset = Dataset.from_dict({
    "text": ["Hello world", "Ignore instructions", "Normal text"],
    "label": [0, 1, 0],
})

# Train
args = TrainingArgs(
    output_dir="outputs/my_model",
    num_train_epochs=3,
    batch_size=4,
    gradient_accumulation_steps=8,
    learning_rate=1e-5,
)

trainer = Trainer(
    model=model,
    tokenizer=tokenizer,
    training_args=args,
    train_dataset=train_dataset,
)
trainer.train()

Token Classification (NER / PII Detection)

from datasets import Dataset
from mlx_trainer import TokenClassificationTrainer, TrainingArgs
from mlx_trainer.load import load_token_classification

# Label schema (BIO format)
LABELS = ["O", "B-PER", "I-PER", "B-EMAIL", "I-EMAIL"]
label2id = {l: i for i, l in enumerate(LABELS)}
id2label = {i: l for i, l in enumerate(LABELS)}

# Load model for token classification
model, tokenizer = load_token_classification(
    "answerdotai/ModernBERT-large",
    train=True,
    num_labels=len(LABELS),
    id2label=id2label,
)

# Dataset with "tokens" and "labels" columns
train_dataset = Dataset.from_dict({
    "tokens": [["Alice", "lives", "in", "Paris"], ["Bob", "emails", "alice@test.com"]],
    "labels": [["B-PER", "O", "O", "O"], ["B-PER", "O", "B-EMAIL", "I-EMAIL"]],
})

# Weighted loss for class imbalance (O class downweighted)
class_weights = [0.2, 1.0, 1.0, 1.0, 1.0]

args = TrainingArgs(
    output_dir="outputs/pii_model",
    num_train_epochs=3,
    batch_size=4,
    gradient_accumulation_steps=8,
    learning_rate=1e-5,
)

trainer = TokenClassificationTrainer(
    model=model,
    tokenizer=tokenizer,
    training_args=args,
    train_dataset=train_dataset,
    id2label=id2label,
    label2id=label2id,
    class_weights=class_weights,
)
trainer.train()

Examples

Sequence Classification

Train a prompt injection guardrail:

python examples/sequence_classification/train.py

Run inference on a trained guardrail:

python examples/sequence_classification/inference.py

Token Classification (PII Detection)

Train a PII detector:

python examples/pii/train_pii.py
python examples/pii/train_pii.py --epochs 5 --lr 5e-5 --o-weight 0.1

Run PII inference:

python examples/pii/inference_pii.py --text "My name is Alice and my email is alice@example.com"
python examples/pii/inference_pii.py --file document.txt --model outputs/modernbert_pii/checkpoint-3941

Architecture

mlx_trainer/
├── args.py                              # TrainingArgs configuration
├── collator.py                          # TextClassificationCollator
├── load.py                              # Model loaders (sequence + token classification)
├── modernbert_config.py                 # ModelArgs dataclass
├── modernbert_model.py                  # ModernBERT + SequenceClassification head
├── token_classification_collator.py     # BIO label alignment via word_ids()
├── token_classification_model.py        # ModernBERT + TokenClassification head
├── token_classification_trainer.py      # Trainer with entity-level eval (seqeval)
├── tokenizer_utils.py                   # HF tokenizer wrapper
├── trainer.py                           # Base trainer (gradient accumulation, checkpointing)
├── train.py                             # Sequence classification training script
└── inference.py                         # Sequence classification inference script

Key Features

  • Pure MLX — No PyTorch dependency. Runs on Metal GPU.
  • fp16 training — 50% memory reduction, no accuracy loss.
  • Gradient checkpointing — Trade compute for memory on long sequences.
  • Entity-level evaluation — Uses seqeval for proper NER metrics (P/R/F1 per entity type).
  • Weighted cross-entropy — Handle class imbalance (e.g., 95% "O" tokens in NER).
  • Subword label alignment — Correctly maps BIO labels through tokenizer subwords via word_ids().
  • Checkpoint rotation — Automatically keeps only N most recent checkpoints.

Training Tips

Class Imbalance (Token Classification)

PII datasets are heavily imbalanced (~95% "O" tokens). Use weighted loss:

# Downweight O class to 0.2, keep entity classes at 1.0
class_weights = [0.2] + [1.0] * (num_labels - 1)

Gradient Checkpointing

Enable for long sequences to reduce memory:

args = TrainingArgs(..., grad_checkpoint=True)

Resume Training

Resume from a checkpoint:

args = TrainingArgs(
    ...,
    resume_from_checkpoint="outputs/my_model/checkpoint-1000",
)

Learning Rate

  • Sequence classification: 1e-5 to 5e-5
  • Token classification: 1e-5 to 5e-5
  • Use warmup: warmup_ratio=0.1 (10% of steps)

License

MIT

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

mlx_modernbert-0.1.0.tar.gz (138.7 kB view details)

Uploaded Source

Built Distribution

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

mlx_modernbert-0.1.0-py3-none-any.whl (22.6 kB view details)

Uploaded Python 3

File details

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

File metadata

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

File hashes

Hashes for mlx_modernbert-0.1.0.tar.gz
Algorithm Hash digest
SHA256 64ac341ccd4a76fdb4ac531ff3fc84e1382c2f566573e19e4559fc9ccf04e117
MD5 1aab57fef293256f28852a49bcd45be4
BLAKE2b-256 624a5eefaa787d72366dc6ee36a0331a4135204a2377f035fe8a1fee7a57dd1f

See more details on using hashes here.

File details

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

File metadata

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

File hashes

Hashes for mlx_modernbert-0.1.0-py3-none-any.whl
Algorithm Hash digest
SHA256 475f4daf147981a8a01d734852cf8463ae7f2c7a713b79a9483e3ac5779a3480
MD5 f157dc4ed5cb8d551cbc6c503f27cff1
BLAKE2b-256 e13579a3fb4ef4ad40c7bb3461ccf663dbae35326425a4a5ca726a7089c16161

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